Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
"""
|
||||
AI 知识库模块
|
||||
"""
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,45 @@
|
||||
"""
|
||||
文档分块策略
|
||||
"""
|
||||
from .base import BaseChunker, ChunkResult
|
||||
from .recursive import RecursiveChunker
|
||||
from .markdown import MarkdownChunker
|
||||
from .fixed import FixedChunker
|
||||
from .qa_chunker import QAChunker
|
||||
from .sentence import SentenceChunker
|
||||
from .auto import AutoChunker
|
||||
|
||||
__all__ = [
|
||||
'BaseChunker',
|
||||
'ChunkResult',
|
||||
'RecursiveChunker',
|
||||
'MarkdownChunker',
|
||||
'FixedChunker',
|
||||
'QAChunker',
|
||||
'SentenceChunker',
|
||||
'AutoChunker',
|
||||
'get_chunker',
|
||||
]
|
||||
|
||||
|
||||
def get_chunker(strategy: str, chunk_size: int = 500, chunk_overlap: int = 50, separator: str = None, **kwargs) -> BaseChunker:
|
||||
"""
|
||||
根据策略名称获取分块器实例
|
||||
|
||||
Args:
|
||||
strategy: 分块策略名称(recursive/markdown/fixed/qa/sentence/auto)
|
||||
chunk_size: 分块大小
|
||||
chunk_overlap: 分块重叠
|
||||
separator: 自定义分隔符
|
||||
**kwargs: 额外参数(如 QAChunker 的 llm_caller)
|
||||
"""
|
||||
chunkers = {
|
||||
'recursive': RecursiveChunker,
|
||||
'markdown': MarkdownChunker,
|
||||
'fixed': FixedChunker,
|
||||
'qa': QAChunker,
|
||||
'sentence': SentenceChunker,
|
||||
'auto': AutoChunker,
|
||||
}
|
||||
chunker_cls = chunkers.get(strategy, RecursiveChunker)
|
||||
return chunker_cls(chunk_size=chunk_size, chunk_overlap=chunk_overlap, separator=separator, **kwargs)
|
||||
@@ -0,0 +1,90 @@
|
||||
"""
|
||||
自动分块策略
|
||||
|
||||
根据文件类型自动选择最佳分块器。
|
||||
参考 Dify 的 auto 分块模式。
|
||||
"""
|
||||
import logging
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 文件类型 → 推荐分块策略
|
||||
FILE_TYPE_STRATEGY_MAP = {
|
||||
# Markdown 文件使用 Markdown 分块器
|
||||
'md': 'markdown',
|
||||
'markdown': 'markdown',
|
||||
# 代码文件使用按句子分块(按行/语句边界)
|
||||
'py': 'sentence',
|
||||
'js': 'sentence',
|
||||
'ts': 'sentence',
|
||||
'java': 'sentence',
|
||||
'go': 'sentence',
|
||||
'rs': 'sentence',
|
||||
'c': 'sentence',
|
||||
'cpp': 'sentence',
|
||||
'h': 'sentence',
|
||||
# 纯文本使用按句子分块
|
||||
'txt': 'sentence',
|
||||
# CSV/Excel 使用固定大小(表格数据按行分割更合理)
|
||||
'csv': 'fixed',
|
||||
'xlsx': 'fixed',
|
||||
'xls': 'fixed',
|
||||
# HTML 使用 Markdown 分块器(HTML 结构类似)
|
||||
'html': 'markdown',
|
||||
'htm': 'markdown',
|
||||
# 其他文档类型使用递归分块
|
||||
'pdf': 'recursive',
|
||||
'docx': 'recursive',
|
||||
'doc': 'recursive',
|
||||
'pptx': 'recursive',
|
||||
'ppt': 'recursive',
|
||||
}
|
||||
|
||||
|
||||
class AutoChunker(BaseChunker):
|
||||
"""
|
||||
自动分块器
|
||||
|
||||
根据文档的文件类型自动选择最佳分块策略。
|
||||
metadata 中需要包含 'file_type' 字段。
|
||||
"""
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""自动选择分块策略并执行"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
metadata = metadata or {}
|
||||
file_type = metadata.get('file_type', '').lower().lstrip('.')
|
||||
|
||||
# 根据文件类型选择策略
|
||||
strategy = FILE_TYPE_STRATEGY_MAP.get(file_type, 'recursive')
|
||||
|
||||
logger.info(f"AutoChunker: file_type={file_type} → strategy={strategy}")
|
||||
|
||||
# 动态创建对应的分块器
|
||||
chunker = self._get_chunker(strategy)
|
||||
return chunker.chunk(text, metadata)
|
||||
|
||||
def _get_chunker(self, strategy: str) -> BaseChunker:
|
||||
"""获取对应策略的分块器实例"""
|
||||
from .recursive import RecursiveChunker
|
||||
from .markdown import MarkdownChunker
|
||||
from .fixed import FixedChunker
|
||||
from .sentence import SentenceChunker
|
||||
|
||||
chunkers = {
|
||||
'recursive': RecursiveChunker,
|
||||
'markdown': MarkdownChunker,
|
||||
'fixed': FixedChunker,
|
||||
'sentence': SentenceChunker,
|
||||
}
|
||||
cls = chunkers.get(strategy, RecursiveChunker)
|
||||
return cls(
|
||||
chunk_size=self.chunk_size,
|
||||
chunk_overlap=self.chunk_overlap,
|
||||
separator=self.separator,
|
||||
)
|
||||
@@ -0,0 +1,92 @@
|
||||
"""
|
||||
分块策略基类
|
||||
"""
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChunkResult:
|
||||
"""分块结果"""
|
||||
content: str
|
||||
position: int = 0
|
||||
char_count: int = 0
|
||||
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
if not self.char_count:
|
||||
self.char_count = len(self.content)
|
||||
|
||||
|
||||
class BaseChunker(ABC):
|
||||
"""
|
||||
分块策略基类
|
||||
|
||||
所有分块策略必须继承此类并实现 chunk 方法
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
chunk_size: int = 500,
|
||||
chunk_overlap: int = 50,
|
||||
separator: Optional[str] = None,
|
||||
):
|
||||
self.chunk_size = chunk_size
|
||||
self.chunk_overlap = chunk_overlap
|
||||
self.separator = separator
|
||||
|
||||
@abstractmethod
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""
|
||||
将文本分块
|
||||
|
||||
Args:
|
||||
text: 原始文本
|
||||
metadata: 文档元数据
|
||||
|
||||
Returns:
|
||||
分块结果列表
|
||||
"""
|
||||
pass
|
||||
|
||||
def _clean_text(self, text: str) -> str:
|
||||
"""清理文本:去除多余空白"""
|
||||
import re
|
||||
# 合并连续空行为单个空行
|
||||
text = re.sub(r'\n{3,}', '\n\n', text)
|
||||
# 去除行尾空白
|
||||
text = '\n'.join(line.rstrip() for line in text.split('\n'))
|
||||
return text.strip()
|
||||
|
||||
def _merge_small_chunks(self, chunks: List[str], min_size: int = 50) -> List[str]:
|
||||
"""合并过小的分块"""
|
||||
if not chunks:
|
||||
return []
|
||||
|
||||
merged = []
|
||||
buffer = ""
|
||||
|
||||
for chunk in chunks:
|
||||
if not chunk.strip():
|
||||
continue
|
||||
if buffer and len(buffer) + len(chunk) <= self.chunk_size:
|
||||
buffer = buffer + "\n" + chunk
|
||||
elif buffer and len(buffer) < min_size:
|
||||
buffer = buffer + "\n" + chunk
|
||||
else:
|
||||
if buffer:
|
||||
merged.append(buffer)
|
||||
buffer = chunk
|
||||
|
||||
if buffer:
|
||||
# 最后一个 buffer 如果太小,合并到前一个
|
||||
if merged and len(buffer) < min_size:
|
||||
merged[-1] = merged[-1] + "\n" + buffer
|
||||
else:
|
||||
merged.append(buffer)
|
||||
|
||||
return merged
|
||||
@@ -0,0 +1,57 @@
|
||||
"""
|
||||
固定大小分块策略
|
||||
|
||||
按固定字符数分割文本,最简单的分块方式
|
||||
"""
|
||||
import logging
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class FixedChunker(BaseChunker):
|
||||
"""
|
||||
固定大小分块器
|
||||
|
||||
按固定字符数分割文本,相邻分块之间有 overlap 重叠
|
||||
"""
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""固定大小分块"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
text = self._clean_text(text)
|
||||
metadata = metadata or {}
|
||||
|
||||
if len(text) <= self.chunk_size:
|
||||
return [ChunkResult(
|
||||
content=text,
|
||||
position=0,
|
||||
metadata={**metadata},
|
||||
)]
|
||||
|
||||
chunks = []
|
||||
start = 0
|
||||
position = 0
|
||||
step = self.chunk_size - self.chunk_overlap
|
||||
|
||||
while start < len(text):
|
||||
end = min(start + self.chunk_size, len(text))
|
||||
chunk_text = text[start:end].strip()
|
||||
|
||||
if chunk_text:
|
||||
chunks.append(ChunkResult(
|
||||
content=chunk_text,
|
||||
position=position,
|
||||
metadata={**metadata},
|
||||
))
|
||||
position += 1
|
||||
|
||||
start += step
|
||||
if step <= 0:
|
||||
break
|
||||
|
||||
return chunks
|
||||
@@ -0,0 +1,172 @@
|
||||
"""
|
||||
Markdown 结构化分块策略
|
||||
|
||||
按 Markdown 标题层级分割文档,保留文档结构信息
|
||||
"""
|
||||
import re
|
||||
import logging
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MarkdownChunker(BaseChunker):
|
||||
"""
|
||||
Markdown 分块器
|
||||
|
||||
按标题层级分割 Markdown 文档,每个标题下的内容作为一个分块
|
||||
如果单个标题下的内容超过 chunk_size,则使用递归分割
|
||||
"""
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""按 Markdown 标题分块"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
text = self._clean_text(text)
|
||||
metadata = metadata or {}
|
||||
|
||||
# 按标题分割
|
||||
sections = self._split_by_headers(text)
|
||||
|
||||
# 处理每个 section
|
||||
raw_chunks = []
|
||||
for section in sections:
|
||||
header = section.get('header', '')
|
||||
content = section.get('content', '')
|
||||
level = section.get('level', 0)
|
||||
|
||||
if not content.strip():
|
||||
continue
|
||||
|
||||
# 组合标题和内容
|
||||
full_text = f"{header}\n{content}" if header else content
|
||||
|
||||
if len(full_text) <= self.chunk_size:
|
||||
raw_chunks.append({
|
||||
'content': full_text.strip(),
|
||||
'metadata': {
|
||||
**metadata,
|
||||
'header': header,
|
||||
'header_level': level,
|
||||
}
|
||||
})
|
||||
else:
|
||||
# 内容超长,递归分割
|
||||
sub_chunks = self._split_long_section(content, header)
|
||||
for i, sub in enumerate(sub_chunks):
|
||||
raw_chunks.append({
|
||||
'content': sub.strip(),
|
||||
'metadata': {
|
||||
**metadata,
|
||||
'header': header,
|
||||
'header_level': level,
|
||||
'sub_chunk': i,
|
||||
}
|
||||
})
|
||||
|
||||
# 合并过小的分块
|
||||
merged = self._merge_small_section_chunks(raw_chunks)
|
||||
|
||||
# 构建结果
|
||||
results = []
|
||||
for i, item in enumerate(merged):
|
||||
if item['content'].strip():
|
||||
results.append(ChunkResult(
|
||||
content=item['content'],
|
||||
position=i,
|
||||
metadata=item.get('metadata', {}),
|
||||
))
|
||||
|
||||
return results
|
||||
|
||||
def _split_by_headers(self, text: str) -> List[Dict[str, Any]]:
|
||||
"""按 Markdown 标题分割"""
|
||||
# 匹配 Markdown 标题: # Title, ## Title, ### Title 等
|
||||
header_pattern = re.compile(r'^(#{1,6})\s+(.+)$', re.MULTILINE)
|
||||
|
||||
sections = []
|
||||
last_end = 0
|
||||
last_header = ''
|
||||
last_level = 0
|
||||
|
||||
for match in header_pattern.finditer(text):
|
||||
# 保存前一个 section 的内容
|
||||
if last_end > 0 or match.start() > 0:
|
||||
content = text[last_end:match.start()]
|
||||
if content.strip() or last_header:
|
||||
sections.append({
|
||||
'header': last_header,
|
||||
'content': content.strip(),
|
||||
'level': last_level,
|
||||
})
|
||||
|
||||
last_header = match.group(0)
|
||||
last_level = len(match.group(1))
|
||||
last_end = match.end()
|
||||
|
||||
# 最后一个 section
|
||||
remaining = text[last_end:]
|
||||
if remaining.strip() or last_header:
|
||||
sections.append({
|
||||
'header': last_header,
|
||||
'content': remaining.strip(),
|
||||
'level': last_level,
|
||||
})
|
||||
|
||||
# 如果没有找到任何标题,整个文本作为一个 section
|
||||
if not sections:
|
||||
sections.append({
|
||||
'header': '',
|
||||
'content': text.strip(),
|
||||
'level': 0,
|
||||
})
|
||||
|
||||
return sections
|
||||
|
||||
def _split_long_section(self, content: str, header: str = '') -> List[str]:
|
||||
"""分割超长的 section 内容"""
|
||||
from .recursive import RecursiveChunker
|
||||
|
||||
chunker = RecursiveChunker(
|
||||
chunk_size=self.chunk_size,
|
||||
chunk_overlap=self.chunk_overlap,
|
||||
)
|
||||
results = chunker.chunk(content)
|
||||
|
||||
chunks = []
|
||||
for i, result in enumerate(results):
|
||||
# 第一个分块带上标题
|
||||
if i == 0 and header:
|
||||
chunks.append(f"{header}\n{result.content}")
|
||||
else:
|
||||
chunks.append(result.content)
|
||||
|
||||
return chunks if chunks else [content]
|
||||
|
||||
def _merge_small_section_chunks(self, chunks: List[Dict], min_size: int = 80) -> List[Dict]:
|
||||
"""合并过小的 section 分块"""
|
||||
if not chunks:
|
||||
return []
|
||||
|
||||
merged = []
|
||||
buffer = None
|
||||
|
||||
for chunk in chunks:
|
||||
if buffer is None:
|
||||
buffer = chunk
|
||||
elif len(buffer['content']) < min_size and len(buffer['content']) + len(chunk['content']) <= self.chunk_size:
|
||||
buffer['content'] = buffer['content'] + "\n\n" + chunk['content']
|
||||
else:
|
||||
merged.append(buffer)
|
||||
buffer = chunk
|
||||
|
||||
if buffer:
|
||||
if merged and len(buffer['content']) < min_size:
|
||||
merged[-1]['content'] = merged[-1]['content'] + "\n\n" + buffer['content']
|
||||
else:
|
||||
merged.append(buffer)
|
||||
|
||||
return merged
|
||||
@@ -0,0 +1,197 @@
|
||||
"""
|
||||
Q&A 自动拆分分块策略
|
||||
|
||||
使用 LLM 将文档内容自动拆分为问答对。
|
||||
每个分段的 content 存储 question,metadata 中存储 answer。
|
||||
检索时用 question 做向量匹配,返回 answer 作为上下文。
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Q&A 拆分的系统提示词
|
||||
QA_SYSTEM_PROMPT = """你是一个专业的知识库问答对生成助手。请根据给定的文本内容,生成高质量的问答对(Q&A pairs)。
|
||||
|
||||
要求:
|
||||
1. 问题应该是用户可能会问的自然语言问题
|
||||
2. 答案应该准确、完整,直接来源于原文
|
||||
3. 每个问答对应该覆盖文本中的一个独立知识点
|
||||
4. 问题要具体明确,避免过于宽泛
|
||||
5. 答案要简洁但完整,包含必要的上下文
|
||||
|
||||
请以 JSON 数组格式输出,每个元素包含 question 和 answer 字段:
|
||||
```json
|
||||
[
|
||||
{"question": "问题1", "answer": "答案1"},
|
||||
{"question": "问题2", "answer": "答案2"}
|
||||
]
|
||||
```
|
||||
|
||||
只输出 JSON 数组,不要输出其他内容。"""
|
||||
|
||||
|
||||
class QAChunker(BaseChunker):
|
||||
"""
|
||||
Q&A 自动拆分分块器
|
||||
|
||||
使用 LLM 将文本拆分为问答对。
|
||||
需要在初始化时传入 LLM 调用函数。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
chunk_size: int = 500,
|
||||
chunk_overlap: int = 50,
|
||||
separator: Optional[str] = None,
|
||||
llm_caller: Optional[Any] = None,
|
||||
):
|
||||
super().__init__(chunk_size, chunk_overlap, separator)
|
||||
self._llm_caller = llm_caller
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""
|
||||
同步分块(Q&A 模式不支持同步调用,返回空列表)
|
||||
请使用 chunk_async 方法。
|
||||
"""
|
||||
logger.warning("QAChunker.chunk() 不支持同步调用,请使用 chunk_async()")
|
||||
return []
|
||||
|
||||
async def chunk_async(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""
|
||||
异步分块:使用 LLM 将文本拆分为 Q&A 对
|
||||
|
||||
Args:
|
||||
text: 原始文本
|
||||
metadata: 文档元数据
|
||||
|
||||
Returns:
|
||||
分块结果列表,每个 ChunkResult 的 content 为 question,
|
||||
metadata 中包含 answer 和 chunk_mode='qa'
|
||||
"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
metadata = metadata or {}
|
||||
text = self._clean_text(text)
|
||||
|
||||
# 如果文本太长,先按段落粗分再逐段生成 Q&A
|
||||
max_input_size = self.chunk_size * 8 # LLM 输入上限
|
||||
if len(text) > max_input_size:
|
||||
segments = self._split_for_qa(text, max_input_size)
|
||||
else:
|
||||
segments = [text]
|
||||
|
||||
all_results = []
|
||||
position = 0
|
||||
|
||||
for segment in segments:
|
||||
qa_pairs = await self._generate_qa_pairs(segment)
|
||||
for qa in qa_pairs:
|
||||
question = qa.get('question', '').strip()
|
||||
answer = qa.get('answer', '').strip()
|
||||
if not question or not answer:
|
||||
continue
|
||||
|
||||
all_results.append(ChunkResult(
|
||||
content=question,
|
||||
position=position,
|
||||
metadata={
|
||||
**metadata,
|
||||
'answer': answer,
|
||||
'chunk_mode': 'qa',
|
||||
},
|
||||
))
|
||||
position += 1
|
||||
|
||||
logger.info(f"Q&A 拆分完成: {len(all_results)} 个问答对")
|
||||
return all_results
|
||||
|
||||
async def _generate_qa_pairs(self, text: str) -> List[Dict[str, str]]:
|
||||
"""调用 LLM 生成 Q&A 对"""
|
||||
if not self._llm_caller:
|
||||
logger.error("QAChunker: 未配置 LLM 调用函数")
|
||||
return []
|
||||
|
||||
try:
|
||||
user_prompt = f"请根据以下文本生成问答对:\n\n{text}"
|
||||
|
||||
response_text = await self._llm_caller(
|
||||
system_prompt=QA_SYSTEM_PROMPT,
|
||||
user_prompt=user_prompt,
|
||||
)
|
||||
|
||||
if not response_text:
|
||||
return []
|
||||
|
||||
# 解析 JSON 响应
|
||||
return self._parse_qa_response(response_text)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Q&A 生成失败: {e}")
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _parse_qa_response(response_text: str) -> List[Dict[str, str]]:
|
||||
"""解析 LLM 返回的 Q&A JSON"""
|
||||
try:
|
||||
# 尝试直接解析
|
||||
result = json.loads(response_text)
|
||||
if isinstance(result, list):
|
||||
return [
|
||||
item for item in result
|
||||
if isinstance(item, dict) and 'question' in item and 'answer' in item
|
||||
]
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# 尝试提取 JSON 代码块
|
||||
import re
|
||||
json_match = re.search(r'```(?:json)?\s*\n?(.*?)\n?```', response_text, re.DOTALL)
|
||||
if json_match:
|
||||
try:
|
||||
result = json.loads(json_match.group(1))
|
||||
if isinstance(result, list):
|
||||
return [
|
||||
item for item in result
|
||||
if isinstance(item, dict) and 'question' in item and 'answer' in item
|
||||
]
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# 尝试找到 [ ... ] 部分
|
||||
bracket_match = re.search(r'\[.*\]', response_text, re.DOTALL)
|
||||
if bracket_match:
|
||||
try:
|
||||
result = json.loads(bracket_match.group(0))
|
||||
if isinstance(result, list):
|
||||
return [
|
||||
item for item in result
|
||||
if isinstance(item, dict) and 'question' in item and 'answer' in item
|
||||
]
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
logger.warning(f"无法解析 Q&A 响应: {response_text[:200]}")
|
||||
return []
|
||||
|
||||
def _split_for_qa(self, text: str, max_size: int) -> List[str]:
|
||||
"""将长文本按段落分割为适合 LLM 处理的片段"""
|
||||
paragraphs = text.split('\n\n')
|
||||
segments = []
|
||||
current = ""
|
||||
|
||||
for para in paragraphs:
|
||||
if current and len(current) + len(para) + 2 > max_size:
|
||||
segments.append(current.strip())
|
||||
current = para
|
||||
else:
|
||||
current = current + "\n\n" + para if current else para
|
||||
|
||||
if current.strip():
|
||||
segments.append(current.strip())
|
||||
|
||||
return segments
|
||||
@@ -0,0 +1,160 @@
|
||||
"""
|
||||
递归字符分块策略
|
||||
|
||||
最常用的分块策略,按照分隔符层级递归分割文本
|
||||
优先按段落 → 句子 → 字符的顺序分割
|
||||
"""
|
||||
import logging
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 默认分隔符层级(从大到小)
|
||||
DEFAULT_SEPARATORS = [
|
||||
"\n\n", # 段落
|
||||
"\n", # 换行
|
||||
"。", # 中文句号
|
||||
"!", # 中文感叹号
|
||||
"?", # 中文问号
|
||||
";", # 中文分号
|
||||
". ", # 英文句号
|
||||
"! ", # 英文感叹号
|
||||
"? ", # 英文问号
|
||||
"; ", # 英文分号
|
||||
",", # 中文逗号
|
||||
", ", # 英文逗号
|
||||
" ", # 空格
|
||||
"", # 逐字符
|
||||
]
|
||||
|
||||
|
||||
class RecursiveChunker(BaseChunker):
|
||||
"""
|
||||
递归字符分块器
|
||||
|
||||
按分隔符层级递归分割文本,确保每个分块不超过 chunk_size,
|
||||
相邻分块之间有 chunk_overlap 的重叠
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
chunk_size: int = 500,
|
||||
chunk_overlap: int = 50,
|
||||
separator: Optional[str] = None,
|
||||
separators: Optional[List[str]] = None,
|
||||
):
|
||||
super().__init__(chunk_size, chunk_overlap, separator)
|
||||
if separator:
|
||||
self.separators = [separator] + DEFAULT_SEPARATORS
|
||||
elif separators:
|
||||
self.separators = separators
|
||||
else:
|
||||
self.separators = DEFAULT_SEPARATORS
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""递归分块"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
text = self._clean_text(text)
|
||||
metadata = metadata or {}
|
||||
|
||||
# 递归分割
|
||||
raw_chunks = self._recursive_split(text, self.separators)
|
||||
|
||||
# 合并过小的分块
|
||||
raw_chunks = self._merge_small_chunks(raw_chunks)
|
||||
|
||||
# 添加重叠
|
||||
chunks_with_overlap = self._add_overlap(raw_chunks)
|
||||
|
||||
# 构建结果
|
||||
results = []
|
||||
for i, content in enumerate(chunks_with_overlap):
|
||||
if content.strip():
|
||||
results.append(ChunkResult(
|
||||
content=content.strip(),
|
||||
position=i,
|
||||
metadata={**metadata},
|
||||
))
|
||||
|
||||
return results
|
||||
|
||||
def _recursive_split(self, text: str, separators: List[str]) -> List[str]:
|
||||
"""递归分割文本"""
|
||||
if len(text) <= self.chunk_size:
|
||||
return [text] if text.strip() else []
|
||||
|
||||
# 找到合适的分隔符
|
||||
separator = ""
|
||||
for sep in separators:
|
||||
if sep == "":
|
||||
separator = sep
|
||||
break
|
||||
if sep in text:
|
||||
separator = sep
|
||||
break
|
||||
|
||||
# 按分隔符分割
|
||||
if separator:
|
||||
splits = text.split(separator)
|
||||
else:
|
||||
# 逐字符分割
|
||||
splits = list(text)
|
||||
|
||||
# 合并分割结果,确保不超过 chunk_size
|
||||
chunks = []
|
||||
current = ""
|
||||
|
||||
for split in splits:
|
||||
piece = split if not separator else split
|
||||
test_piece = current + separator + piece if current else piece
|
||||
|
||||
if len(test_piece) <= self.chunk_size:
|
||||
current = test_piece
|
||||
else:
|
||||
if current:
|
||||
chunks.append(current)
|
||||
# 如果单个片段超过 chunk_size,递归处理
|
||||
if len(piece) > self.chunk_size:
|
||||
remaining_separators = separators[separators.index(separator) + 1:] if separator in separators else separators[1:]
|
||||
if remaining_separators:
|
||||
sub_chunks = self._recursive_split(piece, remaining_separators)
|
||||
chunks.extend(sub_chunks)
|
||||
current = ""
|
||||
else:
|
||||
# 没有更小的分隔符了,强制截断
|
||||
for j in range(0, len(piece), self.chunk_size):
|
||||
chunks.append(piece[j:j + self.chunk_size])
|
||||
current = ""
|
||||
else:
|
||||
current = piece
|
||||
|
||||
if current:
|
||||
chunks.append(current)
|
||||
|
||||
return chunks
|
||||
|
||||
def _add_overlap(self, chunks: List[str]) -> List[str]:
|
||||
"""为相邻分块添加重叠"""
|
||||
if self.chunk_overlap <= 0 or len(chunks) <= 1:
|
||||
return chunks
|
||||
|
||||
result = []
|
||||
for i, chunk in enumerate(chunks):
|
||||
if i == 0:
|
||||
result.append(chunk)
|
||||
else:
|
||||
# 从前一个分块的末尾取 overlap 字符作为前缀
|
||||
prev = chunks[i - 1]
|
||||
overlap_text = prev[-self.chunk_overlap:] if len(prev) > self.chunk_overlap else prev
|
||||
# 确保合并后不超过 chunk_size 太多
|
||||
combined = overlap_text + "\n" + chunk
|
||||
if len(combined) <= self.chunk_size * 1.2:
|
||||
result.append(combined)
|
||||
else:
|
||||
result.append(chunk)
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,105 @@
|
||||
"""
|
||||
按句子分块策略
|
||||
|
||||
按句号/问号/感叹号等句子边界分割文本,
|
||||
然后将小句子合并到不超过 chunk_size 的分块中。
|
||||
参考 Dify 的 sentence 分块模式。
|
||||
"""
|
||||
import logging
|
||||
import re
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
from .base import BaseChunker, ChunkResult
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 句子分隔符正则(中英文句号/问号/感叹号/分号)
|
||||
SENTENCE_PATTERN = re.compile(
|
||||
r'(?<=[。!?;.!?;])\s*'
|
||||
)
|
||||
|
||||
|
||||
class SentenceChunker(BaseChunker):
|
||||
"""
|
||||
按句子分块器
|
||||
|
||||
先按句子边界分割文本,再将相邻句子合并为不超过 chunk_size 的分块。
|
||||
保证每个分块都是完整句子的组合,不会在句子中间截断。
|
||||
"""
|
||||
|
||||
def chunk(self, text: str, metadata: Dict[str, Any] = None) -> List[ChunkResult]:
|
||||
"""按句子分块"""
|
||||
if not text or not text.strip():
|
||||
return []
|
||||
|
||||
text = self._clean_text(text)
|
||||
metadata = metadata or {}
|
||||
|
||||
# 按句子边界分割
|
||||
sentences = SENTENCE_PATTERN.split(text)
|
||||
sentences = [s.strip() for s in sentences if s.strip()]
|
||||
|
||||
if not sentences:
|
||||
return [ChunkResult(content=text, position=0, metadata={**metadata})]
|
||||
|
||||
# 合并句子为分块(不超过 chunk_size)
|
||||
chunks = []
|
||||
current = ""
|
||||
position = 0
|
||||
|
||||
for sentence in sentences:
|
||||
# 如果单个句子就超过 chunk_size,强制作为独立分块
|
||||
if len(sentence) > self.chunk_size:
|
||||
if current:
|
||||
chunks.append(current)
|
||||
current = ""
|
||||
chunks.append(sentence)
|
||||
continue
|
||||
|
||||
test = current + sentence if not current else current + " " + sentence
|
||||
if len(test) <= self.chunk_size:
|
||||
current = test
|
||||
else:
|
||||
if current:
|
||||
chunks.append(current)
|
||||
current = sentence
|
||||
|
||||
if current:
|
||||
chunks.append(current)
|
||||
|
||||
# 合并过小的分块
|
||||
chunks = self._merge_small_chunks(chunks)
|
||||
|
||||
# 添加重叠
|
||||
if self.chunk_overlap > 0 and len(chunks) > 1:
|
||||
chunks = self._add_sentence_overlap(chunks)
|
||||
|
||||
# 构建结果
|
||||
results = []
|
||||
for i, content in enumerate(chunks):
|
||||
if content.strip():
|
||||
results.append(ChunkResult(
|
||||
content=content.strip(),
|
||||
position=i,
|
||||
metadata={**metadata},
|
||||
))
|
||||
|
||||
return results
|
||||
|
||||
def _add_sentence_overlap(self, chunks: List[str]) -> List[str]:
|
||||
"""为相邻分块添加句子级重叠"""
|
||||
result = [chunks[0]]
|
||||
for i in range(1, len(chunks)):
|
||||
prev = chunks[i - 1]
|
||||
# 从前一个分块取最后一个句子作为重叠
|
||||
prev_sentences = SENTENCE_PATTERN.split(prev)
|
||||
prev_sentences = [s.strip() for s in prev_sentences if s.strip()]
|
||||
if prev_sentences:
|
||||
overlap = prev_sentences[-1]
|
||||
if len(overlap) <= self.chunk_overlap:
|
||||
combined = overlap + " " + chunks[i]
|
||||
if len(combined) <= self.chunk_size * 1.2:
|
||||
result.append(combined)
|
||||
continue
|
||||
result.append(chunks[i])
|
||||
return result
|
||||
@@ -0,0 +1,16 @@
|
||||
"""
|
||||
知识库数据模型
|
||||
"""
|
||||
from .knowledge_base_model import KnowledgeBase
|
||||
from .document_model import KnowledgeDocument
|
||||
from .segment_model import KnowledgeSegment
|
||||
from .annotation_model import KnowledgeAnnotation
|
||||
from .retrieval_log_model import KnowledgeRetrievalLog
|
||||
|
||||
__all__ = [
|
||||
'KnowledgeBase',
|
||||
'KnowledgeDocument',
|
||||
'KnowledgeSegment',
|
||||
'KnowledgeAnnotation',
|
||||
'KnowledgeRetrievalLog',
|
||||
]
|
||||
@@ -0,0 +1,34 @@
|
||||
"""
|
||||
知识库标注模型(Q&A 对)
|
||||
|
||||
手动添加的高优先级问答对,检索时优先匹配。
|
||||
"""
|
||||
from sqlalchemy import Column, String, Text, Integer, Boolean, Index
|
||||
|
||||
from app.base_model import BaseModel
|
||||
|
||||
|
||||
class KnowledgeAnnotation(BaseModel):
|
||||
"""
|
||||
知识库标注(Q&A 对)
|
||||
|
||||
用户手动添加的问答对,检索时优先匹配 question,返回 answer。
|
||||
"""
|
||||
__tablename__ = "ai_knowledge_annotation"
|
||||
|
||||
knowledge_base_id = Column(String(21), nullable=False, index=True, comment="所属知识库ID(逻辑外键关联ai_knowledge_base)")
|
||||
|
||||
# Q&A 内容
|
||||
question = Column(Text, nullable=False, comment="问题")
|
||||
answer = Column(Text, nullable=False, comment="答案")
|
||||
|
||||
# 向量化状态
|
||||
embedding_status = Column(String(20), default="pending", comment="向量化状态: pending/completed/failed")
|
||||
|
||||
# 状态
|
||||
enabled = Column(Boolean, default=True, comment="是否启用")
|
||||
hit_count = Column(Integer, default=0, comment="命中次数")
|
||||
|
||||
__table_args__ = (
|
||||
Index('idx_annotation_kb_enabled', 'knowledge_base_id', 'enabled'),
|
||||
)
|
||||
@@ -0,0 +1,39 @@
|
||||
"""
|
||||
知识库文档模型
|
||||
"""
|
||||
from sqlalchemy import Column, String, Text, Integer, Boolean, DateTime, BigInteger
|
||||
|
||||
from app.base_model import BaseModel
|
||||
|
||||
|
||||
class KnowledgeDocument(BaseModel):
|
||||
"""
|
||||
知识库文档
|
||||
|
||||
记录上传到知识库的文档信息及处理状态
|
||||
"""
|
||||
__tablename__ = "ai_knowledge_document"
|
||||
|
||||
knowledge_base_id = Column(String(21), nullable=False, index=True, comment="所属知识库ID(逻辑外键关联ai_knowledge_base)")
|
||||
file_id = Column(String(21), nullable=True, comment="关联文件ID(逻辑外键关联core_file_manager)")
|
||||
name = Column(String(255), nullable=False, comment="文档名称")
|
||||
file_type = Column(String(20), nullable=True, comment="文件类型: pdf/docx/txt/md/xlsx/csv/html/pptx")
|
||||
file_size = Column(BigInteger, default=0, comment="文件大小(字节)")
|
||||
content_hash = Column(String(64), nullable=True, index=True, comment="内容MD5(用于去重)")
|
||||
|
||||
# 处理结果
|
||||
segment_count = Column(Integer, default=0, comment="分段数量")
|
||||
token_count = Column(Integer, default=0, comment="Token 总数")
|
||||
char_count = Column(Integer, default=0, comment="字符总数")
|
||||
|
||||
# 处理状态
|
||||
status = Column(String(20), default="pending", index=True, comment="状态: pending/indexing/completed/failed/disabled")
|
||||
error_message = Column(Text, nullable=True, comment="错误信息")
|
||||
indexing_started_at = Column(DateTime, nullable=True, comment="索引开始时间")
|
||||
indexing_completed_at = Column(DateTime, nullable=True, comment="索引完成时间")
|
||||
|
||||
# 去重
|
||||
duplicate_warning = Column(Text, nullable=True, comment="内容重复警告(跨知识库检测)")
|
||||
|
||||
# 是否启用
|
||||
enabled = Column(Boolean, default=True, comment="是否启用(禁用后不参与检索)")
|
||||
@@ -0,0 +1,55 @@
|
||||
"""
|
||||
知识库模型
|
||||
"""
|
||||
from sqlalchemy import Column, String, Text, Integer, Float, Boolean, JSON
|
||||
|
||||
from app.base_model import BaseModel
|
||||
|
||||
|
||||
class KnowledgeBase(BaseModel):
|
||||
"""
|
||||
知识库
|
||||
|
||||
管理文档集合,配置分块策略和检索参数
|
||||
"""
|
||||
__tablename__ = "ai_knowledge_base"
|
||||
|
||||
application_id = Column(String(21), nullable=True, index=True, comment="所属应用ID(逻辑外键关联core_application)")
|
||||
is_global = Column(Boolean, default=False, comment="是否在子应用中可见")
|
||||
name = Column(String(100), nullable=False, comment="知识库名称")
|
||||
code = Column(String(100), nullable=False, unique=True, comment="知识库编码")
|
||||
description = Column(Text, nullable=True, comment="描述")
|
||||
icon = Column(String(50), default="", comment="图标")
|
||||
|
||||
# Embedding 配置
|
||||
embedding_model_id = Column(String(21), nullable=True, comment="Embedding 模型ID(逻辑外键关联ai_llm_model)")
|
||||
embedding_dimensions = Column(Integer, default=1536, comment="向量维度")
|
||||
|
||||
# 分块策略
|
||||
chunk_strategy = Column(String(20), default="recursive", comment="分块策略: recursive/semantic/markdown/fixed")
|
||||
chunk_size = Column(Integer, default=500, comment="分块大小(字符数)")
|
||||
chunk_overlap = Column(Integer, default=50, comment="分块重叠大小(字符数)")
|
||||
separator = Column(String(50), nullable=True, comment="自定义分隔符")
|
||||
|
||||
# 检索配置
|
||||
retrieval_mode = Column(String(20), default="hybrid", comment="检索模式: vector/fulltext/hybrid")
|
||||
top_k = Column(Integer, default=5, comment="检索返回数量")
|
||||
score_threshold = Column(Float, default=0.5, comment="相似度阈值(0-1)")
|
||||
rerank_enabled = Column(Boolean, default=False, comment="是否启用重排序")
|
||||
rerank_model_id = Column(String(21), nullable=True, comment="重排序模型ID")
|
||||
retrieval_weight = Column(Float, default=1.0, comment="检索权重(多知识库检索时的加权系数,0.1-10.0)")
|
||||
|
||||
# 预处理规则
|
||||
process_rules = Column(JSON, nullable=True, comment="预处理规则(清洗配置)")
|
||||
|
||||
# 索引模式
|
||||
indexing_technique = Column(String(20), default="high_quality", comment="索引模式: high_quality/economy")
|
||||
|
||||
# 统计
|
||||
document_count = Column(Integer, default=0, comment="文档数量")
|
||||
segment_count = Column(Integer, default=0, comment="分段数量")
|
||||
total_token_count = Column(Integer, default=0, comment="总 Token 数")
|
||||
total_char_count = Column(Integer, default=0, comment="总字符数")
|
||||
|
||||
# 状态
|
||||
status = Column(String(20), default="active", comment="状态: active/disabled")
|
||||
@@ -0,0 +1,41 @@
|
||||
"""
|
||||
知识库检索日志模型
|
||||
|
||||
记录每次检索的查询、结果、耗时等信息,用于分析检索质量。
|
||||
"""
|
||||
from sqlalchemy import Column, String, Text, Integer, Float, JSON, Index
|
||||
|
||||
from app.base_model import BaseModel
|
||||
|
||||
|
||||
class KnowledgeRetrievalLog(BaseModel):
|
||||
"""
|
||||
知识库检索日志
|
||||
|
||||
记录每次检索请求的完整信息,用于检索质量分析和优化。
|
||||
"""
|
||||
__tablename__ = "ai_knowledge_retrieval_log"
|
||||
|
||||
# 检索请求
|
||||
query = Column(Text, nullable=False, comment="查询文本")
|
||||
knowledge_base_ids = Column(JSON, nullable=False, comment="检索的知识库ID列表")
|
||||
retrieval_mode = Column(String(20), default="hybrid", comment="检索模式: vector/fulltext/hybrid")
|
||||
top_k = Column(Integer, default=5, comment="请求的返回数量")
|
||||
score_threshold = Column(Float, default=0.5, comment="相似度阈值")
|
||||
|
||||
# 检索结果
|
||||
result_count = Column(Integer, default=0, comment="实际返回结果数")
|
||||
results = Column(JSON, nullable=True, comment="检索结果摘要(segment_id/score/kb_id)")
|
||||
rerank_applied = Column(String(5), default="false", comment="是否应用了重排序")
|
||||
|
||||
# 性能
|
||||
elapsed_time = Column(Integer, default=0, comment="耗时(毫秒)")
|
||||
|
||||
# 来源
|
||||
source = Column(String(50), nullable=True, comment="调用来源: api/workflow/chat")
|
||||
user_id = Column(String(21), nullable=True, comment="操作用户ID")
|
||||
|
||||
__table_args__ = (
|
||||
Index('idx_retrieval_log_query_time', 'sys_create_datetime'),
|
||||
Index('idx_retrieval_log_user', 'user_id'),
|
||||
)
|
||||
@@ -0,0 +1,47 @@
|
||||
"""
|
||||
知识库分段模型
|
||||
"""
|
||||
from sqlalchemy import Column, String, Text, Integer, Boolean, JSON, Index
|
||||
|
||||
from app.base_model import BaseModel
|
||||
|
||||
|
||||
class KnowledgeSegment(BaseModel):
|
||||
"""
|
||||
知识库分段(Chunk)
|
||||
|
||||
文档经过分块后的最小检索单元
|
||||
向量数据存储在 Qdrant 向量数据库中,此表只存业务数据
|
||||
"""
|
||||
__tablename__ = "ai_knowledge_segment"
|
||||
|
||||
knowledge_base_id = Column(String(21), nullable=False, index=True, comment="所属知识库ID(逻辑外键关联ai_knowledge_base)")
|
||||
document_id = Column(String(21), nullable=False, index=True, comment="所属文档ID(逻辑外键关联ai_knowledge_document)")
|
||||
|
||||
# 内容
|
||||
position = Column(Integer, default=0, comment="在文档中的位置序号")
|
||||
content = Column(Text, nullable=False, comment="分段文本内容")
|
||||
answer = Column(Text, nullable=True, comment="Q&A 模式的答案内容")
|
||||
token_count = Column(Integer, default=0, comment="Token 数")
|
||||
char_count = Column(Integer, default=0, comment="字符数")
|
||||
word_count = Column(Integer, default=0, comment="词数")
|
||||
|
||||
# 元数据
|
||||
page_number = Column(Integer, nullable=True, comment="来源页码(PDF/PPT)")
|
||||
keywords = Column(JSON, nullable=True, comment="关键词列表(用于全文检索增强)")
|
||||
extra_metadata = Column(JSON, nullable=True, comment="元数据(标题/来源等)")
|
||||
|
||||
# 向量化状态(向量数据存在 Qdrant 中,这里只记录状态)
|
||||
embedding_status = Column(String(20), default="pending", comment="向量化状态: pending/completed/failed")
|
||||
|
||||
# 父子分段(Small-to-Big)
|
||||
parent_segment_id = Column(String(21), nullable=True, index=True, comment="父分段ID(逻辑外键,用于 Small-to-Big 检索)")
|
||||
|
||||
# 状态
|
||||
enabled = Column(Boolean, default=True, comment="是否启用(禁用后不参与检索)")
|
||||
hit_count = Column(Integer, default=0, comment="命中次数")
|
||||
|
||||
__table_args__ = (
|
||||
Index('idx_segment_kb_doc', 'knowledge_base_id', 'document_id'),
|
||||
Index('idx_segment_kb_enabled', 'knowledge_base_id', 'enabled'),
|
||||
)
|
||||
@@ -0,0 +1,6 @@
|
||||
"""
|
||||
知识库 Schema
|
||||
"""
|
||||
from .knowledge_base_schema import *
|
||||
from .document_schema import *
|
||||
from .segment_schema import *
|
||||
@@ -0,0 +1,37 @@
|
||||
"""
|
||||
知识库标注 Schema
|
||||
"""
|
||||
from typing import Optional, List
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field, ConfigDict
|
||||
|
||||
from app.base_schema import CSTDatetime
|
||||
|
||||
|
||||
class AnnotationCreateInput(BaseModel):
|
||||
"""创建标注"""
|
||||
question: str = Field(..., min_length=1, description="问题")
|
||||
answer: str = Field(..., min_length=1, description="答案")
|
||||
|
||||
|
||||
class AnnotationUpdateInput(BaseModel):
|
||||
"""更新标注"""
|
||||
question: Optional[str] = Field(None, min_length=1, description="问题")
|
||||
answer: Optional[str] = Field(None, min_length=1, description="答案")
|
||||
enabled: Optional[bool] = Field(None, description="是否启用")
|
||||
|
||||
|
||||
class AnnotationResponse(BaseModel):
|
||||
"""标注输出"""
|
||||
id: str
|
||||
knowledge_base_id: str
|
||||
question: str
|
||||
answer: str
|
||||
embedding_status: str = "pending"
|
||||
enabled: bool = True
|
||||
hit_count: int = 0
|
||||
sys_create_datetime: Optional[CSTDatetime] = None
|
||||
sys_update_datetime: Optional[CSTDatetime] = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
@@ -0,0 +1,62 @@
|
||||
"""
|
||||
知识库文档 Schema
|
||||
"""
|
||||
from typing import Optional, List
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field, ConfigDict
|
||||
|
||||
from app.base_schema import CSTDatetime
|
||||
|
||||
|
||||
class DocumentUploadInput(BaseModel):
|
||||
"""文档上传输入(通过文件管理系统上传后传入 file_id)"""
|
||||
file_id: str = Field(..., description="文件ID(来自文件管理系统)")
|
||||
name: Optional[str] = Field(None, description="文档名称(不传则使用文件名)")
|
||||
|
||||
|
||||
class DocumentBatchUploadInput(BaseModel):
|
||||
"""批量文档上传"""
|
||||
file_ids: List[str] = Field(..., min_length=1, description="文件ID列表")
|
||||
|
||||
|
||||
class DocumentResponse(BaseModel):
|
||||
"""文档输出"""
|
||||
id: str
|
||||
knowledge_base_id: str
|
||||
file_id: Optional[str] = None
|
||||
name: str
|
||||
file_type: str = ""
|
||||
file_size: int = 0
|
||||
content_hash: str = ""
|
||||
segment_count: int = 0
|
||||
token_count: int = 0
|
||||
char_count: int = 0
|
||||
status: str = "pending"
|
||||
error_message: str = ""
|
||||
duplicate_warning: Optional[str] = None
|
||||
enabled: bool = True
|
||||
indexing_started_at: Optional[CSTDatetime] = None
|
||||
indexing_completed_at: Optional[CSTDatetime] = None
|
||||
sys_create_datetime: Optional[CSTDatetime] = None
|
||||
sys_update_datetime: Optional[CSTDatetime] = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class DocumentListResponse(BaseModel):
|
||||
"""文档列表输出"""
|
||||
id: str
|
||||
knowledge_base_id: str
|
||||
file_id: Optional[str] = None
|
||||
name: str
|
||||
file_type: str = ""
|
||||
file_size: int = 0
|
||||
segment_count: int = 0
|
||||
token_count: int = 0
|
||||
status: str = "pending"
|
||||
duplicate_warning: Optional[str] = None
|
||||
enabled: bool = True
|
||||
sys_create_datetime: Optional[CSTDatetime] = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
@@ -0,0 +1,115 @@
|
||||
"""
|
||||
知识库 Schema
|
||||
"""
|
||||
from typing import Optional, List, Dict, Any
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field, ConfigDict
|
||||
|
||||
from app.base_schema import CSTDatetime
|
||||
|
||||
|
||||
class KnowledgeBaseCreate(BaseModel):
|
||||
"""创建知识库"""
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
application_id: Optional[str] = Field(None, description="所属应用ID")
|
||||
is_global: bool = Field(default=False, description="是否在子应用中可见")
|
||||
name: str = Field(..., max_length=100, description="知识库名称")
|
||||
code: str = Field(..., max_length=100, description="知识库编码")
|
||||
description: Optional[str] = Field(None, description="描述")
|
||||
icon: str = Field(default="", description="图标")
|
||||
embedding_model_id: Optional[str] = Field(None, description="Embedding 模型ID")
|
||||
embedding_dimensions: int = Field(default=1536, description="向量维度")
|
||||
chunk_strategy: str = Field(default="recursive", description="分块策略")
|
||||
chunk_size: int = Field(default=500, ge=100, le=4000, description="分块大小")
|
||||
chunk_overlap: int = Field(default=50, ge=0, le=500, description="分块重叠")
|
||||
separator: Optional[str] = Field(None, description="自定义分隔符")
|
||||
retrieval_mode: str = Field(default="hybrid", description="检索模式")
|
||||
top_k: int = Field(default=5, ge=1, le=20, description="检索数量")
|
||||
score_threshold: float = Field(default=0.5, ge=0, le=1, description="相似度阈值")
|
||||
rerank_enabled: bool = Field(default=False, description="是否启用重排序")
|
||||
rerank_model_id: Optional[str] = Field(None, description="重排序模型ID")
|
||||
retrieval_weight: float = Field(default=1.0, ge=0.1, le=10.0, description="检索权重")
|
||||
process_rules: Optional[Dict[str, Any]] = Field(None, description="预处理规则")
|
||||
indexing_technique: str = Field(default="high_quality", description="索引模式: high_quality/economy")
|
||||
|
||||
|
||||
class KnowledgeBaseUpdate(BaseModel):
|
||||
"""更新知识库"""
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
name: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
icon: Optional[str] = None
|
||||
embedding_model_id: Optional[str] = None
|
||||
embedding_dimensions: Optional[int] = None
|
||||
chunk_strategy: Optional[str] = None
|
||||
chunk_size: Optional[int] = Field(None, ge=100, le=4000)
|
||||
chunk_overlap: Optional[int] = Field(None, ge=0, le=500)
|
||||
separator: Optional[str] = None
|
||||
retrieval_mode: Optional[str] = None
|
||||
top_k: Optional[int] = Field(None, ge=1, le=20)
|
||||
score_threshold: Optional[float] = Field(None, ge=0, le=1)
|
||||
rerank_enabled: Optional[bool] = None
|
||||
rerank_model_id: Optional[str] = None
|
||||
retrieval_weight: Optional[float] = Field(None, ge=0.1, le=10.0)
|
||||
process_rules: Optional[Dict[str, Any]] = None
|
||||
indexing_technique: Optional[str] = None
|
||||
status: Optional[str] = None
|
||||
is_global: Optional[bool] = None
|
||||
|
||||
|
||||
class KnowledgeBaseResponse(BaseModel):
|
||||
"""知识库详情输出"""
|
||||
id: str
|
||||
application_id: Optional[str] = None
|
||||
is_global: bool = False
|
||||
name: str
|
||||
code: str
|
||||
description: str = ""
|
||||
icon: str = ""
|
||||
embedding_model_id: Optional[str] = None
|
||||
embedding_model_name: str = ""
|
||||
embedding_dimensions: int = 1536
|
||||
chunk_strategy: str = "recursive"
|
||||
chunk_size: int = 500
|
||||
chunk_overlap: int = 50
|
||||
separator: Optional[str] = None
|
||||
retrieval_mode: str = "hybrid"
|
||||
top_k: int = 5
|
||||
score_threshold: float = 0.5
|
||||
rerank_enabled: bool = False
|
||||
rerank_model_id: Optional[str] = None
|
||||
retrieval_weight: float = 1.0
|
||||
process_rules: Optional[Dict[str, Any]] = None
|
||||
indexing_technique: str = "high_quality"
|
||||
document_count: int = 0
|
||||
segment_count: int = 0
|
||||
total_token_count: int = 0
|
||||
total_char_count: int = 0
|
||||
status: str = "active"
|
||||
sort: int = 0
|
||||
sys_create_datetime: Optional[CSTDatetime] = None
|
||||
sys_update_datetime: Optional[CSTDatetime] = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True, protected_namespaces=())
|
||||
|
||||
|
||||
class KnowledgeBaseListResponse(BaseModel):
|
||||
"""知识库列表输出"""
|
||||
id: str
|
||||
application_id: Optional[str] = None
|
||||
application_name: str = ""
|
||||
is_global: bool = False
|
||||
name: str
|
||||
code: str
|
||||
description: str = ""
|
||||
icon: str = ""
|
||||
embedding_model_name: str = ""
|
||||
document_count: int = 0
|
||||
segment_count: int = 0
|
||||
status: str = "active"
|
||||
sys_create_datetime: Optional[CSTDatetime] = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True, protected_namespaces=())
|
||||
@@ -0,0 +1,28 @@
|
||||
"""
|
||||
检索日志 Schema
|
||||
"""
|
||||
from typing import Optional, List, Dict, Any
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field, ConfigDict
|
||||
|
||||
from app.base_schema import CSTDatetime
|
||||
|
||||
|
||||
class RetrievalLogResponse(BaseModel):
|
||||
"""检索日志输出"""
|
||||
id: str
|
||||
query: str
|
||||
knowledge_base_ids: List[str] = []
|
||||
retrieval_mode: str = "hybrid"
|
||||
top_k: int = 5
|
||||
score_threshold: float = 0.5
|
||||
result_count: int = 0
|
||||
results: Optional[List[Dict[str, Any]]] = None
|
||||
rerank_applied: str = "false"
|
||||
elapsed_time: int = 0
|
||||
source: Optional[str] = None
|
||||
user_id: Optional[str] = None
|
||||
sys_create_datetime: Optional[CSTDatetime] = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
@@ -0,0 +1,140 @@
|
||||
"""
|
||||
知识库分段 Schema
|
||||
"""
|
||||
from typing import Optional, List, Dict, Any
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field, ConfigDict
|
||||
|
||||
from app.base_schema import CSTDatetime
|
||||
|
||||
|
||||
class SegmentResponse(BaseModel):
|
||||
"""分段输出"""
|
||||
id: str
|
||||
knowledge_base_id: str
|
||||
document_id: str
|
||||
document_name: str = ""
|
||||
position: int = 0
|
||||
content: str
|
||||
answer: Optional[str] = None
|
||||
token_count: int = 0
|
||||
char_count: int = 0
|
||||
word_count: int = 0
|
||||
page_number: Optional[int] = None
|
||||
keywords: Optional[List[str]] = None
|
||||
metadata: Optional[Dict[str, Any]] = None
|
||||
embedding_status: str = "pending"
|
||||
enabled: bool = True
|
||||
hit_count: int = 0
|
||||
sys_create_datetime: Optional[CSTDatetime] = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class SegmentListResponse(BaseModel):
|
||||
"""分段列表输出"""
|
||||
id: str
|
||||
document_id: str
|
||||
document_name: str = ""
|
||||
position: int = 0
|
||||
content: str
|
||||
answer: Optional[str] = None
|
||||
token_count: int = 0
|
||||
char_count: int = 0
|
||||
word_count: int = 0
|
||||
page_number: Optional[int] = None
|
||||
keywords: Optional[List[str]] = None
|
||||
extra_metadata: Optional[Dict[str, Any]] = None
|
||||
enabled: bool = True
|
||||
hit_count: int = 0
|
||||
embedding_status: str = "pending"
|
||||
sys_create_datetime: Optional[CSTDatetime] = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class SegmentUpdateInput(BaseModel):
|
||||
"""更新分段"""
|
||||
content: Optional[str] = Field(None, description="分段内容")
|
||||
keywords: Optional[List[str]] = Field(None, description="关键词")
|
||||
enabled: Optional[bool] = Field(None, description="是否启用")
|
||||
extra_metadata: Optional[Dict[str, Any]] = Field(None, description="元数据")
|
||||
|
||||
|
||||
class SegmentCreateInput(BaseModel):
|
||||
"""手动创建分段"""
|
||||
content: str = Field(..., min_length=1, description="分段内容")
|
||||
answer: Optional[str] = Field(None, description="Q&A 模式的答案")
|
||||
keywords: Optional[List[str]] = Field(None, description="关键词")
|
||||
|
||||
|
||||
class ChunkPreviewInput(BaseModel):
|
||||
"""分块预览输入"""
|
||||
file_id: str = Field(..., description="文件ID")
|
||||
chunk_strategy: str = Field(default="recursive", description="分块策略")
|
||||
chunk_size: int = Field(default=500, ge=100, le=4000, description="分块大小")
|
||||
chunk_overlap: int = Field(default=50, ge=0, le=500, description="分块重叠")
|
||||
separator: Optional[str] = Field(None, description="自定义分隔符")
|
||||
process_rules: Optional[Dict[str, Any]] = Field(None, description="预处理规则")
|
||||
|
||||
|
||||
class ChunkPreviewItem(BaseModel):
|
||||
"""分块预览结果项"""
|
||||
position: int = 0
|
||||
content: str = ""
|
||||
char_count: int = 0
|
||||
token_count: int = 0
|
||||
word_count: int = 0
|
||||
answer: Optional[str] = None
|
||||
metadata: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class ChunkPreviewResponse(BaseModel):
|
||||
"""分块预览响应"""
|
||||
chunks: List[ChunkPreviewItem] = Field(default_factory=list)
|
||||
total: int = 0
|
||||
strategy: str = ""
|
||||
chunk_size: int = 0
|
||||
chunk_overlap: int = 0
|
||||
|
||||
|
||||
class RetrievalInput(BaseModel):
|
||||
"""检索输入"""
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
query: str = Field(..., min_length=1, description="查询文本")
|
||||
knowledge_base_ids: List[str] = Field(..., min_length=1, description="知识库ID列表")
|
||||
top_k: int = Field(default=5, ge=1, le=20, description="返回数量")
|
||||
score_threshold: float = Field(default=0.5, ge=0, le=1, description="相似度阈值")
|
||||
retrieval_mode: Optional[str] = Field(None, description="检索模式(不传则使用知识库配置)")
|
||||
rerank_enabled: Optional[bool] = Field(None, description="是否启用重排序(不传则使用知识库配置)")
|
||||
rerank_model_id: Optional[str] = Field(None, description="重排序模型ID(不传则使用知识库配置)")
|
||||
metadata_filter: Optional[Dict[str, Any]] = Field(None, description="元数据过滤条件")
|
||||
|
||||
|
||||
class RetrievalResult(BaseModel):
|
||||
"""检索结果"""
|
||||
segment_id: str
|
||||
document_id: str
|
||||
document_name: str = ""
|
||||
knowledge_base_id: str
|
||||
knowledge_base_name: str = ""
|
||||
content: str
|
||||
score: float = 0.0
|
||||
token_count: int = 0
|
||||
page_number: Optional[int] = None
|
||||
metadata: Optional[Dict[str, Any]] = None
|
||||
keywords: Optional[List[str]] = None
|
||||
match_source: Optional[str] = Field(None, description="命中来源: vector/fulltext/annotation")
|
||||
parent_content: Optional[str] = Field(None, description="父分段内容(Small-to-Big 模式)")
|
||||
|
||||
|
||||
class RetrievalResponse(BaseModel):
|
||||
"""检索响应"""
|
||||
results: List[RetrievalResult] = Field(default_factory=list)
|
||||
total: int = 0
|
||||
query: str = ""
|
||||
elapsed_time: int = 0
|
||||
retrieval_mode: str = ""
|
||||
rerank_applied: bool = False
|
||||
@@ -0,0 +1,3 @@
|
||||
"""
|
||||
知识库服务
|
||||
"""
|
||||
@@ -0,0 +1,159 @@
|
||||
"""
|
||||
文档预处理/清洗服务
|
||||
|
||||
参考 Dify 的 DatasetProcessRule,支持可配置的文本清洗规则。
|
||||
在文本提取之后、分块之前执行。
|
||||
"""
|
||||
import logging
|
||||
import re
|
||||
from typing import Dict, Any, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 默认预处理规则
|
||||
DEFAULT_PROCESS_RULES: Dict[str, Any] = {
|
||||
"pre_processing_rules": [
|
||||
{"id": "remove_extra_spaces", "enabled": True},
|
||||
{"id": "remove_urls_emails", "enabled": False},
|
||||
{"id": "remove_html_tags", "enabled": False},
|
||||
{"id": "remove_consecutive_newlines", "enabled": True},
|
||||
{"id": "remove_trailing_whitespace", "enabled": True},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class CleaningService:
|
||||
"""
|
||||
文本清洗服务
|
||||
|
||||
支持的清洗规则:
|
||||
- remove_extra_spaces: 合并连续空格为单个空格
|
||||
- remove_urls_emails: 移除 URL 和邮箱地址
|
||||
- remove_html_tags: 移除 HTML 标签
|
||||
- remove_consecutive_newlines: 合并连续空行(3+)为双空行
|
||||
- remove_trailing_whitespace: 去除行尾空白
|
||||
"""
|
||||
|
||||
# 规则处理器映射
|
||||
RULE_PROCESSORS = {
|
||||
"remove_extra_spaces": "_remove_extra_spaces",
|
||||
"remove_urls_emails": "_remove_urls_emails",
|
||||
"remove_html_tags": "_remove_html_tags",
|
||||
"remove_consecutive_newlines": "_remove_consecutive_newlines",
|
||||
"remove_trailing_whitespace": "_remove_trailing_whitespace",
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def clean(cls, text: str, process_rules: Optional[Dict[str, Any]] = None) -> str:
|
||||
"""
|
||||
根据预处理规则清洗文本
|
||||
|
||||
Args:
|
||||
text: 原始文本
|
||||
process_rules: 预处理规则配置,为 None 则使用默认规则
|
||||
|
||||
Returns:
|
||||
清洗后的文本
|
||||
"""
|
||||
if not text:
|
||||
return text
|
||||
|
||||
rules = process_rules or DEFAULT_PROCESS_RULES
|
||||
pre_rules = rules.get("pre_processing_rules", [])
|
||||
|
||||
original_length = len(text)
|
||||
|
||||
for rule in pre_rules:
|
||||
rule_id = rule.get("id", "")
|
||||
enabled = rule.get("enabled", False)
|
||||
|
||||
if not enabled:
|
||||
continue
|
||||
|
||||
processor_name = cls.RULE_PROCESSORS.get(rule_id)
|
||||
if not processor_name:
|
||||
logger.warning(f"未知的预处理规则: {rule_id}")
|
||||
continue
|
||||
|
||||
processor = getattr(cls, processor_name, None)
|
||||
if processor:
|
||||
text = processor(text)
|
||||
|
||||
cleaned_length = len(text)
|
||||
if original_length != cleaned_length:
|
||||
logger.info(
|
||||
f"文本清洗完成: {original_length} -> {cleaned_length} 字符 "
|
||||
f"(减少 {original_length - cleaned_length})"
|
||||
)
|
||||
|
||||
return text.strip()
|
||||
|
||||
@staticmethod
|
||||
def _remove_extra_spaces(text: str) -> str:
|
||||
"""合并连续空格为单个空格(保留换行符)"""
|
||||
# 只处理同一行内的连续空格,不影响换行
|
||||
lines = text.split('\n')
|
||||
cleaned_lines = []
|
||||
for line in lines:
|
||||
cleaned_lines.append(re.sub(r'[ \t]+', ' ', line))
|
||||
return '\n'.join(cleaned_lines)
|
||||
|
||||
@staticmethod
|
||||
def _remove_urls_emails(text: str) -> str:
|
||||
"""移除 URL 和邮箱地址"""
|
||||
# 移除 URL
|
||||
text = re.sub(
|
||||
r'https?://[^\s<>"{}|\\^`\[\]]+',
|
||||
'',
|
||||
text,
|
||||
)
|
||||
# 移除邮箱
|
||||
text = re.sub(
|
||||
r'[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}',
|
||||
'',
|
||||
text,
|
||||
)
|
||||
return text
|
||||
|
||||
@staticmethod
|
||||
def _remove_html_tags(text: str) -> str:
|
||||
"""移除 HTML 标签,保留文本内容"""
|
||||
# 移除 script 和 style 标签及其内容
|
||||
text = re.sub(r'<script[^>]*>.*?</script>', '', text, flags=re.DOTALL | re.IGNORECASE)
|
||||
text = re.sub(r'<style[^>]*>.*?</style>', '', text, flags=re.DOTALL | re.IGNORECASE)
|
||||
# 移除所有 HTML 标签
|
||||
text = re.sub(r'<[^>]+>', '', text)
|
||||
# 解码常见 HTML 实体
|
||||
text = text.replace(' ', ' ')
|
||||
text = text.replace('<', '<')
|
||||
text = text.replace('>', '>')
|
||||
text = text.replace('&', '&')
|
||||
text = text.replace('"', '"')
|
||||
text = text.replace(''', "'")
|
||||
return text
|
||||
|
||||
@staticmethod
|
||||
def _remove_consecutive_newlines(text: str) -> str:
|
||||
"""合并连续空行(3个以上换行)为双换行"""
|
||||
return re.sub(r'\n{3,}', '\n\n', text)
|
||||
|
||||
@staticmethod
|
||||
def _remove_trailing_whitespace(text: str) -> str:
|
||||
"""去除每行行尾空白"""
|
||||
return '\n'.join(line.rstrip() for line in text.split('\n'))
|
||||
|
||||
@classmethod
|
||||
def get_default_rules(cls) -> Dict[str, Any]:
|
||||
"""获取默认预处理规则"""
|
||||
return DEFAULT_PROCESS_RULES.copy()
|
||||
|
||||
@classmethod
|
||||
def get_available_rules(cls) -> List[Dict[str, str]]:
|
||||
"""获取所有可用的预处理规则"""
|
||||
return [
|
||||
{"id": "remove_extra_spaces", "label": "合并连续空格"},
|
||||
{"id": "remove_urls_emails", "label": "移除 URL 和邮箱"},
|
||||
{"id": "remove_html_tags", "label": "移除 HTML 标签"},
|
||||
{"id": "remove_consecutive_newlines", "label": "合并连续空行"},
|
||||
{"id": "remove_trailing_whitespace", "label": "去除行尾空白"},
|
||||
]
|
||||
@@ -0,0 +1,425 @@
|
||||
"""
|
||||
文档服务
|
||||
|
||||
文档上传、管理、状态控制
|
||||
"""
|
||||
import logging
|
||||
from typing import Optional, List, Tuple
|
||||
|
||||
from sqlalchemy import select, func, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ai_platform.knowledge.models import KnowledgeBase, KnowledgeDocument, KnowledgeSegment
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DocumentService:
|
||||
"""文档服务"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
|
||||
async def get_list(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
name: Optional[str] = None,
|
||||
status: Optional[str] = None,
|
||||
) -> Tuple[List[KnowledgeDocument], int]:
|
||||
"""获取文档列表"""
|
||||
query = select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeDocument.is_deleted == False,
|
||||
)
|
||||
|
||||
if name:
|
||||
query = query.where(KnowledgeDocument.name.ilike(f"%{name}%"))
|
||||
if status:
|
||||
query = query.where(KnowledgeDocument.status == status)
|
||||
|
||||
count_result = await self._db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(KnowledgeDocument.sys_create_datetime.desc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await self._db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return items, total
|
||||
|
||||
async def get_by_id(self, doc_id: str) -> Optional[KnowledgeDocument]:
|
||||
"""获取文档详情"""
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.id == doc_id,
|
||||
KnowledgeDocument.is_deleted == False
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def add_document(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
file_id: str,
|
||||
name: Optional[str] = None,
|
||||
) -> KnowledgeDocument:
|
||||
"""
|
||||
添加文档到知识库
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
file_id: 文件管理系统中的文件 ID
|
||||
name: 文档名称(不传则从文件信息获取)
|
||||
"""
|
||||
# 验证知识库存在
|
||||
kb_result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.id == knowledge_base_id,
|
||||
KnowledgeBase.is_deleted == False
|
||||
)
|
||||
)
|
||||
kb = kb_result.scalar_one_or_none()
|
||||
if not kb:
|
||||
raise ValueError('知识库不存在')
|
||||
|
||||
# 获取文件信息
|
||||
from core.file_manager.model import FileManager
|
||||
file_result = await self._db.execute(
|
||||
select(FileManager).where(
|
||||
FileManager.id == file_id,
|
||||
FileManager.is_deleted == False
|
||||
)
|
||||
)
|
||||
file_info = file_result.scalar_one_or_none()
|
||||
if not file_info:
|
||||
raise ValueError('文件不存在')
|
||||
|
||||
# 检查是否已添加(通过 file_id 去重)
|
||||
existing = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeDocument.file_id == file_id,
|
||||
KnowledgeDocument.is_deleted == False,
|
||||
)
|
||||
)
|
||||
if existing.scalars().first():
|
||||
raise ValueError('该文件已添加到知识库')
|
||||
|
||||
# 通过文件 MD5 检测内容重复(跨知识库)
|
||||
duplicate_warning = None
|
||||
if file_info.md5:
|
||||
dup_result = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.content_hash == file_info.md5,
|
||||
KnowledgeDocument.is_deleted == False,
|
||||
KnowledgeDocument.knowledge_base_id != knowledge_base_id,
|
||||
).limit(1)
|
||||
)
|
||||
dup_doc = dup_result.scalar_one_or_none()
|
||||
if dup_doc:
|
||||
duplicate_warning = f'该文件内容与其他知识库中的文档 "{dup_doc.name}" 重复'
|
||||
logger.info(f'文档内容重复检测: file_id={file_id}, 重复文档={dup_doc.id}')
|
||||
|
||||
# 同知识库内容去重(严格阻止)
|
||||
same_kb_dup = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.content_hash == file_info.md5,
|
||||
KnowledgeDocument.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeDocument.is_deleted == False,
|
||||
).limit(1)
|
||||
)
|
||||
if same_kb_dup.scalar_one_or_none():
|
||||
raise ValueError('该知识库中已存在相同内容的文档')
|
||||
|
||||
doc = KnowledgeDocument(
|
||||
knowledge_base_id=knowledge_base_id,
|
||||
file_id=file_id,
|
||||
name=name or file_info.name,
|
||||
file_type=file_info.file_ext or '',
|
||||
file_size=file_info.size or 0,
|
||||
content_hash=file_info.md5 or '',
|
||||
status='pending',
|
||||
duplicate_warning=duplicate_warning,
|
||||
)
|
||||
self._db.add(doc)
|
||||
await self._db.commit()
|
||||
await self._db.refresh(doc)
|
||||
|
||||
return doc
|
||||
|
||||
async def batch_add_documents(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
file_ids: List[str],
|
||||
) -> List[KnowledgeDocument]:
|
||||
"""批量添加文档"""
|
||||
docs = []
|
||||
for file_id in file_ids:
|
||||
try:
|
||||
doc = await self.add_document(knowledge_base_id, file_id)
|
||||
docs.append(doc)
|
||||
except ValueError as e:
|
||||
logger.warning(f'添加文档失败 (file_id={file_id}): {e}')
|
||||
continue
|
||||
return docs
|
||||
|
||||
async def delete_document(self, doc_id: str) -> bool:
|
||||
"""删除文档(软删除,同时删除分段 + 清理 Qdrant 向量)"""
|
||||
doc = await self.get_by_id(doc_id)
|
||||
if not doc:
|
||||
return False
|
||||
|
||||
doc.is_deleted = True
|
||||
|
||||
# 软删除关联分段
|
||||
await self._db.execute(
|
||||
update(KnowledgeSegment).where(
|
||||
KnowledgeSegment.document_id == doc_id
|
||||
).values(is_deleted=True)
|
||||
)
|
||||
|
||||
# 更新知识库统计
|
||||
from ai_platform.knowledge.services.indexing_service import IndexingService
|
||||
indexing_service = IndexingService(self._db)
|
||||
await indexing_service._update_kb_stats(doc.knowledge_base_id)
|
||||
|
||||
await self._db.commit()
|
||||
|
||||
# 从 Qdrant 删除该文档的所有向量
|
||||
try:
|
||||
from ai_platform.knowledge.vector_store import get_vector_store
|
||||
vector_store = get_vector_store()
|
||||
await vector_store.delete_by_filter(
|
||||
doc.knowledge_base_id,
|
||||
filter_conditions={'document_id': doc_id},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f'从 Qdrant 删除文档向量失败: {e}')
|
||||
|
||||
return True
|
||||
|
||||
async def toggle_document(self, doc_id: str, enabled: bool) -> Optional[KnowledgeDocument]:
|
||||
"""启用/禁用文档"""
|
||||
doc = await self.get_by_id(doc_id)
|
||||
if not doc:
|
||||
return None
|
||||
|
||||
doc.enabled = enabled
|
||||
|
||||
# 同时启用/禁用关联分段
|
||||
await self._db.execute(
|
||||
update(KnowledgeSegment).where(
|
||||
KnowledgeSegment.document_id == doc_id,
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
).values(enabled=enabled)
|
||||
)
|
||||
|
||||
await self._db.commit()
|
||||
await self._db.refresh(doc)
|
||||
return doc
|
||||
|
||||
async def get_segments(
|
||||
self,
|
||||
document_id: str,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
keyword: Optional[str] = None,
|
||||
) -> Tuple[List[KnowledgeSegment], int]:
|
||||
"""获取文档的分段列表"""
|
||||
query = select(KnowledgeSegment).where(
|
||||
KnowledgeSegment.document_id == document_id,
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
)
|
||||
|
||||
if keyword:
|
||||
query = query.where(KnowledgeSegment.content.ilike(f"%{keyword}%"))
|
||||
|
||||
count_result = await self._db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(KnowledgeSegment.position.asc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await self._db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return items, total
|
||||
|
||||
async def get_kb_segments(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
keyword: Optional[str] = None,
|
||||
enabled: Optional[bool] = None,
|
||||
embedding_status: Optional[str] = None,
|
||||
metadata_key: Optional[str] = None,
|
||||
metadata_value: Optional[str] = None,
|
||||
) -> Tuple[List[KnowledgeSegment], int]:
|
||||
"""获取知识库的所有分段"""
|
||||
from app.db_compat import json_extract, json_has_key
|
||||
|
||||
query = select(KnowledgeSegment).where(
|
||||
KnowledgeSegment.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
)
|
||||
|
||||
if keyword:
|
||||
query = query.where(KnowledgeSegment.content.ilike(f"%{keyword}%"))
|
||||
if enabled is not None:
|
||||
query = query.where(KnowledgeSegment.enabled == enabled)
|
||||
if embedding_status:
|
||||
query = query.where(KnowledgeSegment.embedding_status == embedding_status)
|
||||
if metadata_key:
|
||||
if metadata_value:
|
||||
query = query.where(
|
||||
json_extract(KnowledgeSegment.extra_metadata, metadata_key) == metadata_value
|
||||
)
|
||||
else:
|
||||
query = query.where(
|
||||
json_has_key(KnowledgeSegment.extra_metadata, metadata_key)
|
||||
)
|
||||
|
||||
count_result = await self._db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(KnowledgeSegment.document_id, KnowledgeSegment.position.asc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await self._db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return items, total
|
||||
|
||||
async def update_segment(
|
||||
self,
|
||||
segment_id: str,
|
||||
content: Optional[str] = None,
|
||||
keywords: Optional[List[str]] = None,
|
||||
enabled: Optional[bool] = None,
|
||||
extra_metadata: Optional[dict] = None,
|
||||
) -> Optional[KnowledgeSegment]:
|
||||
"""更新分段"""
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeSegment).where(
|
||||
KnowledgeSegment.id == segment_id,
|
||||
KnowledgeSegment.is_deleted == False
|
||||
)
|
||||
)
|
||||
segment = result.scalar_one_or_none()
|
||||
if not segment:
|
||||
return None
|
||||
|
||||
need_reindex = False
|
||||
if content is not None and content != segment.content:
|
||||
segment.content = content
|
||||
segment.char_count = len(content)
|
||||
segment.embedding_status = 'pending'
|
||||
need_reindex = True
|
||||
if keywords is not None:
|
||||
segment.keywords = keywords
|
||||
if enabled is not None:
|
||||
segment.enabled = enabled
|
||||
if extra_metadata is not None:
|
||||
segment.extra_metadata = extra_metadata
|
||||
|
||||
await self._db.commit()
|
||||
|
||||
# 如果内容变更,重新向量化
|
||||
if need_reindex:
|
||||
from ai_platform.knowledge.services.indexing_service import IndexingService
|
||||
indexing_service = IndexingService(self._db)
|
||||
await indexing_service.index_segment(segment.knowledge_base_id, segment_id)
|
||||
|
||||
await self._db.refresh(segment)
|
||||
return segment
|
||||
|
||||
async def add_segment(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
document_id: str,
|
||||
content: str,
|
||||
keywords: Optional[List[str]] = None,
|
||||
) -> KnowledgeSegment:
|
||||
"""手动添加分段"""
|
||||
# 获取当前最大 position
|
||||
max_pos_result = await self._db.execute(
|
||||
select(func.max(KnowledgeSegment.position)).where(
|
||||
KnowledgeSegment.document_id == document_id,
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
)
|
||||
)
|
||||
max_pos = max_pos_result.scalar() or 0
|
||||
|
||||
segment = KnowledgeSegment(
|
||||
knowledge_base_id=knowledge_base_id,
|
||||
document_id=document_id,
|
||||
position=max_pos + 1,
|
||||
content=content,
|
||||
char_count=len(content),
|
||||
keywords=keywords,
|
||||
embedding_status='pending',
|
||||
enabled=True,
|
||||
)
|
||||
self._db.add(segment)
|
||||
await self._db.commit()
|
||||
await self._db.refresh(segment)
|
||||
|
||||
# 向量化
|
||||
from ai_platform.knowledge.services.indexing_service import IndexingService
|
||||
indexing_service = IndexingService(self._db)
|
||||
await indexing_service.index_segment(knowledge_base_id, segment.id)
|
||||
|
||||
# 更新统计
|
||||
await indexing_service._update_kb_stats(knowledge_base_id)
|
||||
await self._db.commit()
|
||||
|
||||
await self._db.refresh(segment)
|
||||
return segment
|
||||
|
||||
async def delete_segment(self, segment_id: str) -> bool:
|
||||
"""删除分段(软删除 + 清理 Qdrant 向量)"""
|
||||
from sqlalchemy import update
|
||||
|
||||
# 先查询获取 kb_id
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeSegment.knowledge_base_id).where(
|
||||
KnowledgeSegment.id == segment_id,
|
||||
KnowledgeSegment.is_deleted == False
|
||||
)
|
||||
)
|
||||
row = result.first()
|
||||
if not row:
|
||||
return False
|
||||
|
||||
kb_id = str(row[0])
|
||||
|
||||
# 直接 SQL UPDATE 避免并发场景下的 StaleDataError
|
||||
await self._db.execute(
|
||||
update(KnowledgeSegment)
|
||||
.where(KnowledgeSegment.id == segment_id)
|
||||
.values(is_deleted=True)
|
||||
)
|
||||
await self._db.commit()
|
||||
|
||||
# 从 Qdrant 删除该分段的向量
|
||||
try:
|
||||
from ai_platform.knowledge.vector_store import get_vector_store
|
||||
vector_store = get_vector_store()
|
||||
await vector_store.delete(kb_id, [str(segment_id)])
|
||||
except Exception as e:
|
||||
logger.warning(f'从 Qdrant 删除分段向量失败: {e}')
|
||||
|
||||
# 更新统计
|
||||
from ai_platform.knowledge.services.indexing_service import IndexingService
|
||||
indexing_service = IndexingService(self._db)
|
||||
await indexing_service._update_kb_stats(kb_id)
|
||||
await self._db.commit()
|
||||
|
||||
return True
|
||||
@@ -0,0 +1,206 @@
|
||||
"""
|
||||
Embedding 服务
|
||||
|
||||
通过 OpenAI 兼容的 Embeddings API 将文本转换为向量
|
||||
支持所有兼容 OpenAI 接口的提供商(OpenAI、Qwen、Ollama 等)
|
||||
"""
|
||||
import logging
|
||||
import math
|
||||
from typing import List, Optional
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 单次批量请求的默认最大文本数
|
||||
DEFAULT_BATCH_SIZE = 50
|
||||
|
||||
# 不同提供商的批次大小限制
|
||||
PROVIDER_BATCH_SIZE = {
|
||||
'qwen': 10, # 阿里云 DashScope 限制 10
|
||||
'dashscope': 10,
|
||||
'siliconflow': 10, # 硅基流动限制较小
|
||||
'ollama': 1, # Ollama 通常逐条处理
|
||||
}
|
||||
|
||||
|
||||
|
||||
class EmbeddingService:
|
||||
"""
|
||||
Embedding 服务
|
||||
|
||||
通过模型 ID 获取对应的提供商,调用 OpenAI 兼容的 Embeddings API
|
||||
"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
self._client_cache = {}
|
||||
|
||||
async def _get_client_and_model(self, model_id: str):
|
||||
"""
|
||||
根据模型 ID 获取异步客户端和模型名称
|
||||
|
||||
Returns:
|
||||
(async_client, model_name, max_tokens, provider_type)
|
||||
"""
|
||||
from ai_platform.models import LLMModel, LLMProvider
|
||||
|
||||
result = await self._db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.id == model_id,
|
||||
LLMModel.is_active == True,
|
||||
LLMModel.is_deleted == False
|
||||
)
|
||||
)
|
||||
model = result.scalar_one_or_none()
|
||||
if not model:
|
||||
raise ValueError(f'Embedding 模型不存在或已禁用: {model_id}')
|
||||
if model.model_type != 'embedding':
|
||||
raise ValueError(f'模型 {model.display_name} 不是 Embedding 类型')
|
||||
|
||||
provider_result = await self._db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == model.provider_id,
|
||||
LLMProvider.is_active == True,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = provider_result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise ValueError('Embedding 模型对应的提供商不存在或已禁用')
|
||||
|
||||
cache_key = str(provider.id)
|
||||
if cache_key not in self._client_cache:
|
||||
import httpx
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
# 根据提供商类型确定 base_url
|
||||
if provider.provider_type == 'ollama':
|
||||
base_url = (provider.ollama_host or 'http://localhost:11434').rstrip('/') + '/v1'
|
||||
else:
|
||||
from ai_platform.providers.registry import ProviderRegistry
|
||||
provider_class = ProviderRegistry.get(provider.provider_type)
|
||||
default_base = getattr(provider_class, 'DEFAULT_API_BASE', 'https://api.openai.com/v1') if provider_class else 'https://api.openai.com/v1'
|
||||
base_url = provider.api_base or default_base
|
||||
|
||||
self._client_cache[cache_key] = AsyncOpenAI(
|
||||
api_key=provider.api_key or 'ollama',
|
||||
base_url=base_url,
|
||||
timeout=httpx.Timeout(120.0, connect=30.0),
|
||||
max_retries=5,
|
||||
)
|
||||
|
||||
return self._client_cache[cache_key], model.model_name, model.context_window or 8191, provider.provider_type
|
||||
|
||||
async def embed_text(
|
||||
self,
|
||||
model_id: str,
|
||||
text: str,
|
||||
dimensions: Optional[int] = None,
|
||||
) -> List[float]:
|
||||
"""
|
||||
将单个文本转换为向量
|
||||
|
||||
Args:
|
||||
model_id: Embedding 模型 ID
|
||||
text: 文本内容
|
||||
dimensions: 向量维度(可选,部分模型支持)
|
||||
|
||||
Returns:
|
||||
向量列表 List[float]
|
||||
"""
|
||||
results = await self.embed_texts(model_id, [text], dimensions)
|
||||
return results[0]
|
||||
|
||||
async def embed_texts(
|
||||
self,
|
||||
model_id: str,
|
||||
texts: List[str],
|
||||
dimensions: Optional[int] = None,
|
||||
) -> List[List[float]]:
|
||||
"""
|
||||
批量将文本转换为向量
|
||||
|
||||
Args:
|
||||
model_id: Embedding 模型 ID
|
||||
texts: 文本列表
|
||||
dimensions: 向量维度(可选)
|
||||
|
||||
Returns:
|
||||
向量列表 List[List[float]]
|
||||
"""
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
client, model_name, max_tokens, provider_type = await self._get_client_and_model(model_id)
|
||||
|
||||
# 根据提供商类型确定批次大小
|
||||
batch_size = PROVIDER_BATCH_SIZE.get(provider_type, DEFAULT_BATCH_SIZE)
|
||||
|
||||
# 预处理:截断超长文本
|
||||
processed_texts = []
|
||||
for text in texts:
|
||||
if not text or not text.strip():
|
||||
processed_texts.append(" ")
|
||||
else:
|
||||
# 粗略估算 token 数(中文约 1 字 = 1.5 token,英文约 4 字符 = 1 token)
|
||||
# 保守截断到 max_tokens * 2 个字符
|
||||
max_chars = max_tokens * 2
|
||||
if len(text) > max_chars:
|
||||
processed_texts.append(text[:max_chars])
|
||||
else:
|
||||
processed_texts.append(text)
|
||||
|
||||
# 分批处理
|
||||
all_embeddings = [None] * len(processed_texts)
|
||||
total_batches = math.ceil(len(processed_texts) / batch_size)
|
||||
|
||||
for batch_idx in range(total_batches):
|
||||
start = batch_idx * batch_size
|
||||
end = min(start + batch_size, len(processed_texts))
|
||||
batch_texts = processed_texts[start:end]
|
||||
|
||||
try:
|
||||
kwargs = {
|
||||
'model': model_name,
|
||||
'input': batch_texts,
|
||||
}
|
||||
# 仅对明确支持 dimensions 参数的模型传递该参数
|
||||
if dimensions:
|
||||
model_lower = model_name.lower()
|
||||
# OpenAI text-embedding-3 系列原生支持任意 dimensions
|
||||
if 'text-embedding-3' in model_lower:
|
||||
kwargs['dimensions'] = dimensions
|
||||
# DashScope text-embedding-v3 只接受 [64,128,256,512,768,1024]
|
||||
elif 'text-embedding-v3' in model_lower and dimensions in (64, 128, 256, 512, 768, 1024):
|
||||
kwargs['dimensions'] = dimensions
|
||||
|
||||
response = await client.embeddings.create(**kwargs)
|
||||
|
||||
for item in response.data:
|
||||
all_embeddings[start + item.index] = item.embedding
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f'Embedding 批次 {batch_idx + 1}/{total_batches} 失败: {e}')
|
||||
raise ValueError(f'Embedding 调用失败: {str(e)}')
|
||||
|
||||
# 检查是否所有向量都已生成
|
||||
for i, emb in enumerate(all_embeddings):
|
||||
if emb is None:
|
||||
raise ValueError(f'第 {i} 个文本的向量未生成')
|
||||
|
||||
return all_embeddings
|
||||
|
||||
async def get_embedding_dimensions(self, model_id: str) -> int:
|
||||
"""
|
||||
获取模型的向量维度(通过嵌入一个测试文本来检测)
|
||||
|
||||
Args:
|
||||
model_id: Embedding 模型 ID
|
||||
|
||||
Returns:
|
||||
向量维度
|
||||
"""
|
||||
test_embedding = await self.embed_text(model_id, "test")
|
||||
return len(test_embedding)
|
||||
@@ -0,0 +1,91 @@
|
||||
"""
|
||||
索引进度推送服务
|
||||
|
||||
通过 Redis Pub/Sub 推送索引进度,前端通过 SSE 订阅。
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Redis 频道前缀
|
||||
CHANNEL_PREFIX = "knowledge:indexing:progress:"
|
||||
|
||||
|
||||
class IndexingProgressService:
|
||||
"""索引进度推送服务"""
|
||||
|
||||
@staticmethod
|
||||
def _channel(knowledge_base_id: str) -> str:
|
||||
return f"{CHANNEL_PREFIX}{knowledge_base_id}"
|
||||
|
||||
@classmethod
|
||||
async def publish(
|
||||
cls,
|
||||
knowledge_base_id: str,
|
||||
document_id: str,
|
||||
step: str,
|
||||
progress: float,
|
||||
message: str = "",
|
||||
document_name: str = "",
|
||||
error: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
发布索引进度事件
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
document_id: 文档 ID
|
||||
step: 当前步骤 (extracting/cleaning/chunking/vectorizing/completed/failed)
|
||||
progress: 进度 0.0 ~ 1.0
|
||||
message: 进度描述
|
||||
document_name: 文档名称
|
||||
error: 错误信息(仅 failed 步骤)
|
||||
"""
|
||||
try:
|
||||
from utils.redis import RedisClient
|
||||
client = await RedisClient.get_client()
|
||||
|
||||
event = {
|
||||
"document_id": document_id,
|
||||
"document_name": document_name,
|
||||
"step": step,
|
||||
"progress": round(progress, 2),
|
||||
"message": message,
|
||||
}
|
||||
if error:
|
||||
event["error"] = error
|
||||
|
||||
await client.publish(
|
||||
cls._channel(knowledge_base_id),
|
||||
json.dumps(event, ensure_ascii=False),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"发布索引进度失败: {e}")
|
||||
|
||||
@classmethod
|
||||
async def subscribe(cls, knowledge_base_id: str):
|
||||
"""
|
||||
订阅索引进度事件(异步生成器,用于 SSE)
|
||||
|
||||
Yields:
|
||||
dict: 进度事件
|
||||
"""
|
||||
from utils.redis import RedisClient
|
||||
client = await RedisClient.get_client()
|
||||
pubsub = client.pubsub()
|
||||
channel = cls._channel(knowledge_base_id)
|
||||
|
||||
await pubsub.subscribe(channel)
|
||||
try:
|
||||
async for message in pubsub.listen():
|
||||
if message["type"] == "message":
|
||||
try:
|
||||
data = json.loads(message["data"])
|
||||
yield data
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
continue
|
||||
finally:
|
||||
await pubsub.unsubscribe(channel)
|
||||
await pubsub.close()
|
||||
@@ -0,0 +1,540 @@
|
||||
"""
|
||||
文档索引服务
|
||||
|
||||
负责文档处理管道:文本提取 → 分块 → 向量化 → 入库
|
||||
分段数据存入业务数据库,向量数据存入 Qdrant 向量数据库
|
||||
"""
|
||||
import hashlib
|
||||
import logging
|
||||
import math
|
||||
from datetime import datetime
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from sqlalchemy import select, func, delete
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ai_platform.knowledge.models import KnowledgeBase, KnowledgeDocument, KnowledgeSegment
|
||||
from ai_platform.knowledge.chunking import get_chunker
|
||||
from ai_platform.knowledge.services.embedding_service import EmbeddingService
|
||||
from ai_platform.knowledge.vector_store import get_vector_store, VectorPoint
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 向量化批次大小
|
||||
EMBEDDING_BATCH_SIZE = 50
|
||||
|
||||
|
||||
class IndexingService:
|
||||
"""
|
||||
文档索引服务
|
||||
|
||||
处理管道:
|
||||
1. 从文件管理系统提取文本内容
|
||||
2. 按知识库配置的策略分块
|
||||
3. 调用 Embedding 模型向量化
|
||||
4. 将分段写入业务数据库,向量写入 Qdrant
|
||||
"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
self._embedding_service = EmbeddingService(db)
|
||||
self._vector_store = get_vector_store()
|
||||
|
||||
async def index_document(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
document_id: str,
|
||||
) -> Tuple[int, int]:
|
||||
"""
|
||||
索引单个文档
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
document_id: 文档 ID
|
||||
|
||||
Returns:
|
||||
(segment_count, token_count) 分段数和 Token 数
|
||||
"""
|
||||
# 1. 获取知识库配置
|
||||
kb_result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.id == knowledge_base_id,
|
||||
KnowledgeBase.is_deleted == False
|
||||
)
|
||||
)
|
||||
kb = kb_result.scalar_one_or_none()
|
||||
if not kb:
|
||||
raise ValueError(f'知识库不存在: {knowledge_base_id}')
|
||||
|
||||
if not kb.embedding_model_id:
|
||||
raise ValueError('知识库未配置 Embedding 模型')
|
||||
|
||||
# 2. 获取文档
|
||||
doc_result = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.id == document_id,
|
||||
KnowledgeDocument.is_deleted == False
|
||||
)
|
||||
)
|
||||
doc = doc_result.scalar_one_or_none()
|
||||
if not doc:
|
||||
raise ValueError(f'文档不存在: {document_id}')
|
||||
|
||||
# 更新状态为 indexing
|
||||
doc.status = 'indexing'
|
||||
doc.indexing_started_at = datetime.now()
|
||||
doc.error_message = None
|
||||
await self._db.commit()
|
||||
|
||||
try:
|
||||
from ai_platform.knowledge.services.indexing_progress_service import IndexingProgressService
|
||||
|
||||
# 3. 提取文本
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='extracting', progress=0.1,
|
||||
message='正在提取文本内容...', document_name=doc.name,
|
||||
)
|
||||
text_content = await self._extract_text(doc.file_id)
|
||||
if not text_content or not text_content.strip():
|
||||
raise ValueError('文档内容为空,无法索引')
|
||||
|
||||
# 3.5 预处理/清洗
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='cleaning', progress=0.2,
|
||||
message='正在预处理/清洗文本...', document_name=doc.name,
|
||||
)
|
||||
from ai_platform.knowledge.services.cleaning_service import CleaningService
|
||||
text_content = CleaningService.clean(text_content, kb.process_rules)
|
||||
|
||||
# 计算内容哈希(用于去重)
|
||||
content_hash = hashlib.md5(text_content.encode('utf-8')).hexdigest()
|
||||
doc.content_hash = content_hash
|
||||
|
||||
# 4. 分块
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='chunking', progress=0.3,
|
||||
message='正在分块...', document_name=doc.name,
|
||||
)
|
||||
chunk_strategy = kb.chunk_strategy or 'recursive'
|
||||
chunk_kwargs = {}
|
||||
|
||||
# Q&A 模式需要 LLM 调用函数
|
||||
if chunk_strategy == 'qa':
|
||||
chunk_kwargs['llm_caller'] = self._create_llm_caller(kb)
|
||||
|
||||
chunker = get_chunker(
|
||||
strategy=chunk_strategy,
|
||||
chunk_size=kb.chunk_size or 500,
|
||||
chunk_overlap=kb.chunk_overlap or 50,
|
||||
separator=kb.separator,
|
||||
**chunk_kwargs,
|
||||
)
|
||||
|
||||
doc_metadata = {
|
||||
'document_id': document_id,
|
||||
'document_name': doc.name,
|
||||
'file_type': doc.file_type,
|
||||
}
|
||||
|
||||
# Q&A 模式使用异步分块
|
||||
if chunk_strategy == 'qa' and hasattr(chunker, 'chunk_async'):
|
||||
chunks = await chunker.chunk_async(text_content, metadata=doc_metadata)
|
||||
else:
|
||||
chunks = chunker.chunk(text_content, metadata=doc_metadata)
|
||||
|
||||
if not chunks:
|
||||
raise ValueError('文档分块结果为空')
|
||||
|
||||
# 5. 删除旧的分段(重新索引场景)
|
||||
await self._delete_document_segments(document_id, knowledge_base_id)
|
||||
|
||||
# 判断索引模式
|
||||
is_economy = (kb.indexing_technique == 'economy')
|
||||
|
||||
# 6. 确保 Qdrant collection 存在(经济模式跳过)
|
||||
if not is_economy:
|
||||
# 自动检测并修正 embedding 维度
|
||||
try:
|
||||
real_dim = await self._embedding_service.get_embedding_dimensions(kb.embedding_model_id)
|
||||
if real_dim != kb.embedding_dimensions:
|
||||
logger.info(f'修正 embedding 维度: {kb.embedding_dimensions} → {real_dim}')
|
||||
kb.embedding_dimensions = real_dim
|
||||
await self._db.commit()
|
||||
except Exception as e:
|
||||
logger.warning(f'自动检测 embedding 维度失败: {e}')
|
||||
|
||||
vector_size = kb.embedding_dimensions or 1536
|
||||
await self._vector_store.ensure_collection(knowledge_base_id, vector_size)
|
||||
|
||||
# 7. 入库(分批处理)
|
||||
segment_count = 0
|
||||
total_token_count = 0
|
||||
total_char_count = 0
|
||||
failed_embedding_count = 0
|
||||
|
||||
total_batches = math.ceil(len(chunks) / EMBEDDING_BATCH_SIZE)
|
||||
|
||||
for batch_idx in range(total_batches):
|
||||
start = batch_idx * EMBEDDING_BATCH_SIZE
|
||||
end = min(start + EMBEDDING_BATCH_SIZE, len(chunks))
|
||||
batch_chunks = chunks[start:end]
|
||||
|
||||
# 高质量模式:批量向量化;经济模式:跳过
|
||||
if is_economy:
|
||||
embeddings = [None] * len(batch_chunks)
|
||||
else:
|
||||
batch_texts = [c.content for c in batch_chunks]
|
||||
try:
|
||||
embeddings = await self._embedding_service.embed_texts(
|
||||
model_id=kb.embedding_model_id,
|
||||
texts=batch_texts,
|
||||
dimensions=kb.embedding_dimensions if kb.embedding_dimensions else None,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f'向量化批次 {batch_idx + 1}/{total_batches} 失败: {e}')
|
||||
embeddings = [None] * len(batch_texts)
|
||||
failed_embedding_count += len(batch_texts)
|
||||
|
||||
# 创建分段记录(业务数据库)+ 收集向量点(Qdrant)
|
||||
vector_points = []
|
||||
for i, chunk in enumerate(batch_chunks):
|
||||
embedding = embeddings[i] if i < len(embeddings) else None
|
||||
char_count = len(chunk.content)
|
||||
token_count = self._estimate_tokens(chunk.content)
|
||||
word_count = self._count_words(chunk.content)
|
||||
|
||||
# 自动提取关键词(高质量和经济模式均提取,增强全文检索)
|
||||
keywords = chunk.metadata.get('keywords')
|
||||
if not keywords:
|
||||
keywords = self._extract_keywords(chunk.content)
|
||||
|
||||
# 经济模式下 embedding_status 标记为 'skipped'
|
||||
if is_economy:
|
||||
emb_status = 'skipped'
|
||||
else:
|
||||
emb_status = 'completed' if embedding else 'failed'
|
||||
|
||||
segment = KnowledgeSegment(
|
||||
knowledge_base_id=knowledge_base_id,
|
||||
document_id=document_id,
|
||||
position=start + i,
|
||||
content=chunk.content,
|
||||
answer=chunk.metadata.get('answer'),
|
||||
token_count=token_count,
|
||||
char_count=char_count,
|
||||
word_count=word_count,
|
||||
page_number=chunk.metadata.get('page_number'),
|
||||
keywords=keywords,
|
||||
extra_metadata=chunk.metadata,
|
||||
embedding_status=emb_status,
|
||||
enabled=True,
|
||||
)
|
||||
self._db.add(segment)
|
||||
await self._db.flush()
|
||||
|
||||
# 收集向量点,稍后批量写入 Qdrant(经济模式跳过)
|
||||
if embedding and not is_economy:
|
||||
vector_points.append(VectorPoint(
|
||||
id=str(segment.id),
|
||||
vector=embedding,
|
||||
payload={
|
||||
'document_id': document_id,
|
||||
'knowledge_base_id': knowledge_base_id,
|
||||
'position': start + i,
|
||||
},
|
||||
))
|
||||
|
||||
segment_count += 1
|
||||
total_token_count += token_count
|
||||
total_char_count += char_count
|
||||
|
||||
# 提交业务数据库
|
||||
await self._db.commit()
|
||||
|
||||
# 批量写入 Qdrant(经济模式跳过)
|
||||
if vector_points:
|
||||
await self._vector_store.upsert(knowledge_base_id, vector_points)
|
||||
|
||||
step_label = '关键词提取中' if is_economy else '向量化中'
|
||||
batch_progress = 0.3 + 0.6 * (end / len(chunks))
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='vectorizing',
|
||||
progress=batch_progress,
|
||||
message=f'{step_label} {end}/{len(chunks)}',
|
||||
document_name=doc.name,
|
||||
)
|
||||
logger.info(f'文档 {doc.name} 索引进度: {end}/{len(chunks)}')
|
||||
|
||||
# 7. 更新文档状态
|
||||
if not is_economy and failed_embedding_count > 0:
|
||||
if failed_embedding_count >= segment_count:
|
||||
doc.status = 'failed'
|
||||
doc.error_message = f'所有 {segment_count} 个分段向量化失败'
|
||||
else:
|
||||
doc.status = 'completed'
|
||||
doc.error_message = f'{failed_embedding_count}/{segment_count} 个分段向量化失败'
|
||||
else:
|
||||
doc.status = 'completed'
|
||||
doc.segment_count = segment_count
|
||||
doc.token_count = total_token_count
|
||||
doc.char_count = total_char_count
|
||||
doc.indexing_completed_at = datetime.now()
|
||||
|
||||
# 8. 更新知识库统计
|
||||
await self._update_kb_stats(knowledge_base_id)
|
||||
|
||||
await self._db.commit()
|
||||
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='completed', progress=1.0,
|
||||
message=f'索引完成: {segment_count} 个分段',
|
||||
document_name=doc.name,
|
||||
)
|
||||
logger.info(f'文档 {doc.name} 索引完成: {segment_count} 个分段, {total_token_count} tokens')
|
||||
return segment_count, total_token_count
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f'文档索引失败: {e}')
|
||||
doc.status = 'failed'
|
||||
doc.error_message = str(e)[:500]
|
||||
await self._db.commit()
|
||||
await IndexingProgressService.publish(
|
||||
knowledge_base_id, document_id, step='failed', progress=0.0,
|
||||
message='索引失败', document_name=doc.name,
|
||||
error=str(e)[:200],
|
||||
)
|
||||
raise
|
||||
|
||||
async def reindex_document(self, knowledge_base_id: str, document_id: str) -> Tuple[int, int]:
|
||||
"""重新索引文档(删除旧分段后重新处理)"""
|
||||
return await self.index_document(knowledge_base_id, document_id)
|
||||
|
||||
async def index_segment(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
segment_id: str,
|
||||
) -> bool:
|
||||
"""
|
||||
为单个分段生成向量(用于手动添加或更新分段后)
|
||||
"""
|
||||
kb_result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.id == knowledge_base_id,
|
||||
KnowledgeBase.is_deleted == False
|
||||
)
|
||||
)
|
||||
kb = kb_result.scalar_one_or_none()
|
||||
if not kb or not kb.embedding_model_id:
|
||||
return False
|
||||
|
||||
seg_result = await self._db.execute(
|
||||
select(KnowledgeSegment).where(
|
||||
KnowledgeSegment.id == segment_id,
|
||||
KnowledgeSegment.is_deleted == False
|
||||
)
|
||||
)
|
||||
segment = seg_result.scalar_one_or_none()
|
||||
if not segment:
|
||||
return False
|
||||
|
||||
try:
|
||||
embedding = await self._embedding_service.embed_text(
|
||||
model_id=kb.embedding_model_id,
|
||||
text=segment.content,
|
||||
dimensions=kb.embedding_dimensions if kb.embedding_dimensions else None,
|
||||
)
|
||||
|
||||
# 确保 collection 存在
|
||||
vector_size = kb.embedding_dimensions or 1536
|
||||
await self._vector_store.ensure_collection(knowledge_base_id, vector_size)
|
||||
|
||||
# 写入 Qdrant
|
||||
await self._vector_store.upsert(knowledge_base_id, [VectorPoint(
|
||||
id=str(segment_id),
|
||||
vector=embedding,
|
||||
payload={
|
||||
'document_id': str(segment.document_id),
|
||||
'knowledge_base_id': knowledge_base_id,
|
||||
'position': segment.position or 0,
|
||||
},
|
||||
)])
|
||||
|
||||
segment.embedding_status = 'completed'
|
||||
await self._db.commit()
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f'分段向量化失败: {e}')
|
||||
segment.embedding_status = 'failed'
|
||||
await self._db.commit()
|
||||
return False
|
||||
|
||||
async def _extract_text(self, file_id: str) -> str:
|
||||
"""从文件管理系统提取文本内容(启用 OCR 支持图片和扫描版 PDF)"""
|
||||
from core.file_manager.service import FileManagerService
|
||||
|
||||
text_content = await FileManagerService.get_file_text_content(
|
||||
self._db, file_id, enable_ocr=True
|
||||
)
|
||||
if not text_content:
|
||||
raise ValueError('无法提取文件文本内容')
|
||||
return text_content
|
||||
|
||||
async def _delete_document_segments(self, document_id: str, knowledge_base_id: str):
|
||||
"""删除文档的所有分段(业务数据库 + Qdrant)"""
|
||||
# 先从 Qdrant 删除该文档的所有向量
|
||||
try:
|
||||
await self._vector_store.delete_by_filter(
|
||||
knowledge_base_id,
|
||||
filter_conditions={'document_id': document_id},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f'从 Qdrant 删除文档向量失败: {e}')
|
||||
|
||||
# 再从业务数据库删除分段记录
|
||||
await self._db.execute(
|
||||
delete(KnowledgeSegment).where(
|
||||
KnowledgeSegment.document_id == document_id
|
||||
)
|
||||
)
|
||||
|
||||
async def _update_kb_stats(self, knowledge_base_id: str):
|
||||
"""更新知识库统计信息"""
|
||||
# 文档数
|
||||
doc_count_result = await self._db.execute(
|
||||
select(func.count()).select_from(KnowledgeDocument).where(
|
||||
KnowledgeDocument.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeDocument.is_deleted == False,
|
||||
)
|
||||
)
|
||||
doc_count = doc_count_result.scalar() or 0
|
||||
|
||||
# 分段数和 Token 数
|
||||
seg_stats = await self._db.execute(
|
||||
select(
|
||||
func.count(),
|
||||
func.coalesce(func.sum(KnowledgeSegment.token_count), 0),
|
||||
func.coalesce(func.sum(KnowledgeSegment.char_count), 0),
|
||||
).where(
|
||||
KnowledgeSegment.knowledge_base_id == knowledge_base_id,
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
)
|
||||
)
|
||||
row = seg_stats.one()
|
||||
seg_count = row[0] or 0
|
||||
total_tokens = row[1] or 0
|
||||
total_chars = row[2] or 0
|
||||
|
||||
# 更新知识库
|
||||
kb_result = await self._db.execute(
|
||||
select(KnowledgeBase).where(KnowledgeBase.id == knowledge_base_id)
|
||||
)
|
||||
kb = kb_result.scalar_one_or_none()
|
||||
if kb:
|
||||
kb.document_count = doc_count
|
||||
kb.segment_count = seg_count
|
||||
kb.total_token_count = total_tokens
|
||||
kb.total_char_count = total_chars
|
||||
|
||||
def _create_llm_caller(self, kb: KnowledgeBase):
|
||||
"""
|
||||
创建 Q&A 分块所需的 LLM 调用函数
|
||||
|
||||
使用知识库所属应用中配置的第一个 chat 类型模型。
|
||||
"""
|
||||
db = self._db
|
||||
|
||||
async def llm_caller(system_prompt: str, user_prompt: str) -> str:
|
||||
from ai_platform.services.llm_service import LLMService
|
||||
|
||||
# 查找可用的 chat 模型
|
||||
from ai_platform.models import LLMModel
|
||||
model_result = await db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.model_type == 'chat',
|
||||
LLMModel.is_active == True,
|
||||
LLMModel.is_deleted == False,
|
||||
).limit(1)
|
||||
)
|
||||
model = model_result.scalar_one_or_none()
|
||||
if not model:
|
||||
raise ValueError('未找到可用的 chat 模型,无法进行 Q&A 拆分')
|
||||
|
||||
llm_service = LLMService(db)
|
||||
messages = []
|
||||
if system_prompt:
|
||||
messages.append({'role': 'system', 'content': system_prompt})
|
||||
messages.append({'role': 'user', 'content': user_prompt})
|
||||
|
||||
response = await llm_service.chat_async(
|
||||
model_id=str(model.id),
|
||||
messages=messages,
|
||||
temperature=0.3,
|
||||
max_tokens=4096,
|
||||
)
|
||||
return response.content or ''
|
||||
|
||||
return llm_caller
|
||||
|
||||
@staticmethod
|
||||
def _extract_keywords(text: str, max_keywords: int = 10) -> List[str]:
|
||||
"""
|
||||
从文本中提取关键词(经济模式使用)
|
||||
|
||||
使用简单的 TF 统计提取高频词,无需外部依赖。
|
||||
"""
|
||||
import re
|
||||
if not text:
|
||||
return []
|
||||
|
||||
# 中文分词(简单按标点和空格分割)
|
||||
# 提取中文词组(2-4字)和英文单词
|
||||
chinese_words = re.findall(r'[\u4e00-\u9fff]{2,4}', text)
|
||||
english_words = [w.lower() for w in re.findall(r'[a-zA-Z]{3,}', text)]
|
||||
|
||||
all_words = chinese_words + english_words
|
||||
|
||||
# 停用词(简单列表)
|
||||
stop_words = {
|
||||
'的', '了', '在', '是', '我', '有', '和', '就', '不', '人', '都', '一',
|
||||
'一个', '上', '也', '很', '到', '说', '要', '去', '你', '会', '着',
|
||||
'没有', '看', '好', '自己', '这', '他', '她', '它', '我们', '他们',
|
||||
'可以', '这个', '那个', '什么', '如果', '因为', '所以', '但是', '而且',
|
||||
'the', 'and', 'for', 'are', 'but', 'not', 'you', 'all', 'can',
|
||||
'had', 'her', 'was', 'one', 'our', 'out', 'has', 'have', 'been',
|
||||
'this', 'that', 'with', 'from', 'they', 'will', 'would', 'there',
|
||||
}
|
||||
|
||||
# 词频统计
|
||||
word_freq = {}
|
||||
for word in all_words:
|
||||
if word in stop_words or len(word) < 2:
|
||||
continue
|
||||
word_freq[word] = word_freq.get(word, 0) + 1
|
||||
|
||||
# 按频率排序取 top
|
||||
sorted_words = sorted(word_freq.items(), key=lambda x: x[1], reverse=True)
|
||||
return [w for w, _ in sorted_words[:max_keywords]]
|
||||
|
||||
@staticmethod
|
||||
def _count_words(text: str) -> int:
|
||||
"""计算词数(中文按字计算,英文按空格分词)"""
|
||||
if not text:
|
||||
return 0
|
||||
import re
|
||||
# 中文字符数
|
||||
chinese_chars = len(re.findall(r'[\u4e00-\u9fff]', text))
|
||||
# 英文单词数
|
||||
english_words = len(re.findall(r'[a-zA-Z]+', text))
|
||||
return chinese_chars + english_words
|
||||
|
||||
@staticmethod
|
||||
def _estimate_tokens(text: str) -> int:
|
||||
"""粗略估算文本的 Token 数"""
|
||||
if not text:
|
||||
return 0
|
||||
# 中文约 1 字 = 1.5 token,英文约 4 字符 = 1 token
|
||||
# 简单混合估算
|
||||
chinese_chars = sum(1 for c in text if '\u4e00' <= c <= '\u9fff')
|
||||
other_chars = len(text) - chinese_chars
|
||||
return int(chinese_chars * 1.5 + other_chars / 4)
|
||||
@@ -0,0 +1,267 @@
|
||||
"""
|
||||
知识库服务
|
||||
|
||||
知识库 CRUD 操作
|
||||
|
||||
数据权限:
|
||||
- 使用 get_list_with_data_scope() 自动应用数据权限
|
||||
- 支持本人、本部门、本部门及下级、全部等数据范围
|
||||
"""
|
||||
import logging
|
||||
from typing import Optional, List, Tuple
|
||||
|
||||
from sqlalchemy import select, func, or_, and_
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ai_platform.knowledge.models import KnowledgeBase, KnowledgeDocument, KnowledgeSegment
|
||||
from ai_platform.knowledge.schemas.knowledge_base_schema import KnowledgeBaseCreate, KnowledgeBaseUpdate
|
||||
from app.data_scope_utils import get_data_scope_filter, apply_data_scope_to_conditions
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 资源类型(用于数据权限配置)
|
||||
RESOURCE_TYPE = "knowledge_base"
|
||||
RESOURCE_DISPLAY_NAME = "知识库管理"
|
||||
|
||||
|
||||
class KnowledgeService:
|
||||
"""知识库服务"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
|
||||
async def get_list(
|
||||
self,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
name: Optional[str] = None,
|
||||
status: Optional[str] = None,
|
||||
application_id: Optional[str] = None,
|
||||
) -> Tuple[List[KnowledgeBase], int]:
|
||||
"""获取知识库列表"""
|
||||
query = select(KnowledgeBase).where(KnowledgeBase.is_deleted == False)
|
||||
|
||||
if application_id:
|
||||
query = query.where(or_(
|
||||
KnowledgeBase.application_id == application_id,
|
||||
and_(KnowledgeBase.application_id.is_(None), KnowledgeBase.is_global == True)
|
||||
))
|
||||
if name:
|
||||
query = query.where(KnowledgeBase.name.ilike(f"%{name}%"))
|
||||
if status:
|
||||
query = query.where(KnowledgeBase.status == status)
|
||||
|
||||
# 总数
|
||||
count_result = await self._db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(KnowledgeBase.sort.desc(), KnowledgeBase.sys_create_datetime.desc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await self._db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return items, total
|
||||
|
||||
async def get_list_with_data_scope(
|
||||
self,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
name: Optional[str] = None,
|
||||
status: Optional[str] = None,
|
||||
application_id: Optional[str] = None,
|
||||
) -> Tuple[List[KnowledgeBase], int]:
|
||||
"""
|
||||
获取知识库列表(带数据权限过滤)
|
||||
|
||||
自动从上下文获取当前用户信息,应用数据权限过滤
|
||||
"""
|
||||
conditions = [KnowledgeBase.is_deleted == False]
|
||||
|
||||
if application_id:
|
||||
conditions.append(or_(
|
||||
KnowledgeBase.application_id == application_id,
|
||||
and_(KnowledgeBase.application_id.is_(None), KnowledgeBase.is_global == True)
|
||||
))
|
||||
if name:
|
||||
conditions.append(KnowledgeBase.name.ilike(f"%{name}%"))
|
||||
if status:
|
||||
conditions.append(KnowledgeBase.status == status)
|
||||
|
||||
# 获取数据权限过滤条件并应用
|
||||
data_scope_filter = await get_data_scope_filter(self._db, RESOURCE_TYPE)
|
||||
scope_conditions = apply_data_scope_to_conditions(KnowledgeBase, data_scope_filter)
|
||||
conditions.extend(scope_conditions)
|
||||
|
||||
# 总数
|
||||
query = select(KnowledgeBase).where(and_(*conditions))
|
||||
count_result = await self._db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar() or 0
|
||||
|
||||
# 分页
|
||||
offset = (page - 1) * page_size
|
||||
query = query.order_by(KnowledgeBase.sort.desc(), KnowledgeBase.sys_create_datetime.desc())
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
result = await self._db.execute(query)
|
||||
items = result.scalars().all()
|
||||
|
||||
return items, total
|
||||
|
||||
async def get_by_id(self, kb_id: str) -> Optional[KnowledgeBase]:
|
||||
"""获取知识库详情"""
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.id == kb_id,
|
||||
KnowledgeBase.is_deleted == False
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def get_by_code(self, code: str) -> Optional[KnowledgeBase]:
|
||||
"""根据编码获取知识库"""
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.code == code,
|
||||
KnowledgeBase.is_deleted == False
|
||||
)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
async def create(self, data: KnowledgeBaseCreate) -> KnowledgeBase:
|
||||
"""创建知识库"""
|
||||
# 检查编码唯一性
|
||||
existing = await self.get_by_code(data.code)
|
||||
if existing:
|
||||
raise ValueError(f'知识库编码 {data.code} 已存在')
|
||||
|
||||
kb_data = data.model_dump()
|
||||
|
||||
# 自动检测 embedding 模型的真实维度
|
||||
if data.embedding_model_id:
|
||||
try:
|
||||
from ai_platform.knowledge.services.embedding_service import EmbeddingService
|
||||
embedding_service = EmbeddingService(self._db)
|
||||
real_dim = await embedding_service.get_embedding_dimensions(data.embedding_model_id)
|
||||
kb_data['embedding_dimensions'] = real_dim
|
||||
logger.info(f'自动检测 embedding 维度: {real_dim}')
|
||||
except Exception as e:
|
||||
logger.warning(f'自动检测 embedding 维度失败,使用默认值: {e}')
|
||||
|
||||
kb = KnowledgeBase(**kb_data)
|
||||
|
||||
# 自动填充创建人和部门
|
||||
from utils.context import get_current_user_info_from_context
|
||||
user_info = get_current_user_info_from_context()
|
||||
if user_info:
|
||||
if not kb.sys_creator_id:
|
||||
kb.sys_creator_id = user_info.get('user_id')
|
||||
if not kb.sys_dept_id and user_info.get('dept_id'):
|
||||
kb.sys_dept_id = user_info.get('dept_id')
|
||||
|
||||
self._db.add(kb)
|
||||
await self._db.commit()
|
||||
await self._db.refresh(kb)
|
||||
return kb
|
||||
|
||||
async def update(self, kb_id: str, data: KnowledgeBaseUpdate) -> Optional[KnowledgeBase]:
|
||||
"""更新知识库"""
|
||||
kb = await self.get_by_id(kb_id)
|
||||
if not kb:
|
||||
return None
|
||||
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
|
||||
# 如果更换了 embedding 模型,自动重新检测维度
|
||||
new_model_id = update_data.get('embedding_model_id')
|
||||
if new_model_id and new_model_id != kb.embedding_model_id:
|
||||
try:
|
||||
from ai_platform.knowledge.services.embedding_service import EmbeddingService
|
||||
embedding_service = EmbeddingService(self._db)
|
||||
real_dim = await embedding_service.get_embedding_dimensions(new_model_id)
|
||||
update_data['embedding_dimensions'] = real_dim
|
||||
logger.info(f'更换模型后自动检测 embedding 维度: {real_dim}')
|
||||
except Exception as e:
|
||||
logger.warning(f'自动检测 embedding 维度失败: {e}')
|
||||
|
||||
for key, value in update_data.items():
|
||||
setattr(kb, key, value)
|
||||
|
||||
await self._db.commit()
|
||||
await self._db.refresh(kb)
|
||||
return kb
|
||||
|
||||
async def delete(self, kb_id: str) -> bool:
|
||||
"""删除知识库(软删除 + 清理 Qdrant collection)"""
|
||||
kb = await self.get_by_id(kb_id)
|
||||
if not kb:
|
||||
return False
|
||||
|
||||
kb.is_deleted = True
|
||||
|
||||
# 同时软删除所有文档和分段
|
||||
doc_result = await self._db.execute(
|
||||
select(KnowledgeDocument).where(
|
||||
KnowledgeDocument.knowledge_base_id == kb_id,
|
||||
KnowledgeDocument.is_deleted == False
|
||||
)
|
||||
)
|
||||
docs = doc_result.scalars().all()
|
||||
for doc in docs:
|
||||
doc.is_deleted = True
|
||||
|
||||
# 软删除分段
|
||||
from sqlalchemy import update
|
||||
await self._db.execute(
|
||||
update(KnowledgeSegment).where(
|
||||
KnowledgeSegment.knowledge_base_id == kb_id
|
||||
).values(is_deleted=True)
|
||||
)
|
||||
|
||||
await self._db.commit()
|
||||
|
||||
# 删除 Qdrant 中对应的 collection
|
||||
try:
|
||||
from ai_platform.knowledge.vector_store import get_vector_store
|
||||
vector_store = get_vector_store()
|
||||
await vector_store.delete_collection(kb_id)
|
||||
except Exception as e:
|
||||
logger.warning(f'删除 Qdrant collection 失败: {e}')
|
||||
|
||||
return True
|
||||
|
||||
async def get_simple_list(self, application_id: Optional[str] = None) -> List[dict]:
|
||||
"""获取知识库简单列表(用于下拉选择)"""
|
||||
query = select(
|
||||
KnowledgeBase.id,
|
||||
KnowledgeBase.name,
|
||||
KnowledgeBase.code,
|
||||
KnowledgeBase.document_count,
|
||||
KnowledgeBase.segment_count,
|
||||
).where(
|
||||
KnowledgeBase.is_deleted == False,
|
||||
KnowledgeBase.status == 'active',
|
||||
)
|
||||
|
||||
if application_id:
|
||||
query = query.where(or_(
|
||||
KnowledgeBase.application_id == application_id,
|
||||
and_(KnowledgeBase.application_id.is_(None), KnowledgeBase.is_global == True)
|
||||
))
|
||||
|
||||
query = query.order_by(KnowledgeBase.sort.desc(), KnowledgeBase.sys_create_datetime.desc())
|
||||
result = await self._db.execute(query)
|
||||
rows = result.all()
|
||||
|
||||
return [
|
||||
{
|
||||
'id': row.id,
|
||||
'name': row.name,
|
||||
'code': row.code,
|
||||
'document_count': row.document_count or 0,
|
||||
'segment_count': row.segment_count or 0,
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
@@ -0,0 +1,182 @@
|
||||
"""
|
||||
Rerank 重排序服务
|
||||
|
||||
通过 Rerank 模型对检索结果进行重新排序,提升检索质量。
|
||||
支持 Jina/Cohere 风格的 Rerank API(大多数提供商兼容此接口)。
|
||||
"""
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RerankResult:
|
||||
"""重排序结果"""
|
||||
index: int
|
||||
relevance_score: float
|
||||
|
||||
|
||||
class RerankService:
|
||||
"""
|
||||
Rerank 重排序服务
|
||||
|
||||
通过模型 ID 获取对应的提供商,调用 Rerank API 对文档进行重排序。
|
||||
支持两种 API 风格:
|
||||
- Jina/Cohere 风格:POST /v1/rerank
|
||||
- OpenAI 兼容风格(部分提供商)
|
||||
"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
self._client_cache = {}
|
||||
|
||||
async def _get_client_config(self, model_id: str):
|
||||
"""
|
||||
根据模型 ID 获取 API 配置
|
||||
|
||||
Returns:
|
||||
(base_url, api_key, model_name)
|
||||
"""
|
||||
from ai_platform.models import LLMModel, LLMProvider
|
||||
|
||||
result = await self._db.execute(
|
||||
select(LLMModel).where(
|
||||
LLMModel.id == model_id,
|
||||
LLMModel.is_active == True,
|
||||
LLMModel.is_deleted == False
|
||||
)
|
||||
)
|
||||
model = result.scalar_one_or_none()
|
||||
if not model:
|
||||
raise ValueError(f'Rerank 模型不存在或已禁用: {model_id}')
|
||||
if model.model_type != 'rerank':
|
||||
raise ValueError(f'模型 {model.display_name} 不是 Rerank 类型')
|
||||
|
||||
provider_result = await self._db.execute(
|
||||
select(LLMProvider).where(
|
||||
LLMProvider.id == model.provider_id,
|
||||
LLMProvider.is_active == True,
|
||||
LLMProvider.is_deleted == False
|
||||
)
|
||||
)
|
||||
provider = provider_result.scalar_one_or_none()
|
||||
if not provider:
|
||||
raise ValueError('Rerank 模型对应的提供商不存在或已禁用')
|
||||
|
||||
if provider.provider_type == 'ollama':
|
||||
base_url = (provider.ollama_host or 'http://localhost:11434').rstrip('/') + '/v1'
|
||||
else:
|
||||
base_url = provider.api_base or 'https://api.openai.com/v1'
|
||||
|
||||
api_key = provider.api_key or 'ollama'
|
||||
|
||||
return base_url, api_key, model.model_name
|
||||
|
||||
async def rerank(
|
||||
self,
|
||||
model_id: str,
|
||||
query: str,
|
||||
documents: List[str],
|
||||
top_n: Optional[int] = None,
|
||||
) -> List[RerankResult]:
|
||||
"""
|
||||
对文档列表进行重排序
|
||||
|
||||
Args:
|
||||
model_id: Rerank 模型 ID
|
||||
query: 查询文本
|
||||
documents: 待排序的文档列表
|
||||
top_n: 返回前 N 个结果(默认返回全部)
|
||||
|
||||
Returns:
|
||||
按相关性降序排列的 RerankResult 列表
|
||||
"""
|
||||
if not documents:
|
||||
return []
|
||||
|
||||
if top_n is None:
|
||||
top_n = len(documents)
|
||||
|
||||
base_url, api_key, model_name = await self._get_client_config(model_id)
|
||||
|
||||
try:
|
||||
return await self._call_rerank_api(
|
||||
base_url=base_url,
|
||||
api_key=api_key,
|
||||
model_name=model_name,
|
||||
query=query,
|
||||
documents=documents,
|
||||
top_n=top_n,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f'Rerank 调用失败: {e}')
|
||||
raise ValueError(f'Rerank 调用失败: {str(e)}')
|
||||
|
||||
async def _call_rerank_api(
|
||||
self,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
model_name: str,
|
||||
query: str,
|
||||
documents: List[str],
|
||||
top_n: int,
|
||||
) -> List[RerankResult]:
|
||||
"""
|
||||
调用 Rerank API(Jina/Cohere 兼容风格)
|
||||
|
||||
POST {base_url}/rerank
|
||||
{
|
||||
"model": "...",
|
||||
"query": "...",
|
||||
"documents": ["...", "..."],
|
||||
"top_n": 5
|
||||
}
|
||||
|
||||
Response:
|
||||
{
|
||||
"results": [
|
||||
{"index": 0, "relevance_score": 0.95},
|
||||
{"index": 2, "relevance_score": 0.87},
|
||||
...
|
||||
]
|
||||
}
|
||||
"""
|
||||
import httpx
|
||||
|
||||
url = base_url.rstrip('/') + '/rerank'
|
||||
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'Authorization': f'Bearer {api_key}',
|
||||
}
|
||||
|
||||
payload = {
|
||||
'model': model_name,
|
||||
'query': query,
|
||||
'documents': documents,
|
||||
'top_n': top_n,
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient(timeout=60) as client:
|
||||
response = await client.post(url, json=payload, headers=headers)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
# 解析结果(兼容 Jina/Cohere/通义千问 等格式)
|
||||
raw_results = data.get('results', [])
|
||||
results = []
|
||||
for item in raw_results:
|
||||
results.append(RerankResult(
|
||||
index=item.get('index', 0),
|
||||
relevance_score=item.get('relevance_score', 0.0),
|
||||
))
|
||||
|
||||
# 按相关性降序排序
|
||||
results.sort(key=lambda r: r.relevance_score, reverse=True)
|
||||
|
||||
return results
|
||||
@@ -0,0 +1,666 @@
|
||||
"""
|
||||
检索服务
|
||||
|
||||
支持向量检索、全文检索、混合检索(RRF 融合)
|
||||
向量检索通过 Qdrant 向量数据库实现,全文检索通过业务数据库 SQL 实现
|
||||
"""
|
||||
import logging
|
||||
import time
|
||||
from typing import List, Optional, Dict, Any
|
||||
|
||||
from sqlalchemy import select, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from ai_platform.knowledge.models import KnowledgeBase, KnowledgeDocument, KnowledgeSegment
|
||||
from ai_platform.knowledge.services.embedding_service import EmbeddingService
|
||||
from ai_platform.knowledge.schemas.segment_schema import RetrievalResult
|
||||
from ai_platform.knowledge.vector_store import get_vector_store
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# RRF 融合常数
|
||||
RRF_K = 60
|
||||
|
||||
|
||||
class RetrievalService:
|
||||
"""
|
||||
检索服务
|
||||
|
||||
支持三种检索模式:
|
||||
- vector: 纯向量检索(通过 Qdrant)
|
||||
- fulltext: 纯全文检索(通过业务数据库 LIKE + 关键词匹配)
|
||||
- hybrid: 混合检索(向量 + 全文,RRF 融合排序)
|
||||
"""
|
||||
|
||||
def __init__(self, db: AsyncSession):
|
||||
self._db = db
|
||||
self._embedding_service = EmbeddingService(db)
|
||||
self._vector_store = get_vector_store()
|
||||
|
||||
async def retrieve(
|
||||
self,
|
||||
query: str,
|
||||
knowledge_base_ids: List[str],
|
||||
top_k: int = 5,
|
||||
score_threshold: float = 0.5,
|
||||
retrieval_mode: Optional[str] = None,
|
||||
rerank_enabled: Optional[bool] = None,
|
||||
rerank_model_id: Optional[str] = None,
|
||||
metadata_filter: Optional[Dict[str, Any]] = None,
|
||||
) -> List[RetrievalResult]:
|
||||
"""
|
||||
检索知识库
|
||||
|
||||
Args:
|
||||
query: 查询文本
|
||||
knowledge_base_ids: 知识库 ID 列表
|
||||
top_k: 返回数量
|
||||
score_threshold: 相似度阈值
|
||||
retrieval_mode: 检索模式(不传则使用第一个知识库的配置)
|
||||
rerank_enabled: 是否启用重排序(不传则使用知识库配置)
|
||||
rerank_model_id: 重排序模型 ID(不传则使用知识库配置)
|
||||
metadata_filter: 元数据过滤条件
|
||||
|
||||
Returns:
|
||||
检索结果列表
|
||||
"""
|
||||
if not query or not knowledge_base_ids:
|
||||
return []
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
# 获取知识库配置
|
||||
kb_map = await self._get_knowledge_bases(knowledge_base_ids)
|
||||
if not kb_map:
|
||||
return []
|
||||
|
||||
first_kb = list(kb_map.values())[0]
|
||||
|
||||
# 经济模式强制使用全文检索
|
||||
is_economy = getattr(first_kb, 'indexing_technique', 'high_quality') == 'economy'
|
||||
|
||||
# 确定检索模式
|
||||
if is_economy:
|
||||
retrieval_mode = 'fulltext'
|
||||
elif not retrieval_mode:
|
||||
retrieval_mode = first_kb.retrieval_mode or 'hybrid'
|
||||
|
||||
# 确定 rerank 配置(参数优先,否则使用知识库配置;经济模式禁用 rerank)
|
||||
if is_economy:
|
||||
rerank_enabled = False
|
||||
else:
|
||||
if rerank_enabled is None:
|
||||
rerank_enabled = first_kb.rerank_enabled or False
|
||||
if not rerank_model_id:
|
||||
rerank_model_id = first_kb.rerank_model_id
|
||||
|
||||
# 获取 embedding 模型(使用第一个知识库的配置)
|
||||
embedding_model_id = first_kb.embedding_model_id
|
||||
|
||||
# 如果启用了 rerank,初始检索多取一些候选结果
|
||||
candidate_multiplier = 3 if rerank_enabled and rerank_model_id else 2
|
||||
candidate_top_k = top_k * candidate_multiplier
|
||||
|
||||
results = []
|
||||
|
||||
if retrieval_mode == 'vector':
|
||||
results = await self._vector_search(
|
||||
query, knowledge_base_ids, embedding_model_id,
|
||||
top_k=candidate_top_k, score_threshold=score_threshold,
|
||||
dimensions=first_kb.embedding_dimensions,
|
||||
metadata_filter=metadata_filter,
|
||||
)
|
||||
for r in results:
|
||||
r.match_source = 'vector'
|
||||
elif retrieval_mode == 'fulltext':
|
||||
results = await self._fulltext_search(
|
||||
query, knowledge_base_ids, top_k=candidate_top_k,
|
||||
metadata_filter=metadata_filter,
|
||||
)
|
||||
for r in results:
|
||||
r.match_source = 'fulltext'
|
||||
elif retrieval_mode == 'hybrid':
|
||||
# 混合检索:向量 + 全文,RRF 融合
|
||||
vector_results = await self._vector_search(
|
||||
query, knowledge_base_ids, embedding_model_id,
|
||||
top_k=candidate_top_k, score_threshold=score_threshold,
|
||||
dimensions=first_kb.embedding_dimensions,
|
||||
metadata_filter=metadata_filter,
|
||||
)
|
||||
for r in vector_results:
|
||||
r.match_source = 'vector'
|
||||
fulltext_results = await self._fulltext_search(
|
||||
query, knowledge_base_ids, top_k=candidate_top_k,
|
||||
metadata_filter=metadata_filter,
|
||||
)
|
||||
for r in fulltext_results:
|
||||
r.match_source = 'fulltext'
|
||||
results = self._rrf_merge(vector_results, fulltext_results)
|
||||
|
||||
# 多知识库权重加权
|
||||
if len(kb_map) > 1:
|
||||
for r in results:
|
||||
kb = kb_map.get(r.knowledge_base_id)
|
||||
weight = getattr(kb, 'retrieval_weight', 1.0) or 1.0 if kb else 1.0
|
||||
if weight != 1.0:
|
||||
r.score = round(r.score * weight, 4)
|
||||
results.sort(key=lambda x: x.score, reverse=True)
|
||||
|
||||
# 过滤低分结果(rerank 前先粗筛)
|
||||
if not rerank_enabled:
|
||||
results = [r for r in results if r.score >= score_threshold]
|
||||
|
||||
# Rerank 重排序
|
||||
if rerank_enabled and rerank_model_id and results:
|
||||
results = await self._rerank_results(query, results, rerank_model_id, top_k)
|
||||
# rerank 后再按阈值过滤
|
||||
results = [r for r in results if r.score >= score_threshold]
|
||||
|
||||
# 截断到 top_k
|
||||
results = results[:top_k]
|
||||
|
||||
# 内容级去重(多知识库检索时可能有重复内容)
|
||||
results = self._deduplicate_results(results)
|
||||
|
||||
# 标注优先匹配:将匹配到的标注结果插入到最前面
|
||||
annotation_results = await self._match_annotations(
|
||||
query, knowledge_base_ids, embedding_model_id,
|
||||
score_threshold=score_threshold,
|
||||
dimensions=first_kb.embedding_dimensions,
|
||||
)
|
||||
if annotation_results:
|
||||
for r in annotation_results:
|
||||
r.match_source = 'annotation'
|
||||
# 标注结果置顶,去重后合并
|
||||
existing_ids = {r.segment_id for r in annotation_results}
|
||||
results = annotation_results + [r for r in results if r.segment_id not in existing_ids]
|
||||
results = results[:top_k]
|
||||
|
||||
# 填充知识库名称和文档名称
|
||||
await self._fill_names(results, kb_map)
|
||||
|
||||
# 填充父分段内容(Small-to-Big 模式)
|
||||
await self._fill_parent_content(results)
|
||||
|
||||
# 更新命中次数
|
||||
segment_ids = [r.segment_id for r in results]
|
||||
if segment_ids:
|
||||
await self._update_hit_counts(segment_ids)
|
||||
|
||||
elapsed = int((time.time() - start_time) * 1000)
|
||||
rerank_info = ', rerank=ON' if rerank_enabled else ''
|
||||
logger.info(f'检索完成: {len(results)} 条结果, 耗时 {elapsed}ms, 模式={retrieval_mode}{rerank_info}')
|
||||
|
||||
return results
|
||||
|
||||
async def _rerank_results(
|
||||
self,
|
||||
query: str,
|
||||
results: List[RetrievalResult],
|
||||
rerank_model_id: str,
|
||||
top_n: int,
|
||||
) -> List[RetrievalResult]:
|
||||
"""使用 Rerank 模型对检索结果重排序"""
|
||||
from ai_platform.knowledge.services.rerank_service import RerankService
|
||||
|
||||
try:
|
||||
rerank_service = RerankService(self._db)
|
||||
documents = [r.content for r in results]
|
||||
|
||||
rerank_results = await rerank_service.rerank(
|
||||
model_id=rerank_model_id,
|
||||
query=query,
|
||||
documents=documents,
|
||||
top_n=top_n,
|
||||
)
|
||||
|
||||
# 按 rerank 分数重新排列结果
|
||||
reranked = []
|
||||
for rr in rerank_results:
|
||||
if 0 <= rr.index < len(results):
|
||||
result = results[rr.index]
|
||||
result.score = round(rr.relevance_score, 4)
|
||||
reranked.append(result)
|
||||
|
||||
logger.info(f'Rerank 完成: {len(results)} -> {len(reranked)} 条结果')
|
||||
return reranked
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f'Rerank 失败,使用原始排序: {e}')
|
||||
return results
|
||||
|
||||
async def _vector_search(
|
||||
self,
|
||||
query: str,
|
||||
knowledge_base_ids: List[str],
|
||||
embedding_model_id: str,
|
||||
top_k: int = 10,
|
||||
score_threshold: float = 0.0,
|
||||
dimensions: Optional[int] = None,
|
||||
metadata_filter: Optional[Dict[str, Any]] = None,
|
||||
) -> List[RetrievalResult]:
|
||||
"""向量检索(通过 Qdrant 余弦相似度)"""
|
||||
if not embedding_model_id:
|
||||
logger.warning('未配置 Embedding 模型,跳过向量检索')
|
||||
return []
|
||||
|
||||
try:
|
||||
query_embedding = await self._embedding_service.embed_text(
|
||||
model_id=embedding_model_id,
|
||||
text=query,
|
||||
dimensions=dimensions,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f'查询向量化失败: {e}')
|
||||
return []
|
||||
|
||||
# 对每个知识库分别搜索(每个知识库对应一个 Qdrant collection)
|
||||
all_hits = []
|
||||
for kb_id in knowledge_base_ids:
|
||||
hits = await self._vector_store.search(
|
||||
knowledge_base_id=kb_id,
|
||||
query_vector=query_embedding,
|
||||
top_k=top_k,
|
||||
score_threshold=score_threshold,
|
||||
)
|
||||
all_hits.extend(hits)
|
||||
|
||||
if not all_hits:
|
||||
return []
|
||||
|
||||
# 按分数排序
|
||||
all_hits.sort(key=lambda h: h.score, reverse=True)
|
||||
all_hits = all_hits[:top_k]
|
||||
|
||||
# 从业务数据库获取分段详情
|
||||
segment_ids = [h.id for h in all_hits]
|
||||
score_map = {h.id: h.score for h in all_hits}
|
||||
|
||||
seg_result = await self._db.execute(
|
||||
select(KnowledgeSegment).where(
|
||||
KnowledgeSegment.id.in_(segment_ids),
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
KnowledgeSegment.enabled == True,
|
||||
)
|
||||
)
|
||||
segments = {str(s.id): s for s in seg_result.scalars().all()}
|
||||
|
||||
results = []
|
||||
for hit in all_hits:
|
||||
seg = segments.get(hit.id)
|
||||
if not seg:
|
||||
continue
|
||||
# Q&A 模式:question 用于匹配,返回 answer 作为 content
|
||||
content = seg.answer if seg.answer else seg.content
|
||||
meta = dict(seg.extra_metadata) if seg.extra_metadata else {}
|
||||
if seg.answer:
|
||||
meta['question'] = seg.content
|
||||
meta['chunk_mode'] = 'qa'
|
||||
results.append(RetrievalResult(
|
||||
segment_id=str(seg.id),
|
||||
document_id=str(seg.document_id),
|
||||
knowledge_base_id=str(seg.knowledge_base_id),
|
||||
content=content,
|
||||
score=round(score_map.get(hit.id, 0.0), 4),
|
||||
token_count=seg.token_count or 0,
|
||||
page_number=seg.page_number,
|
||||
metadata=meta,
|
||||
keywords=seg.keywords,
|
||||
))
|
||||
|
||||
return results
|
||||
|
||||
async def _fulltext_search(
|
||||
self,
|
||||
query: str,
|
||||
knowledge_base_ids: List[str],
|
||||
top_k: int = 10,
|
||||
metadata_filter: Optional[Dict[str, Any]] = None,
|
||||
) -> List[RetrievalResult]:
|
||||
"""全文检索(基于 ORM LIKE 和关键词匹配,兼容所有数据库)"""
|
||||
import re
|
||||
keywords = re.split(r'[\s,,。.!!??;;、]+', query)
|
||||
keywords = [k.strip() for k in keywords if k.strip() and len(k.strip()) >= 2]
|
||||
|
||||
if not keywords:
|
||||
keywords = [query.strip()]
|
||||
|
||||
# 使用 SQLAlchemy ORM 构建查询(兼容 PG / MySQL 等)
|
||||
from sqlalchemy import or_
|
||||
keyword_conditions = [
|
||||
func.lower(KnowledgeSegment.content).contains(kw.lower())
|
||||
for kw in keywords[:5]
|
||||
]
|
||||
|
||||
conditions = [
|
||||
KnowledgeSegment.knowledge_base_id.in_(knowledge_base_ids),
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
KnowledgeSegment.enabled == True,
|
||||
or_(*keyword_conditions),
|
||||
]
|
||||
|
||||
# 元数据过滤
|
||||
if metadata_filter:
|
||||
conditions.extend(self._build_metadata_conditions(metadata_filter))
|
||||
|
||||
stmt = (
|
||||
select(KnowledgeSegment)
|
||||
.where(*conditions)
|
||||
.order_by(KnowledgeSegment.char_count.asc())
|
||||
.limit(top_k)
|
||||
)
|
||||
|
||||
result = await self._db.execute(stmt)
|
||||
segments = result.scalars().all()
|
||||
|
||||
# 计算简单的关键词匹配分数
|
||||
results = []
|
||||
for seg in segments:
|
||||
content_lower = seg.content.lower()
|
||||
match_count = sum(1 for kw in keywords if kw.lower() in content_lower)
|
||||
score = match_count / len(keywords) if keywords else 0
|
||||
# Q&A 模式:返回 answer 作为 content
|
||||
content = seg.answer if seg.answer else seg.content
|
||||
meta = dict(seg.extra_metadata) if seg.extra_metadata else {}
|
||||
if seg.answer:
|
||||
meta['question'] = seg.content
|
||||
meta['chunk_mode'] = 'qa'
|
||||
results.append(RetrievalResult(
|
||||
segment_id=str(seg.id),
|
||||
document_id=str(seg.document_id),
|
||||
knowledge_base_id=str(seg.knowledge_base_id),
|
||||
content=content,
|
||||
score=round(score, 4),
|
||||
token_count=seg.token_count or 0,
|
||||
page_number=seg.page_number,
|
||||
metadata=meta,
|
||||
keywords=seg.keywords,
|
||||
))
|
||||
|
||||
results.sort(key=lambda x: x.score, reverse=True)
|
||||
return results
|
||||
|
||||
def _rrf_merge(
|
||||
self,
|
||||
vector_results: List[RetrievalResult],
|
||||
fulltext_results: List[RetrievalResult],
|
||||
) -> List[RetrievalResult]:
|
||||
"""
|
||||
RRF (Reciprocal Rank Fusion) 融合排序
|
||||
|
||||
RRF_score = sum(1 / (k + rank_i)) for each result list
|
||||
"""
|
||||
scores = {} # segment_id -> (rrf_score, result)
|
||||
|
||||
# 向量检索结果排名
|
||||
for rank, result in enumerate(vector_results):
|
||||
rrf_score = 1.0 / (RRF_K + rank + 1)
|
||||
if result.segment_id in scores:
|
||||
old_score, old_result = scores[result.segment_id]
|
||||
scores[result.segment_id] = (old_score + rrf_score, old_result)
|
||||
else:
|
||||
scores[result.segment_id] = (rrf_score, result)
|
||||
|
||||
# 全文检索结果排名
|
||||
for rank, result in enumerate(fulltext_results):
|
||||
rrf_score = 1.0 / (RRF_K + rank + 1)
|
||||
if result.segment_id in scores:
|
||||
old_score, old_result = scores[result.segment_id]
|
||||
scores[result.segment_id] = (old_score + rrf_score, old_result)
|
||||
else:
|
||||
scores[result.segment_id] = (rrf_score, result)
|
||||
|
||||
# 按 RRF 分数排序
|
||||
sorted_items = sorted(scores.values(), key=lambda x: x[0], reverse=True)
|
||||
|
||||
if not sorted_items:
|
||||
return []
|
||||
|
||||
# 归一化分数到 0-1
|
||||
# RRF 单条结果的理论最大分数为 2/(k+1)(同时出现在两个列表的第一名)
|
||||
# 使用理论最大值归一化,避免单条结果被归一化为 100%
|
||||
theoretical_max = 2.0 / (RRF_K + 1)
|
||||
|
||||
results = []
|
||||
for rrf_score, result in sorted_items:
|
||||
normalized_score = min(rrf_score / theoretical_max, 1.0)
|
||||
result.score = round(normalized_score, 4)
|
||||
results.append(result)
|
||||
|
||||
return results
|
||||
|
||||
@staticmethod
|
||||
def _deduplicate_results(results: List[RetrievalResult], similarity_threshold: float = 0.95) -> List[RetrievalResult]:
|
||||
"""
|
||||
内容级去重(多知识库检索时可能有重复内容)
|
||||
|
||||
使用内容前 200 字符的相似度判断是否重复,保留分数最高的。
|
||||
"""
|
||||
if len(results) <= 1:
|
||||
return results
|
||||
|
||||
deduplicated = []
|
||||
seen_contents = []
|
||||
|
||||
for r in results:
|
||||
content_key = r.content[:200].strip().lower()
|
||||
is_dup = False
|
||||
for seen in seen_contents:
|
||||
# 简单的字符重叠率判断
|
||||
if content_key == seen:
|
||||
is_dup = True
|
||||
break
|
||||
# 如果前 200 字符有 95% 以上重叠,视为重复
|
||||
shorter = min(len(content_key), len(seen))
|
||||
if shorter > 0:
|
||||
common = sum(1 for a, b in zip(content_key, seen) if a == b)
|
||||
if common / shorter >= similarity_threshold:
|
||||
is_dup = True
|
||||
break
|
||||
if not is_dup:
|
||||
deduplicated.append(r)
|
||||
seen_contents.append(content_key)
|
||||
|
||||
return deduplicated
|
||||
|
||||
async def _fill_parent_content(self, results: List[RetrievalResult]):
|
||||
"""填充父分段内容(Small-to-Big 模式)"""
|
||||
if not results:
|
||||
return
|
||||
|
||||
# 获取所有 segment_id,查询是否有 parent_segment_id
|
||||
segment_ids = [r.segment_id for r in results if r.segment_id]
|
||||
if not segment_ids:
|
||||
return
|
||||
|
||||
seg_result = await self._db.execute(
|
||||
select(KnowledgeSegment.id, KnowledgeSegment.parent_segment_id).where(
|
||||
KnowledgeSegment.id.in_(segment_ids),
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
)
|
||||
)
|
||||
parent_map = {}
|
||||
for row in seg_result:
|
||||
if row.parent_segment_id:
|
||||
parent_map[str(row.id)] = row.parent_segment_id
|
||||
|
||||
if not parent_map:
|
||||
return
|
||||
|
||||
# 批量获取父分段内容
|
||||
parent_ids = list(set(parent_map.values()))
|
||||
parent_result = await self._db.execute(
|
||||
select(KnowledgeSegment.id, KnowledgeSegment.content).where(
|
||||
KnowledgeSegment.id.in_(parent_ids),
|
||||
KnowledgeSegment.is_deleted == False,
|
||||
)
|
||||
)
|
||||
parent_content_map = {str(row.id): row.content for row in parent_result}
|
||||
|
||||
# 填充到结果中
|
||||
for r in results:
|
||||
parent_id = parent_map.get(r.segment_id)
|
||||
if parent_id:
|
||||
r.parent_content = parent_content_map.get(str(parent_id))
|
||||
|
||||
@staticmethod
|
||||
def _build_metadata_conditions(metadata_filter: Dict[str, Any]) -> list:
|
||||
"""构建元数据过滤条件(基于 JSON 字段,跨数据库兼容)"""
|
||||
from app.db_compat import json_extract
|
||||
|
||||
conditions = []
|
||||
for key, value in metadata_filter.items():
|
||||
if value is not None:
|
||||
# 使用跨数据库兼容的 json_extract 函数
|
||||
try:
|
||||
conditions.append(
|
||||
json_extract(KnowledgeSegment.extra_metadata, key) == str(value)
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return conditions
|
||||
|
||||
async def _get_knowledge_bases(self, kb_ids: List[str]) -> Dict[str, KnowledgeBase]:
|
||||
"""批量获取知识库"""
|
||||
result = await self._db.execute(
|
||||
select(KnowledgeBase).where(
|
||||
KnowledgeBase.id.in_(kb_ids),
|
||||
KnowledgeBase.is_deleted == False,
|
||||
)
|
||||
)
|
||||
kbs = result.scalars().all()
|
||||
return {str(kb.id): kb for kb in kbs}
|
||||
|
||||
async def _fill_names(self, results: List[RetrievalResult], kb_map: Dict[str, KnowledgeBase]):
|
||||
"""填充知识库名称和文档名称"""
|
||||
if not results:
|
||||
return
|
||||
|
||||
# 获取文档名称
|
||||
doc_ids = list({r.document_id for r in results})
|
||||
doc_result = await self._db.execute(
|
||||
select(KnowledgeDocument.id, KnowledgeDocument.name).where(
|
||||
KnowledgeDocument.id.in_(doc_ids)
|
||||
)
|
||||
)
|
||||
doc_name_map = {row.id: row.name for row in doc_result}
|
||||
|
||||
for result in results:
|
||||
result.document_name = doc_name_map.get(result.document_id, '')
|
||||
kb = kb_map.get(result.knowledge_base_id)
|
||||
result.knowledge_base_name = kb.name if kb else ''
|
||||
|
||||
async def _update_hit_counts(self, segment_ids: List[str]):
|
||||
"""更新分段命中次数(兼容所有数据库)"""
|
||||
if not segment_ids:
|
||||
return
|
||||
try:
|
||||
from sqlalchemy import update
|
||||
await self._db.execute(
|
||||
update(KnowledgeSegment)
|
||||
.where(KnowledgeSegment.id.in_(segment_ids))
|
||||
.values(hit_count=func.coalesce(KnowledgeSegment.hit_count, 0) + 1)
|
||||
)
|
||||
await self._db.commit()
|
||||
except Exception as e:
|
||||
logger.warning(f'更新命中次数失败: {e}')
|
||||
|
||||
async def _match_annotations(
|
||||
self,
|
||||
query: str,
|
||||
knowledge_base_ids: List[str],
|
||||
embedding_model_id: Optional[str],
|
||||
score_threshold: float = 0.5,
|
||||
dimensions: Optional[int] = None,
|
||||
max_results: int = 3,
|
||||
) -> List[RetrievalResult]:
|
||||
"""
|
||||
匹配标注(Q&A 对)
|
||||
|
||||
通过向量相似度匹配标注的 question,返回对应的 answer。
|
||||
标注结果优先级高于普通分段。
|
||||
"""
|
||||
from ai_platform.knowledge.models import KnowledgeAnnotation
|
||||
|
||||
if not embedding_model_id:
|
||||
return []
|
||||
|
||||
try:
|
||||
# 向量化查询
|
||||
query_embedding = await self._embedding_service.embed_text(
|
||||
model_id=embedding_model_id,
|
||||
text=query,
|
||||
dimensions=dimensions if dimensions else None,
|
||||
)
|
||||
|
||||
# 在 Qdrant 中搜索标注向量(payload.type == 'annotation')
|
||||
all_hits = []
|
||||
for kb_id in knowledge_base_ids:
|
||||
try:
|
||||
hits = await self._vector_store.search(
|
||||
knowledge_base_id=kb_id,
|
||||
query_vector=query_embedding,
|
||||
top_k=max_results,
|
||||
score_threshold=score_threshold,
|
||||
filter_conditions={'type': 'annotation'},
|
||||
)
|
||||
all_hits.extend(hits)
|
||||
except Exception as e:
|
||||
logger.warning(f'标注向量搜索失败 (kb={kb_id}): {e}')
|
||||
continue
|
||||
|
||||
if not all_hits:
|
||||
return []
|
||||
|
||||
# 按分数排序取 top
|
||||
all_hits.sort(key=lambda h: h.score, reverse=True)
|
||||
all_hits = all_hits[:max_results]
|
||||
|
||||
# 从数据库获取标注详情
|
||||
annotation_ids = [h.id for h in all_hits]
|
||||
score_map = {h.id: h.score for h in all_hits}
|
||||
|
||||
ann_result = await self._db.execute(
|
||||
select(KnowledgeAnnotation).where(
|
||||
KnowledgeAnnotation.id.in_(annotation_ids),
|
||||
KnowledgeAnnotation.is_deleted == False,
|
||||
KnowledgeAnnotation.enabled == True,
|
||||
)
|
||||
)
|
||||
annotations = {str(a.id): a for a in ann_result.scalars().all()}
|
||||
|
||||
results = []
|
||||
for hit in all_hits:
|
||||
ann = annotations.get(hit.id)
|
||||
if not ann:
|
||||
continue
|
||||
# 标注结果:content 返回 answer,segment_id 用 annotation id
|
||||
results.append(RetrievalResult(
|
||||
segment_id=str(ann.id),
|
||||
document_id='',
|
||||
document_name='[Q&A]',
|
||||
knowledge_base_id=str(ann.knowledge_base_id),
|
||||
content=ann.answer,
|
||||
score=round(score_map.get(hit.id, 0.0), 4),
|
||||
token_count=0,
|
||||
metadata={'type': 'annotation', 'question': ann.question},
|
||||
))
|
||||
|
||||
# 更新标注命中次数
|
||||
if annotation_ids:
|
||||
try:
|
||||
from sqlalchemy import update
|
||||
await self._db.execute(
|
||||
update(KnowledgeAnnotation)
|
||||
.where(KnowledgeAnnotation.id.in_(annotation_ids))
|
||||
.values(hit_count=func.coalesce(KnowledgeAnnotation.hit_count, 0) + 1)
|
||||
)
|
||||
await self._db.commit()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return results
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f'标注匹配失败: {e}')
|
||||
return []
|
||||
@@ -0,0 +1,45 @@
|
||||
"""
|
||||
向量存储模块
|
||||
|
||||
提供可插拔的向量存储后端,支持 Qdrant 等专业向量数据库。
|
||||
与业务数据库完全解耦,segment 表只存业务数据,向量数据存在向量数据库中。
|
||||
"""
|
||||
from ai_platform.knowledge.vector_store.base import BaseVectorStore, VectorPoint, VectorSearchResult
|
||||
from ai_platform.knowledge.vector_store.qdrant_store import QdrantVectorStore
|
||||
|
||||
__all__ = [
|
||||
'BaseVectorStore',
|
||||
'VectorPoint',
|
||||
'VectorSearchResult',
|
||||
'QdrantVectorStore',
|
||||
'get_vector_store',
|
||||
]
|
||||
|
||||
# 单例缓存
|
||||
_vector_store_instance: BaseVectorStore | None = None
|
||||
|
||||
|
||||
def get_vector_store() -> BaseVectorStore:
|
||||
"""
|
||||
工厂函数:根据配置获取向量存储实例(单例)
|
||||
"""
|
||||
global _vector_store_instance
|
||||
if _vector_store_instance is not None:
|
||||
return _vector_store_instance
|
||||
|
||||
from app.config import settings
|
||||
|
||||
store_type = getattr(settings, 'VECTOR_STORE_TYPE', 'qdrant')
|
||||
|
||||
if store_type == 'qdrant':
|
||||
_vector_store_instance = QdrantVectorStore(
|
||||
host=getattr(settings, 'QDRANT_HOST', 'localhost'),
|
||||
port=getattr(settings, 'QDRANT_PORT', 6333),
|
||||
api_key=getattr(settings, 'QDRANT_API_KEY', None),
|
||||
grpc_port=getattr(settings, 'QDRANT_GRPC_PORT', 6334),
|
||||
prefer_grpc=getattr(settings, 'QDRANT_PREFER_GRPC', False),
|
||||
)
|
||||
else:
|
||||
raise ValueError(f'不支持的向量存储类型: {store_type}')
|
||||
|
||||
return _vector_store_instance
|
||||
@@ -0,0 +1,138 @@
|
||||
"""
|
||||
向量存储抽象基类
|
||||
|
||||
定义向量存储的统一接口,所有向量存储后端必须实现这些方法。
|
||||
"""
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class VectorPoint:
|
||||
"""向量数据点"""
|
||||
id: str
|
||||
vector: List[float]
|
||||
payload: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class VectorSearchResult:
|
||||
"""向量搜索结果"""
|
||||
id: str
|
||||
score: float
|
||||
payload: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
class BaseVectorStore(ABC):
|
||||
"""
|
||||
向量存储抽象基类
|
||||
|
||||
每个知识库对应一个 collection,collection 名称格式: kb_{knowledge_base_id}
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def collection_name(knowledge_base_id: str) -> str:
|
||||
"""生成 collection 名称"""
|
||||
return f"kb_{knowledge_base_id.replace('-', '_')}"
|
||||
|
||||
@abstractmethod
|
||||
async def ensure_collection(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
vector_size: int,
|
||||
) -> None:
|
||||
"""
|
||||
确保 collection 存在,不存在则创建
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
vector_size: 向量维度
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete_collection(self, knowledge_base_id: str) -> None:
|
||||
"""
|
||||
删除 collection(删除知识库时调用)
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def upsert(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
points: List[VectorPoint],
|
||||
) -> None:
|
||||
"""
|
||||
批量写入/更新向量
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
points: 向量数据点列表
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
point_ids: List[str],
|
||||
) -> None:
|
||||
"""
|
||||
批量删除向量
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
point_ids: 要删除的向量 ID 列表
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def search(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
query_vector: List[float],
|
||||
top_k: int = 10,
|
||||
score_threshold: float = 0.0,
|
||||
filter_conditions: Optional[Dict[str, Any]] = None,
|
||||
) -> List[VectorSearchResult]:
|
||||
"""
|
||||
向量相似度搜索
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
query_vector: 查询向量
|
||||
top_k: 返回数量
|
||||
score_threshold: 最低相似度阈值
|
||||
filter_conditions: 过滤条件(如 {"document_id": "xxx"})
|
||||
|
||||
Returns:
|
||||
搜索结果列表,按相似度降序排列
|
||||
"""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def delete_by_filter(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
filter_conditions: Dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
按条件删除向量(如删除某个文档的所有向量)
|
||||
|
||||
Args:
|
||||
knowledge_base_id: 知识库 ID
|
||||
filter_conditions: 过滤条件(如 {"document_id": "xxx"})
|
||||
"""
|
||||
...
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""健康检查"""
|
||||
return True
|
||||
@@ -0,0 +1,260 @@
|
||||
"""
|
||||
Qdrant 向量存储实现
|
||||
|
||||
使用 Qdrant 作为向量数据库后端,通过 qdrant-client 进行交互。
|
||||
每个知识库对应一个 Qdrant collection。
|
||||
"""
|
||||
import logging
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from qdrant_client import AsyncQdrantClient
|
||||
from qdrant_client.models import (
|
||||
Distance,
|
||||
FieldCondition,
|
||||
Filter,
|
||||
FilterSelector,
|
||||
MatchValue,
|
||||
PointIdsList,
|
||||
PointStruct,
|
||||
VectorParams,
|
||||
)
|
||||
|
||||
from ai_platform.knowledge.vector_store.base import (
|
||||
BaseVectorStore,
|
||||
VectorPoint,
|
||||
VectorSearchResult,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class QdrantVectorStore(BaseVectorStore):
|
||||
"""
|
||||
Qdrant 向量存储
|
||||
|
||||
特性:
|
||||
- 高性能向量检索(HNSW 索引)
|
||||
- 支持 payload 过滤
|
||||
- 支持 REST 和 gRPC 协议
|
||||
- 与业务数据库完全解耦
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
host: str = "localhost",
|
||||
port: int = 6333,
|
||||
api_key: Optional[str] = None,
|
||||
grpc_port: int = 6334,
|
||||
prefer_grpc: bool = False,
|
||||
):
|
||||
self._host = host
|
||||
self._port = port
|
||||
self._api_key = api_key
|
||||
self._grpc_port = grpc_port
|
||||
self._prefer_grpc = prefer_grpc
|
||||
self._client: Optional[AsyncQdrantClient] = None
|
||||
|
||||
async def _get_client(self) -> AsyncQdrantClient:
|
||||
"""获取或创建 Qdrant 客户端(懒初始化)"""
|
||||
if self._client is None:
|
||||
# 如果 host 已包含协议前缀,直接作为 url 使用;
|
||||
# 否则拼接 http:// 避免 qdrant-client 对非 localhost 域名自动走 HTTPS
|
||||
if self._host.startswith("http://") or self._host.startswith("https://"):
|
||||
url = f"{self._host}:{self._port}"
|
||||
else:
|
||||
url = f"http://{self._host}:{self._port}"
|
||||
self._client = AsyncQdrantClient(
|
||||
url=url,
|
||||
api_key=self._api_key,
|
||||
grpc_port=self._grpc_port,
|
||||
prefer_grpc=self._prefer_grpc,
|
||||
timeout=30,
|
||||
)
|
||||
return self._client
|
||||
|
||||
async def ensure_collection(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
vector_size: int,
|
||||
) -> None:
|
||||
"""确保 collection 存在且维度匹配"""
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
|
||||
collections = await client.get_collections()
|
||||
existing_names = {c.name for c in collections.collections}
|
||||
|
||||
if name in existing_names:
|
||||
# 检查已有 collection 的维度是否匹配
|
||||
info = await client.get_collection(collection_name=name)
|
||||
existing_size = info.config.params.vectors.size
|
||||
if existing_size != vector_size:
|
||||
logger.warning(
|
||||
f"Qdrant collection {name} 维度不匹配: "
|
||||
f"已有={existing_size}, 期望={vector_size},删除重建"
|
||||
)
|
||||
await client.delete_collection(collection_name=name)
|
||||
else:
|
||||
return
|
||||
|
||||
await client.create_collection(
|
||||
collection_name=name,
|
||||
vectors_config=VectorParams(
|
||||
size=vector_size,
|
||||
distance=Distance.COSINE,
|
||||
),
|
||||
)
|
||||
# 创建 payload 索引,加速过滤查询
|
||||
await client.create_payload_index(
|
||||
collection_name=name,
|
||||
field_name="document_id",
|
||||
field_schema="keyword",
|
||||
)
|
||||
logger.info(f"Qdrant collection 已创建: {name} (dim={vector_size})")
|
||||
|
||||
async def delete_collection(self, knowledge_base_id: str) -> None:
|
||||
"""删除 collection"""
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
try:
|
||||
await client.delete_collection(collection_name=name)
|
||||
logger.info(f"Qdrant collection 已删除: {name}")
|
||||
except Exception as e:
|
||||
logger.warning(f"删除 Qdrant collection 失败: {name}, {e}")
|
||||
|
||||
async def upsert(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
points: List[VectorPoint],
|
||||
) -> None:
|
||||
"""批量写入/更新向量"""
|
||||
if not points:
|
||||
return
|
||||
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
|
||||
qdrant_points = [
|
||||
PointStruct(
|
||||
id=self._to_uuid(p.id),
|
||||
vector=p.vector,
|
||||
payload={**p.payload, 'segment_id': p.id},
|
||||
)
|
||||
for p in points
|
||||
]
|
||||
|
||||
# Qdrant 单次 upsert 建议不超过 100 个点
|
||||
batch_size = 100
|
||||
for i in range(0, len(qdrant_points), batch_size):
|
||||
batch = qdrant_points[i:i + batch_size]
|
||||
await client.upsert(
|
||||
collection_name=name,
|
||||
points=batch,
|
||||
)
|
||||
|
||||
async def delete(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
point_ids: List[str],
|
||||
) -> None:
|
||||
"""批量删除向量"""
|
||||
if not point_ids:
|
||||
return
|
||||
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
|
||||
uuid_ids = [self._to_uuid(pid) for pid in point_ids]
|
||||
await client.delete(
|
||||
collection_name=name,
|
||||
points_selector=PointIdsList(points=uuid_ids),
|
||||
)
|
||||
|
||||
async def search(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
query_vector: List[float],
|
||||
top_k: int = 10,
|
||||
score_threshold: float = 0.0,
|
||||
filter_conditions: Optional[Dict[str, Any]] = None,
|
||||
) -> List[VectorSearchResult]:
|
||||
"""向量相似度搜索"""
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
|
||||
# 构建过滤条件
|
||||
query_filter = self._build_filter(filter_conditions) if filter_conditions else None
|
||||
|
||||
try:
|
||||
results = await client.search(
|
||||
collection_name=name,
|
||||
query_vector=query_vector,
|
||||
limit=top_k,
|
||||
score_threshold=score_threshold,
|
||||
query_filter=query_filter,
|
||||
with_payload=True,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Qdrant 搜索失败: {e}")
|
||||
return []
|
||||
|
||||
return [
|
||||
VectorSearchResult(
|
||||
id=(hit.payload or {}).get('segment_id', str(hit.id)),
|
||||
score=hit.score,
|
||||
payload=hit.payload or {},
|
||||
)
|
||||
for hit in results
|
||||
]
|
||||
|
||||
async def delete_by_filter(
|
||||
self,
|
||||
knowledge_base_id: str,
|
||||
filter_conditions: Dict[str, Any],
|
||||
) -> None:
|
||||
"""按条件删除向量"""
|
||||
client = await self._get_client()
|
||||
name = self.collection_name(knowledge_base_id)
|
||||
|
||||
query_filter = self._build_filter(filter_conditions)
|
||||
if query_filter:
|
||||
await client.delete(
|
||||
collection_name=name,
|
||||
points_selector=FilterSelector(filter=query_filter),
|
||||
)
|
||||
|
||||
async def health_check(self) -> bool:
|
||||
"""健康检查"""
|
||||
try:
|
||||
client = await self._get_client()
|
||||
await client.get_collections()
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Qdrant 健康检查失败: {e}")
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _to_uuid(string_id: str) -> str:
|
||||
"""将任意字符串 ID 确定性转换为 UUID5(Qdrant 要求 point ID 为 UUID 或整数)"""
|
||||
return str(uuid.uuid5(uuid.NAMESPACE_DNS, string_id))
|
||||
|
||||
@staticmethod
|
||||
def _build_filter(conditions: Dict[str, Any]) -> Optional[Filter]:
|
||||
"""构建 Qdrant 过滤条件"""
|
||||
if not conditions:
|
||||
return None
|
||||
|
||||
must = []
|
||||
for key, value in conditions.items():
|
||||
if isinstance(value, list):
|
||||
# 列表值:任一匹配(OR 语义),用 should 包裹后作为一个 must 条件
|
||||
should_conditions = [
|
||||
FieldCondition(key=key, match=MatchValue(value=v))
|
||||
for v in value
|
||||
]
|
||||
must.append(Filter(should=should_conditions))
|
||||
else:
|
||||
must.append(FieldCondition(key=key, match=MatchValue(value=value)))
|
||||
|
||||
return Filter(must=must) if must else None
|
||||
Reference in New Issue
Block a user