Build lightweight AI agent admin

This commit is contained in:
Codex
2026-06-08 18:14:59 +08:00
commit e164840f43
2530 changed files with 435693 additions and 0 deletions
@@ -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 存储 questionmetadata 中存储 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('&nbsp;', ' ')
text = text.replace('&lt;', '<')
text = text.replace('&gt;', '>')
text = text.replace('&amp;', '&')
text = text.replace('&quot;', '"')
text = text.replace('&#39;', "'")
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 APIJina/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 返回 answersegment_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):
"""
向量存储抽象基类
每个知识库对应一个 collectioncollection 名称格式: 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 确定性转换为 UUID5Qdrant 要求 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