ForcePilot/backend/package/yuxi/knowledge/implementations/milvus.py

1002 lines
42 KiB
Python
Raw Normal View History

import asyncio
import os
import time
import traceback
from functools import partial
from typing import Any
from pymilvus import (
AnnSearchRequest,
Collection,
CollectionSchema,
DataType,
FieldSchema,
Function,
FunctionType,
WeightedRanker,
connections,
db,
utility,
)
from yuxi.knowledge.base import FileStatus, KnowledgeBase
from yuxi.knowledge.chunking.ragflow_like.dispatcher import chunk_markdown
from yuxi.knowledge.utils.kb_utils import resolve_processing_params
from yuxi.repositories.knowledge_chunk_repository import KnowledgeChunkRepository
from yuxi.services.model_cache import model_cache
from yuxi.plugins.parser.unified import Parser
from yuxi.utils import hashstr, logger
from yuxi.utils.datetime_utils import utc_isoformat
MILVUS_AVAILABLE = True
CONTENT_SPARSE_FIELD = "content_sparse"
CONTENT_ANALYZER_PARAMS = {"type": "chinese"}
VECTOR_METRIC_TYPE = "COSINE"
class MilvusKB(KnowledgeBase):
"""基于 Milvus 的生产级向量库"""
def __init__(self, work_dir: str, **kwargs):
"""
初始化 Milvus 知识库
Args:
work_dir: 工作目录
**kwargs: 其他配置参数
"""
super().__init__(work_dir)
if not MILVUS_AVAILABLE:
raise ImportError("pymilvus is not installed. Please install it with: pip install pymilvus")
# Milvus 配置
# self.milvus_host = kwargs.get('milvus_host', os.getenv('MILVUS_HOST', 'localhost'))
# self.milvus_port = kwargs.get('milvus_port', int(os.getenv('MILVUS_PORT', '19530')))
self.milvus_token = kwargs.get("milvus_token", os.getenv("MILVUS_TOKEN") or "")
self.milvus_uri = kwargs.get("milvus_uri", os.getenv("MILVUS_URI") or "http://localhost:19530")
self.milvus_db = kwargs.get("milvus_db") or "yuxi_know"
# 连接名称
self.connection_alias = f"milvus_{hashstr(work_dir, 6)}"
# 存储集合映射 {db_id: Collection}
self.collections: dict[str, Any] = {}
# 元数据锁
self._metadata_lock = asyncio.Lock()
# 初始化连接
self._init_connection()
logger.info("MilvusKB initialized")
@property
def kb_type(self) -> str:
"""知识库类型标识"""
return "milvus"
def _init_connection(self):
"""初始化 Milvus 连接"""
try:
# 连接到 Milvus
connections.connect(alias=self.connection_alias, uri=self.milvus_uri, token=self.milvus_token)
# 创建数据库(如果不存在)
try:
if self.milvus_db not in db.list_database():
db.create_database(self.milvus_db)
db.using_database(self.milvus_db)
except Exception as e:
logger.warning(f"Database operation failed, using default: {e}")
logger.info(f"Connected to Milvus at {self.milvus_uri}")
except Exception as e:
logger.error(f"Failed to connect to Milvus: {e}")
raise
async def _create_kb_instance(self, db_id: str, kb_config: dict) -> Any:
"""创建 Milvus 集合"""
logger.info(f"Creating Milvus collection for {db_id}")
if not (metadata := self.databases_meta.get(db_id)):
raise ValueError(f"Database {db_id} not found")
embedding_model_spec = metadata.get("embedding_model_spec")
if not embedding_model_spec:
raise ValueError(f"Embedding model spec not found for database {db_id}")
embedding_info = model_cache.get_model_info(embedding_model_spec)
if not embedding_info or embedding_info.model_type != "embedding":
raise ValueError(f"Unsupported embedding model: {embedding_model_spec}")
collection_name = db_id
try:
# 检查集合是否存在
if utility.has_collection(collection_name, using=self.connection_alias):
collection = Collection(name=collection_name, using=self.connection_alias)
# 检查嵌入模型是否匹配
description = collection.description
expected_model = embedding_info.model_id
if expected_model not in description:
logger.warning(
f"Collection {collection_name} model mismatch: "
f"expected='{expected_model}', found_in_description='{description}'"
)
utility.drop_collection(collection_name, using=self.connection_alias)
return self._create_new_collection(collection_name, embedding_info, db_id)
if not self._collection_supports_bm25(collection):
logger.warning(f"Collection {collection_name} schema does not support BM25, recreating")
utility.drop_collection(collection_name, using=self.connection_alias)
return self._create_new_collection(collection_name, embedding_info, db_id)
logger.info(f"Retrieved existing collection: {collection_name}")
return collection
else:
logger.info(f"Collection {collection_name} not found, creating new one")
return self._create_new_collection(collection_name, embedding_info, db_id)
except (connections.MilvusException, RuntimeError) as e:
logger.error(f"Error checking collection {collection_name}: {e}")
raise
except Exception as e:
logger.error(f"Unexpected error while managing collection {collection_name}: {e}")
logger.debug(f"Traceback: {traceback.format_exc()}")
raise
def _create_new_collection(self, collection_name: str, embedding_info: Any, db_id: str) -> Collection:
"""创建新的 Milvus 集合"""
embedding_dim = embedding_info.dimension or 1024
model_name = embedding_info.model_id
# 定义集合Schema
fields = [
FieldSchema(name="id", dtype=DataType.VARCHAR, max_length=100, is_primary=True),
FieldSchema(
name="content",
dtype=DataType.VARCHAR,
max_length=65535,
enable_analyzer=True,
analyzer_params=CONTENT_ANALYZER_PARAMS,
),
FieldSchema(name="chunk_id", dtype=DataType.VARCHAR, max_length=100),
FieldSchema(name="file_id", dtype=DataType.VARCHAR, max_length=100),
FieldSchema(name="chunk_index", dtype=DataType.INT64),
FieldSchema(name="embedding", dtype=DataType.FLOAT_VECTOR, dim=embedding_dim),
FieldSchema(name=CONTENT_SPARSE_FIELD, dtype=DataType.SPARSE_FLOAT_VECTOR),
]
bm25_function = Function(
name="content_bm25",
input_field_names=["content"],
output_field_names=[CONTENT_SPARSE_FIELD],
function_type=FunctionType.BM25,
)
schema = CollectionSchema(
fields=fields,
description=f"Knowledge base collection for {db_id} using {model_name}",
functions=[bm25_function],
)
# 创建集合
collection = Collection(name=collection_name, schema=schema, using=self.connection_alias)
# 创建索引
index_params = {"metric_type": VECTOR_METRIC_TYPE, "index_type": "IVF_FLAT", "params": {"nlist": 1024}}
collection.create_index("embedding", index_params)
sparse_index_params = {
"metric_type": "BM25",
"index_type": "SPARSE_INVERTED_INDEX",
"params": {"inverted_index_algo": "DAAT_MAXSCORE"},
}
collection.create_index(CONTENT_SPARSE_FIELD, sparse_index_params)
logger.info(f"Created new Milvus collection: {collection_name} '{model_name=}', {embedding_dim=}")
return collection
def _collection_supports_bm25(self, collection: Collection) -> bool:
"""检查集合是否具备 Milvus 内置 BM25 所需的 schema。"""
fields = {field.name: field for field in collection.schema.fields}
content_field = fields.get("content")
sparse_field = fields.get(CONTENT_SPARSE_FIELD)
if not content_field or content_field.dtype != DataType.VARCHAR:
return False
if content_field.params.get("enable_analyzer") is not True:
return False
if not sparse_field or sparse_field.dtype != DataType.SPARSE_FLOAT_VECTOR:
return False
for function in collection.schema.functions:
if (
function.type == FunctionType.BM25
and function.input_field_names == ["content"]
and function.output_field_names == [CONTENT_SPARSE_FIELD]
):
return True
return False
async def _initialize_kb_instance(self, instance: Any) -> None:
"""初始化 Milvus 集合(加载到内存)"""
try:
instance.load()
logger.info("Milvus collection loaded into memory")
except Exception as e:
logger.warning(f"Failed to load collection into memory: {e}")
def _get_embedding_function(self, embedding_model_spec: str, *, sync: bool = False):
"""获取 embedding 编码函数。sync=True 返回同步版本,否则返回异步版本。"""
from yuxi.models.embed import select_embedding_model
model = select_embedding_model(embedding_model_spec)
batch_size = int(getattr(model, "batch_size", 40) or 40)
method = model.batch_encode if sync else model.abatch_encode
return partial(method, batch_size=batch_size)
async def _get_milvus_collection(self, db_id: str):
"""获取或创建 Milvus 集合"""
if db_id in self.collections:
return self.collections[db_id]
if db_id not in self.databases_meta:
return None
try:
# 创建集合
collection = await self._create_kb_instance(db_id, {})
await self._initialize_kb_instance(collection)
self.collections[db_id] = collection
return collection
except Exception as e:
logger.error(f"Failed to create Milvus collection for {db_id}: {e}")
logger.error(f"Traceback: {traceback.format_exc()}")
return None
def _split_text_into_chunks(self, text: str, file_id: str, filename: str, params: dict) -> list[dict]:
"""将文本分割成块"""
return chunk_markdown(text, file_id, filename, params)
def _build_chunk_pg_records(self, db_id: str, chunks: list[dict]) -> list[dict[str, Any]]:
return [
{
"chunk_id": chunk["chunk_id"],
"file_id": chunk["file_id"],
"db_id": db_id,
"chunk_index": chunk["chunk_index"],
"content": chunk["content"],
"start_char_pos": chunk.get("start_char_pos"),
"end_char_pos": chunk.get("end_char_pos"),
"start_token_pos": chunk.get("start_token_pos"),
"end_token_pos": chunk.get("end_token_pos"),
"graph_indexed": bool(chunk.get("graph_indexed", False)),
"ent_ids": chunk.get("ent_ids"),
"tags": chunk.get("tags"),
"extraction_result": chunk.get("extraction_result"),
}
for chunk in chunks
]
async def _insert_chunks_to_stores(
self,
db_id: str,
file_id: str,
collection: Collection,
chunks: list[dict],
embeddings: list,
) -> None:
if not chunks:
return
entities = [
[chunk["id"] for chunk in chunks],
[chunk["content"] for chunk in chunks],
[chunk["chunk_id"] for chunk in chunks],
[chunk["file_id"] for chunk in chunks],
[chunk["chunk_index"] for chunk in chunks],
embeddings,
]
chunk_repo = KnowledgeChunkRepository()
def _insert_milvus_records():
collection.insert(entities)
pg_task = chunk_repo.batch_upsert(self._build_chunk_pg_records(db_id, chunks))
milvus_task = asyncio.to_thread(_insert_milvus_records)
results = await asyncio.gather(pg_task, milvus_task, return_exceptions=True)
errors = [result for result in results if isinstance(result, Exception)]
if not errors:
return
logger.error(f"Chunk double-write failed for file {file_id}, rolling back PostgreSQL and Milvus chunks")
try:
await chunk_repo.delete_by_file_id(file_id)
except Exception as cleanup_error:
logger.error(f"Failed to rollback PostgreSQL chunks for {file_id}: {cleanup_error}")
try:
await self._delete_file_chunks_from_milvus(collection, file_id)
except Exception as cleanup_error:
logger.error(f"Failed to rollback Milvus chunks for {file_id}: {cleanup_error}")
raise errors[0]
async def _delete_file_chunks_from_milvus(self, collection: Collection, file_id: str) -> None:
expr = f'file_id == "{file_id}"'
results = collection.query(expr=expr, output_fields=["id"], limit=1)
if not results:
logger.info(f"File {file_id} not found in Milvus, skipping delete operation")
return
def _delete_from_milvus():
collection.delete(expr)
logger.info(f"Deleted chunks for file {file_id} from Milvus")
await asyncio.to_thread(_delete_from_milvus)
def _get_filename_for_file_id(self, file_id: str | None) -> str:
if not file_id:
return "未知来源"
return self.files_meta.get(file_id, {}).get("filename") or "未知来源"
def _build_file_name_expr(self, db_id: str, file_name: str | None) -> str | None:
if not file_name:
return None
file_name_pattern = file_name.replace("%", "")
matched_file_ids = [
file_id
for file_id, file_meta in self.files_meta.items()
if file_meta.get("database_id") == db_id and file_name_pattern in (file_meta.get("filename") or "")
]
if not matched_file_ids:
return 'file_id == "__no_matching_file__"'
escaped_ids = [file_id.replace('"', '\\"') for file_id in matched_file_ids]
if len(escaped_ids) == 1:
return f'file_id == "{escaped_ids[0]}"'
joined_ids = '", "'.join(escaped_ids)
return f'file_id in ["{joined_ids}"]'
2026-04-27 02:24:52 +08:00
async def index_file(
self, db_id: str, file_id: str, operator_id: str | None = None, params: dict | None = None
) -> dict:
"""
Index parsed file (Status: INDEXING -> INDEXED/ERROR_INDEXING)
Args:
db_id: Database ID
file_id: File ID
operator_id: ID of the user performing the operation
params: Override processing params to apply during indexing (merged on top of stored params)
Returns:
Updated file metadata
"""
if db_id not in self.databases_meta:
raise ValueError(f"Database {db_id} not found")
# Get/Create collection
collection = await self._get_milvus_collection(db_id)
if not collection:
raise ValueError(f"Failed to get Milvus collection for {db_id}")
embedding_model_spec = self.databases_meta[db_id].get("embedding_model_spec")
embedding_function = self._get_embedding_function(embedding_model_spec)
# Get file meta
async with self._metadata_lock:
if file_id not in self.files_meta:
raise ValueError(f"File {file_id} not found")
file_meta = self.files_meta[file_id]
# Validate current status - only allow indexing from these states
current_status = file_meta.get("status")
allowed_statuses = {
FileStatus.PARSED,
FileStatus.ERROR_INDEXING,
FileStatus.INDEXED, # For re-indexing
"done", # Legacy status
}
if current_status not in allowed_statuses:
raise ValueError(
f"Cannot index file with status '{current_status}'. "
f"File must be parsed first (status should be one of: {', '.join(allowed_statuses)})"
)
# Check markdown file exists
if not file_meta.get("markdown_file"):
await self._mark_file_unparsed(file_id, operator_id)
raise ValueError("File has not been parsed yet (no markdown_file)")
# Clear previous error if any
if "error" in file_meta:
self.files_meta[file_id].pop("error", None)
# Update status and add to processing queue
self.files_meta[file_id]["status"] = FileStatus.INDEXING
self.files_meta[file_id]["updated_at"] = utc_isoformat()
if operator_id:
self.files_meta[file_id]["updated_by"] = operator_id
# 将传入的 params 作为 request_params确保用户指定的参数始终覆盖存储的参数
params = resolve_processing_params(
kb_additional_params=self.databases_meta.get(db_id, {}).get("metadata"),
file_processing_params=file_meta.get("processing_params"),
request_params=params,
)
self.files_meta[file_id]["processing_params"] = params
await self._save_metadata()
logger.debug(f"[index_file] file_id={file_id}, processing_params={params}")
# Add to processing queue
self._add_to_processing_queue(file_id)
try:
# Read markdown
markdown_content = await self._read_markdown_from_minio(file_meta["markdown_file"])
filename = file_meta.get("filename")
# Split
chunks = self._split_text_into_chunks(markdown_content, file_id, filename, params)
logger.info(
f"Split {filename} into {len(chunks)} chunks with params: "
f"chunk_preset_id={params.get('chunk_preset_id')}, "
f"chunk_parser_config={params.get('chunk_parser_config')}"
)
if chunks:
texts = [chunk["content"] for chunk in chunks]
embeddings = await embedding_function(texts)
# Clean up existing chunks if any (for re-indexing)
await self.delete_file_chunks_only(db_id, file_id)
await self._insert_chunks_to_stores(db_id, file_id, collection, chunks, embeddings)
logger.info(f"Indexed file {file_id} into Milvus")
# Update status
async with self._metadata_lock:
self.files_meta[file_id]["status"] = FileStatus.INDEXED
self.files_meta[file_id]["updated_at"] = utc_isoformat()
if operator_id:
self.files_meta[file_id]["updated_by"] = operator_id
await self._persist_file(file_id)
return self.files_meta[file_id]
except Exception as e:
logger.error(f"Indexing failed for {file_id}: {e}")
async with self._metadata_lock:
self.files_meta[file_id]["status"] = FileStatus.ERROR_INDEXING
self.files_meta[file_id]["error"] = str(e)
self.files_meta[file_id]["updated_at"] = utc_isoformat()
if operator_id:
self.files_meta[file_id]["updated_by"] = operator_id
await self._persist_file(file_id)
raise
finally:
# Remove from processing queue
self._remove_from_processing_queue(file_id)
async def update_content(self, db_id: str, file_ids: list[str], params: dict | None = None) -> list[dict]:
"""更新内容 - 根据file_ids重新解析文件并更新向量库"""
if db_id not in self.databases_meta:
raise ValueError(f"Database {db_id} not found")
collection = await self._get_milvus_collection(db_id)
if not collection:
raise ValueError(f"Failed to get Milvus collection for {db_id}")
embedding_model_spec = self.databases_meta[db_id].get("embedding_model_spec")
embedding_function = self._get_embedding_function(embedding_model_spec)
# 处理默认参数
if params is None:
params = {}
processed_items_info = []
for file_id in file_ids:
# 从元数据中获取文件信息
async with self._metadata_lock:
if file_id not in self.files_meta:
logger.warning(f"File {file_id} not found in metadata, skipping")
continue
file_meta = self.files_meta[file_id]
file_path = file_meta.get("path")
filename = file_meta.get("filename")
if not file_path:
logger.warning(f"File path not found for {file_id}, skipping")
continue
# 添加到处理队列
self._add_to_processing_queue(file_id)
try:
# 更新状态为处理中
async with self._metadata_lock:
resolved_params = resolve_processing_params(
kb_additional_params=self.databases_meta.get(db_id, {}).get("metadata"),
file_processing_params=self.files_meta[file_id].get("processing_params"),
request_params=params,
)
self.files_meta[file_id]["processing_params"] = resolved_params
self.files_meta[file_id]["status"] = FileStatus.INDEXING
await self._persist_file(file_id)
# 重新解析文件为 markdown
parse_params = {**resolved_params, "image_bucket": "public", "image_prefix": f"{db_id}/kb-images"}
markdown_content = await Parser.aparse(source=file_path, params=parse_params)
# 重新生成 chunks
chunks = self._split_text_into_chunks(markdown_content, file_id, filename, resolved_params)
logger.info(f"Split {filename} into {len(chunks)} chunks")
# 先删除现有 chunks保留文件元数据
await self.delete_file_chunks_only(db_id, file_id)
if chunks:
texts = [chunk["content"] for chunk in chunks]
embeddings = await embedding_function(texts)
await self._insert_chunks_to_stores(db_id, file_id, collection, chunks, embeddings)
logger.info(f"Updated file {file_path} in Milvus. Done.")
# 更新元数据状态
async with self._metadata_lock:
self.files_meta[file_id]["status"] = "done"
await self._persist_file(file_id)
# 从处理队列中移除
self._remove_from_processing_queue(file_id)
# 返回更新后的文件信息
updated_file_meta = file_meta.copy()
updated_file_meta["status"] = "done"
updated_file_meta["file_id"] = file_id
processed_items_info.append(updated_file_meta)
except Exception as e:
logger.error(f"更新file {file_path} 失败: {e}, {traceback.format_exc()}")
async with self._metadata_lock:
self.files_meta[file_id]["status"] = "failed"
await self._persist_file(file_id)
# 从处理队列中移除
self._remove_from_processing_queue(file_id)
# 返回失败的文件信息
failed_file_meta = file_meta.copy()
failed_file_meta["status"] = "failed"
failed_file_meta["file_id"] = file_id
processed_items_info.append(failed_file_meta)
return processed_items_info
def _build_chunk_from_hit(
self,
hit: Any,
score: float,
include_distances: bool,
score_field: str | None = None,
) -> dict:
"""将 Milvus Hit 转成知识库统一返回的 Chunk 结构。"""
entity = hit.entity
file_id = entity.get("file_id")
metadata = {
"source": self._get_filename_for_file_id(file_id),
"chunk_id": entity.get("chunk_id"),
"file_id": file_id,
"chunk_index": entity.get("chunk_index"),
}
chunk = {"content": entity.get("content", ""), "metadata": metadata, "score": float(score or 0.0)}
if score_field:
chunk[score_field] = float(score or 0.0)
if include_distances:
chunk["distance"] = hit.distance
return chunk
async def aquery(self, query_text: str, db_id: str, agent_call: bool = False, **kwargs) -> list[dict]:
"""异步查询知识库"""
collection = await self._get_milvus_collection(db_id)
if not collection:
raise ValueError(f"Database {db_id} not found")
query_params = self._get_query_params(db_id)
# 合并查询参数kwargs临时参数优先级高于 query_params持久化参数
# 这样允许用户在单次查询中临时覆盖持久化配置
merged_kwargs = {**query_params, **kwargs}
try:
# 查询参数(从 merged_kwargs 读取)
logger.debug(f"Query params: {merged_kwargs}")
final_top_k = int(merged_kwargs.get("final_top_k", 10))
final_top_k = max(final_top_k, 1)
similarity_threshold = float(merged_kwargs.get("similarity_threshold", 0.2))
metric_type = VECTOR_METRIC_TYPE
include_distances = bool(merged_kwargs.get("include_distances", True))
search_mode = str(merged_kwargs.get("search_mode", "vector")).lower()
if search_mode not in {"vector", "keyword", "hybrid"}:
search_mode = "vector"
use_reranker = bool(merged_kwargs.get("use_reranker", False))
if use_reranker:
recall_top_k = int(merged_kwargs.get("recall_top_k", 50))
recall_top_k = max(recall_top_k, final_top_k)
else:
recall_top_k = final_top_k
file_expr = self._build_file_name_expr(db_id, merged_kwargs.get("file_name"))
if file_expr:
logger.debug(f"Using filter expression: {file_expr}")
output_fields = ["content", "chunk_id", "file_id", "chunk_index"]
retrieved_chunks: list[dict] = []
if search_mode == "vector":
embedding_model_spec = self.databases_meta[db_id].get("embedding_model_spec")
embedding_function = self._get_embedding_function(embedding_model_spec, sync=True)
query_embedding = embedding_function([query_text])
search_params = {"metric_type": metric_type, "params": {"nprobe": 10}}
results = collection.search(
data=query_embedding,
anns_field="embedding",
param=search_params,
limit=recall_top_k,
expr=file_expr,
output_fields=output_fields,
)
if results and len(results) > 0 and len(results[0]) > 0:
for hit in results[0]:
similarity = hit.distance if metric_type == VECTOR_METRIC_TYPE else 1 / (1 + hit.distance)
if similarity < similarity_threshold:
continue
retrieved_chunks.append(self._build_chunk_from_hit(hit, similarity, include_distances))
logger.debug(
f"Milvus vector query response: {len(retrieved_chunks)} chunks found (after similarity filtering)"
)
elif search_mode == "keyword":
bm25_top_k = int(merged_kwargs.get("bm25_top_k", recall_top_k))
bm25_top_k = max(bm25_top_k, 1)
bm25_drop_ratio_search = float(merged_kwargs.get("bm25_drop_ratio_search", 0.0))
bm25_search_params = {
"metric_type": "BM25",
"params": {"drop_ratio_search": bm25_drop_ratio_search},
}
results = collection.search(
data=[query_text],
anns_field=CONTENT_SPARSE_FIELD,
param=bm25_search_params,
limit=bm25_top_k,
expr=file_expr,
output_fields=output_fields,
)
if results and len(results) > 0 and len(results[0]) > 0:
for hit in results[0]:
retrieved_chunks.append(
self._build_chunk_from_hit(hit, hit.distance, include_distances, score_field="bm25_score")
)
logger.debug(f"Milvus BM25 query response: {len(retrieved_chunks)} chunks found")
else:
embedding_model_spec = self.databases_meta[db_id].get("embedding_model_spec")
embedding_function = self._get_embedding_function(embedding_model_spec, sync=True)
query_embedding = embedding_function([query_text])
bm25_top_k = int(merged_kwargs.get("bm25_top_k", recall_top_k))
bm25_top_k = max(bm25_top_k, 1)
bm25_drop_ratio_search = float(merged_kwargs.get("bm25_drop_ratio_search", 0.0))
vector_weight = float(merged_kwargs.get("vector_weight", 0.7))
bm25_weight = float(merged_kwargs.get("bm25_weight", 0.3))
vector_request = AnnSearchRequest(
data=query_embedding,
anns_field="embedding",
param={"metric_type": metric_type, "params": {"nprobe": 10}},
limit=recall_top_k,
expr=file_expr,
)
bm25_request = AnnSearchRequest(
data=[query_text],
anns_field=CONTENT_SPARSE_FIELD,
param={
"metric_type": "BM25",
"params": {"drop_ratio_search": bm25_drop_ratio_search},
},
limit=bm25_top_k,
expr=file_expr,
)
results = collection.hybrid_search(
reqs=[vector_request, bm25_request],
rerank=WeightedRanker(vector_weight, bm25_weight),
limit=recall_top_k,
output_fields=output_fields,
)
if results and len(results) > 0 and len(results[0]) > 0:
for hit in results[0]:
score = float(hit.distance or 0.0)
if score < similarity_threshold:
continue
retrieved_chunks.append(
self._build_chunk_from_hit(hit, score, include_distances, score_field="hybrid_score")
)
logger.debug(f"Milvus hybrid query response: {len(retrieved_chunks)} chunks found")
if not retrieved_chunks:
return []
if not use_reranker:
return retrieved_chunks[:final_top_k]
# 使用重排序模型
reranker_model = merged_kwargs.get("reranker_model")
if not reranker_model:
raise ValueError(
"Reranker model must be specified when use_reranker=True. "
"Please provide reranker_model in query parameters."
)
try:
from yuxi.models.rerank import get_reranker
reranker = get_reranker(reranker_model)
try:
rerank_start = time.time()
documents_text = [chunk["content"] for chunk in retrieved_chunks]
rerank_scores = await reranker.acompute_score([query_text, documents_text], normalize=True)
for chunk, rerank_score in zip(retrieved_chunks, rerank_scores):
chunk["rerank_score"] = float(rerank_score)
retrieved_chunks.sort(
key=lambda item: item.get("rerank_score", item.get("score", 0.0)), reverse=True
)
elapsed = time.time() - rerank_start
2026-01-09 01:23:00 +08:00
logger.info(f"Reranking completed for {db_id} in {elapsed:.3f}s with model {reranker_model}")
finally:
await reranker.aclose()
except Exception as exc: # noqa: BLE001
logger.error(f"Reranking failed: {exc}, falling back to vector scores")
# 统一返回结果
return retrieved_chunks[:final_top_k]
except Exception as e:
logger.error(f"Milvus query error: {e}, {traceback.format_exc()}")
return []
async def delete_file_chunks_only(self, db_id: str, file_id: str) -> None:
"""仅删除文件的chunks数据保留元数据用于更新操作"""
await KnowledgeChunkRepository().delete_by_file_id(file_id)
collection = await self._get_milvus_collection(db_id)
if collection:
# 先查询文件是否存在,避免不必要的删除操作
try:
await self._delete_file_chunks_from_milvus(collection, file_id)
except Exception as e:
logger.error(f"Error checking file existence in Milvus: {e}")
# 注意:这里不删除 files_meta[file_id],保留元数据用于后续操作
async def delete_file(self, db_id: str, file_id: str) -> None:
"""删除文件(包括元数据)"""
# 先删除 Milvus 中的 chunks 数据
await self.delete_file_chunks_only(db_id, file_id)
# 使用锁确保元数据操作的原子性
async with self._metadata_lock:
if file_id in self.files_meta:
del self.files_meta[file_id]
from yuxi.repositories.knowledge_file_repository import KnowledgeFileRepository
await KnowledgeFileRepository().delete(file_id)
async def get_file_basic_info(self, db_id: str, file_id: str) -> dict:
"""获取文件基本信息(仅元数据)"""
if file_id not in self.files_meta:
raise Exception(f"File not found: {file_id}")
return {"meta": self.files_meta[file_id]}
async def get_file_content(self, db_id: str, file_id: str) -> dict:
"""获取文件内容信息chunks和lines"""
if file_id not in self.files_meta:
raise Exception(f"File not found: {file_id}")
content_info = {"lines": []}
try:
chunks = await KnowledgeChunkRepository().list_by_file_id(file_id)
content_info["lines"] = [
{
"id": chunk.chunk_id,
"content": chunk.content,
"chunk_order_index": chunk.chunk_index,
"start_char_pos": chunk.start_char_pos,
"end_char_pos": chunk.end_char_pos,
"start_token_pos": chunk.start_token_pos,
"end_token_pos": chunk.end_token_pos,
"graph_indexed": chunk.graph_indexed,
"ent_ids": chunk.ent_ids,
"tags": chunk.tags,
"extraction_result": chunk.extraction_result,
}
for chunk in chunks
]
except Exception as e:
logger.error(f"Failed to get file content from PostgreSQL: {e}")
if not content_info["lines"]:
logger.warning(f"No chunks found in PostgreSQL for file {file_id}, file may not have been indexed")
# Try to read markdown content if available
file_meta = self.files_meta[file_id]
if file_meta.get("markdown_file"):
try:
content = await self._read_markdown_from_minio(file_meta["markdown_file"])
content_info["content"] = content
except Exception as e:
logger.error(f"Failed to read markdown file for {file_id}: {e}")
return content_info
async def get_file_info(self, db_id: str, file_id: str) -> dict:
"""获取文件完整信息(基本信息+内容信息)"""
if file_id not in self.files_meta:
raise Exception(f"File not found: {file_id}")
# 合并基本信息和内容信息
basic_info = await self.get_file_basic_info(db_id, file_id)
content_info = await self.get_file_content(db_id, file_id)
return {**basic_info, **content_info}
def delete_database(self, db_id: str) -> dict:
"""删除数据库同时清除Milvus中的集合"""
# Drop Milvus collection
try:
if utility.has_collection(db_id, using=self.connection_alias):
utility.drop_collection(db_id, using=self.connection_alias)
logger.info(f"Dropped Milvus collection for {db_id}")
else:
logger.info(f"Milvus collection {db_id} does not exist, skipping")
except Exception as e:
logger.error(f"Failed to drop Milvus collection {db_id}: {e}")
# Call base method to delete local files and metadata
return super().delete_database(db_id)
def get_query_params_config(self, db_id: str, **kwargs) -> dict:
"""获取 Milvus 知识库的查询参数配置"""
# 构建 Milvus 特定查询参数
# TODO 使用 dataclass 重新配置并解析
options = [
{
"key": "search_mode",
"label": "检索模式",
"type": "select",
"default": "vector",
"options": [
{"value": "vector", "label": "向量检索", "description": "仅使用向量相似度检索"},
{"value": "keyword", "label": "BM25 全文检索", "description": "仅使用 Milvus BM25 检索"},
{"value": "hybrid", "label": "混合检索", "description": "Milvus 向量检索与 BM25 融合检索"},
],
"description": "选择检索模式",
},
{
"key": "final_top_k",
"label": "最终返回 Chunk 数",
"type": "number",
"default": 10,
"min": 1,
"max": 100,
"description": "重排序后返回给前端的文档数量",
},
{
"key": "similarity_threshold",
"label": "相似度阈值0-1",
"type": "number",
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"description": "过滤相似度低于此值的结果",
},
{
"key": "bm25_top_k",
"label": "BM25 召回数量",
"type": "number",
"default": 50,
"min": 1,
"max": 200,
"description": "BM25 全文检索和混合检索中的 BM25 候选数量",
},
{
"key": "vector_weight",
"label": "向量检索权重",
"type": "number",
"default": 0.7,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"description": "混合检索中向量召回结果的融合权重",
},
{
"key": "bm25_weight",
"label": "BM25 权重",
"type": "number",
"default": 0.3,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"description": "混合检索中 BM25 召回结果的融合权重",
},
{
"key": "bm25_drop_ratio_search",
"label": "BM25 稀疏项丢弃比例",
"type": "number",
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"description": "BM25 检索时丢弃低分稀疏项的比例,数值越大检索越快但可能降低召回",
},
{
"key": "include_distances",
"label": "显示相似度",
"type": "boolean",
"default": True,
"description": "在结果中显示相似度分数",
},
{
"key": "use_reranker",
"label": "启用重排序",
"type": "boolean",
"default": False,
"description": "是否使用精排模型对检索结果进行重排序",
},
{
"key": "reranker_model",
"label": "重排序模型",
"type": "select",
"default": "",
"options": [
{"label": info.display_name, "value": info.spec} for info in model_cache.get_all_specs("rerank")
],
"description": "选择用于本次查询的重排序模型",
},
{
"key": "recall_top_k",
"label": "召回数量",
"type": "number",
"default": 50,
"min": 10,
"max": 200,
"description": "向量检索或混合检索保留的候选数量(启用重排序时有效)",
},
]
return {"type": "milvus", "options": options}
def __del__(self):
"""清理连接"""
try:
if hasattr(self, "connection_alias"):
connections.disconnect(self.connection_alias)
except Exception: # noqa: S110
pass