kgrag/app_pkg/database/milvus_api.py
2026-07-29 18:10:19 +08:00

861 lines
37 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import json
import os
from pymilvus import (
connections,
utility,
FieldSchema,
CollectionSchema,
DataType,
Collection,
MilvusClient,
db,
Function,
FunctionType,
AnnSearchRequest,
WeightedRanker
)
from typing import List, Dict, Union, Optional, Any
from pypinyin import lazy_pinyin
from typing import List, Dict, Tuple, Optional
class Milvus_Database:
def __init__(self, user=os.getenv("Milvus_USER", "root"), password=os.getenv("Milvus_PASSWORD", "Milvus"),
uri=os.getenv("Milvus_URI")):
self.client = MilvusClient(
uri=uri,
token=f"{user}:{password}"
)
# self.conn = connections.connect(host="172.18.127.124",port=19530,user=user,password=password)
# database = db.create_database(database_name)
# self.client.list_databases()
def create_database(self, database_name) -> List[str]:
database = self.client.create_database(database_name)
return self.client.list_databases()
def list_databases(self) -> List[str]:
return self.client.list_databases()
def using_database(self, database_name):
self.client.use_database(database_name)
def create_hierarchical_collection(self, collection_name, database_name):
self.client.use_database(database_name)
schema = MilvusClient.create_schema(
auto_id=True,
enable_dynamic_field=True,
)
analyzer_params_custom = {
"tokenizer": "jieba",
"type": "chinese",
"filter": ["cnalphanumonly"]
}
schema.add_field(field_name="id", datatype=DataType.INT64, is_primary=True)
# 层次召回中的摘要部分
schema.add_field(field_name="text", datatype=DataType.VARCHAR, max_length=65535, enable_analyzer=True,
analyzer_params=analyzer_params_custom)
# 层次召回中的完整文档
schema.add_field(field_name="origin_text", datatype=DataType.VARCHAR, max_length=65535, enable_analyzer=True,
analyzer_params=analyzer_params_custom)
schema.add_field(field_name="bm_25", datatype=DataType.SPARSE_FLOAT_VECTOR)
# Qwen3-8B-Embedding 维度4096
schema.add_field(field_name="vector", datatype=DataType.FLOAT_VECTOR, dim=4096, )
schema.add_field(field_name="resource", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY'),
schema.add_field(field_name="img_path", datatype=DataType.VARCHAR, max_length=4096, default_value="EMPTY")
schema.add_field(field_name="page_idx", datatype=DataType.INT64, default_value=0)
schema.add_field(field_name="origin_id", datatype=DataType.VARCHAR, max_length=256, default_value="EMPTY")
# schema.add_field(field_name="img_information",datatype=DataType.Json,default_value="EMPTY")
schema.add_field(field_name="header_1", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
schema.add_field(field_name="header_2", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
schema.add_field(field_name="header_3", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
bm25_function = Function(
name="bm25_emb", # Function name
input_field_names=["text"], # Name of the VARCHAR field containing raw text data
output_field_names=["bm_25"],
# Name of the SPARSE_FLOAT_VECTOR field reserved to store generated embeddings
function_type=FunctionType.BM25,
)
schema.add_function(bm25_function)
index_params = self.client.prepare_index_params()
# Add indexes
index_params.add_index(
field_name="vector",
index_name="vector_index",
index_type="HNSW", # "IVF_FLAT". AUTOINDEX
metric_type="IP", # COSINE
# params={"nlist": 128},
)
index_params.add_index(
field_name="bm_25",
index_name="bm_25_index",
index_type="SPARSE_INVERTED_INDEX", # Index type for sparse vectors
metric_type="BM25", # Set to `BM25` when using function to generate sparse vectors
params={"inverted_index_algo": "DAAT_MAXSCORE"},
# The ratio of small vector values to be dropped during indexing
)
self.client.create_collection(
collection_name=collection_name,
schema=schema,
index_params=index_params
)
self.release_collection(collection_name)
return "collections created"
def create_hybrid_collection(self, collection_name, database_name: str, save_title: bool = True):
"""
中文专用混合检索 表
"""
self.client.use_database(database_name)
schema = MilvusClient.create_schema(
auto_id=True,
enable_dynamic_field=True,
)
analyzer_params_custom = {
"tokenizer": "jieba",
"type": "chinese",
"filter": ["cnalphanumonly"]
}
schema.add_field(field_name="id", datatype=DataType.INT64, is_primary=True)
schema.add_field(field_name="text", datatype=DataType.VARCHAR, max_length=65535, enable_analyzer=True,
analyzer_params=analyzer_params_custom)
schema.add_field(field_name="origin_text", datatype=DataType.VARCHAR, max_length=65535, default_value="EMPTY")
schema.add_field(field_name="bm_25", datatype=DataType.SPARSE_FLOAT_VECTOR)
schema.add_field(field_name="vector", datatype=DataType.FLOAT_VECTOR, dim=1024, )
schema.add_field(field_name="resource", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
schema.add_field(field_name="img_path", datatype=DataType.VARCHAR, max_length=4096, default_value="EMPTY")
schema.add_field(field_name="page_idx", datatype=DataType.INT64, default_value=0)
schema.add_field(field_name="origin_id", datatype=DataType.VARCHAR, max_length=256, default_value="EMPTY")
# schema.add_field(field_name="img_information",datatype=DataType.Json,default_value="EMPTY")
if save_title:
schema.add_field(field_name="header_1", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
schema.add_field(field_name="header_2", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
schema.add_field(field_name="header_3", datatype=DataType.VARCHAR, max_length=1024, default_value='EMPTY')
bm25_function = Function(
name="bm25_emb", # Function name
input_field_names=["text"], # Name of the VARCHAR field containing raw text data
output_field_names=["bm_25"],
# Name of the SPARSE_FLOAT_VECTOR field reserved to store generated embeddings
function_type=FunctionType.BM25,
)
schema.add_function(bm25_function)
index_params = self.client.prepare_index_params()
# Add indexes
index_params.add_index(
field_name="vector",
index_name="vector_index",
index_type="HNSW", # "IVF_FLAT". AUTOINDEX
metric_type="IP", # COSINE
# params={"nlist": 128},
)
index_params.add_index(
field_name="bm_25",
index_name="bm_25_index",
index_type="SPARSE_INVERTED_INDEX", # Index type for sparse vectors
metric_type="BM25", # Set to `BM25` when using function to generate sparse vectors
params={"inverted_index_algo": "DAAT_MAXSCORE"},
# The ratio of small vector values to be dropped during indexing
)
self.client.create_collection(
collection_name=collection_name,
schema=schema,
index_params=index_params
)
self.release_collection(collection_name)
return "collections created"
def drop_collection(self, collection_name, database_name: str):
self.client.use_database(database_name)
self.load_collection(collection_name=collection_name)
self.client.drop_collection(collection_name)
def load_collection(self, collection_name):
self.client.load_collection(collection_name=collection_name)
res = self.client.get_load_state(
collection_name=collection_name
)
return res
def release_collection(self, collection_name):
self.client.release_collection(collection_name=collection_name)
res = self.client.get_load_state(
collection_name=collection_name
)
return res
def delete_resource(self, collection_name, resource, database_name):
self.client.use_database(database_name)
self.load_collection(collection_name)
res = self.client.delete(
collection_name=collection_name,
filter="resource == '{}'".format(resource),
)
return res
def has_collection(self, collection_name, database_name):
"""
Before using list collection, You need to using using_database method first
"""
self.using_database(database_name=database_name)
return self.client.has_collection(collection_name)
def list_collections(self, database_name):
"""
Before using list collection, You need to using using_database method first
"""
self.client.use_database(database_name)
return self.client.list_collections()
def hierarchical_insert(self, params, collection_name: str, database_name: str):
try:
self.client.use_database(database_name)
self.load_collection(collection_name)
res = self.client.insert(
collection_name=collection_name,
data=params
)
return res
finally:
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")
def hybrid_insert(self, params, collection_name: str, database_name: str):
try:
self.client.use_database(database_name)
self.load_collection(collection_name)
res = self.client.insert(
collection_name=collection_name,
data=params
)
return res
finally:
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")
# def insert(self, params,collection_name:str = "MarkdownTest",database_name:str = "JXTest"):
# """
# params: List[List]
# 例子: [[id, header_1, header_2, header_3, text, vector]]
# """
# # self.client.insert(collection_name=collection_name, data=params)
# db.using_database(database_name)
# collection = Collection(collection_name)
# mr = collection.insert(params)
# return mr
def get_all_chunk_by_resource(self, collection_name, database_name, resource):
"""
return List :
"""
try:
self.client.using_database(db_name=database_name)
if not self.client.has_collection(collection_name):
return []
else:
self.client.load_collection(collection_name)
iterator = self.client.query_iterator(
collection_name=collection_name,
batch_size=5,
filter=f"resource like \"{resource}\"",
output_fields=["id", "text", "resource", "Header_1", "Header_2", "Header_3"],
)
results = []
while True:
result = iterator.next()
if not result:
iterator.close()
break
results += result
return results
finally:
# 3. 无论成功还是失败,最终释放集合
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")
def get_chunk_by_resource(self, collection_name, database_name, resource):
"""
return List :
"""
try:
self.client.using_database(db_name=database_name)
if not self.client.has_collection(collection_name):
return []
else:
self.client.load_collection(collection_name)
iterator = self.client.query_iterator(
collection_name=collection_name,
batch_size=5,
filter=f"resource like \"{resource}\"",
output_fields=["id", "text", "resource"],
)
results = []
while True:
result = iterator.next()
if not result:
iterator.close()
break
results += result
return results
finally:
# 3. 无论成功还是失败,最终释放集合
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")
def get_chunk_by_id(self, collection_name, database_name, id):
"""
return List :
"""
try:
self.client.using_database(db_name=database_name)
if not self.client.has_collection(collection_name):
return []
else:
self.client.load_collection(collection_name)
iterator = self.client.query_iterator(
collection_name=collection_name,
batch_size=5,
filter=f"id=={id + 1}",
output_fields=["id", "text", "resource"],
)
results = []
while True:
result = iterator.next()
if not result:
iterator.close()
break
results += result
return results
finally:
# 3. 无论成功还是失败,最终释放集合
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")
def get_text_by_resource(self, collection_name, database_name, resource, logger):
"""
return List :
"""
# def prepare_milvus_filter( resource):
# """安全构造 Milvus 过滤条件"""
# escaped_resource = json.dumps(resource)[1:-1]
# return f'resource="{escaped_resource}"'
def prepare_milvus_filter(resource):
"""安全构造 Milvus 过滤条件"""
# 替换资源名称中的双引号
safe_resource = resource.replace('"', '\\"')
return f'resource == "{safe_resource}"'
self.client.using_database(db_name=database_name)
# 检查集合是否存在
if not self.client.has_collection(collection_name):
logger.info(f"集合 {collection_name} 不存在")
return []
else:
self.client.load_collection(collection_name)
try:
# 2. 执行查询
filter_expr = prepare_milvus_filter(resource)
iterator = self.client.query_iterator(
collection_name=collection_name,
batch_size=5,
filter=filter_expr,
output_fields=["id", "resource", "text"],
)
results = []
while True:
try:
result = iterator.next()
if not result:
break
results.extend(result)
except StopIteration:
break
except Exception as e:
print(f"查询迭代出错: {e}")
break
return results
except Exception as e:
print(f"查询失败: {e}")
return []
finally:
# 3. 无论成功还是失败,最终释放集合
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")
def embedding_search(self, query_vector: List, top_k=3, collection_name: str = "MarkdownTest",
database_name: str = "JXTest") -> List[dict | str]:
"""
纯向量检索
"""
self.client.use_database(database_name)
self.load_collection(collection_name)
res = self.client.search(
collection_name=collection_name,
anns_field="vector",
limit=top_k,
data=[query_vector],
search_params={
"metric_type": "IP", # COSINE
"params": {"ef": 64} # "params": {"nprobe": 10}
},
output_fields=["id", "Header_1", "Header_2", "Header_3", "text", "resource", "img_path", "page_idx",
"origin_id", "*"])
search_results = []
search_resources = []
recall_results = []
try:
for i, hits in enumerate(res):
for j, hit in enumerate(hits):
id = hit["entity"].get("id")
search_result_text = hit["entity"].get("text")
distance = hit["distance"]
page_idx = hit["entity"].get("page_idx")
search_result_header_1 = hit["entity"].get("header_1")
search_result_header_2 = hit["entity"].get("header_2")
search_result_header_3 = hit["entity"].get("header_3")
search_result_resource = hit["entity"].get("resource")
search_result_img_path = hit["entity"].get("img_path")
origin_id = hit["entity"].get("origin_id")
dynamic_field_name = hit["entity"].get_properties()
# search_score = hit["entity"].get("score")
# search_result_img_information = hit["entity"].get("img_information")
if search_result_header_1 is None:
search_result_header_2 = ""
else:
search_result_header_1 = f"#{search_result_header_1}"
if search_result_header_2 is None:
search_result_header_2 = ""
else:
search_result_header_2 = f"##{search_result_header_2}"
if search_result_header_3 is None:
search_result_header_3 = ""
else:
search_result_header_3 = f"###{search_result_header_3}"
# search_result = f"#{search_result_header_1}\n{search_result_header_2}\n{search_result_header_3}\n\n{search_result_text}"
search_result = search_result_text
search_results.append(search_result)
search_resources.append({"resource": search_result_resource, "img_path": search_result_img_path})
# search_resources.append({"resource":search_result_resource,"img_path":search_result_img_path,"img_information":search_result_img_imformation})
recall_results.append(
{"id": origin_id, "index": i + j + 1, "text": search_result, "resource": search_result_resource,
"img_path": search_result_img_path, "score": distance, "page_idx": page_idx,
"dynamic_field_name": dynamic_field_name})
# recall_results.append({"index":i+j+1,"text":search_result,"resource":search_result_resource,"img_path":search_result_img_path,"img_information":search_result_img_imformation})
return search_results, search_resources, recall_results
finally:
# 3. 无论成功还是失败,最终释放集合
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")
def hierarchical_search(self, query_vector: List, query, top_k=3, collection_name: str = "MarkdownTest",
database_name: str = "JXTest") -> List[dict | str]:
"""
层级检索
"""
try:
self.client.use_database(database_name)
self.load_collection(collection_name)
search_param_1 = {
"data": [query_vector],
"anns_field": "vector",
"param": {
"metric_type": "IP", # COSINE
"params": {"ef": 64} # "params": {"nprobe": 10}
},
"limit": top_k
}
request_1 = AnnSearchRequest(**search_param_1)
search_param_2 = {
"data": [query],
"anns_field": "bm_25",
"param": {
"metric_type": "BM25",
},
"limit": top_k
}
request_2 = AnnSearchRequest(**search_param_2)
reqs = [request_1, request_2]
ranker = WeightedRanker(0.3, 0.7)
res = self.client.hybrid_search(
collection_name=collection_name,
reqs=reqs,
ranker=ranker,
limit=top_k,
output_fields=["id", "Header_1", "Header_2", "Header_3", "text", "resource", "img_path", "page_idx",
"origin_id"]
# output_fields=["Header_1","Header_2","Header_3","text","resource","img_path","img_information"]
)
search_results = []
search_resources = []
recall_results = []
for i, hits in enumerate(res):
for j, hit in enumerate(hits):
id = hit["entity"].get("id")
search_result_text = hit["entity"].get("text")
distance = hit["distance"]
page_idx = hit["entity"].get("page_idx")
text = hit["entity"].get("text")
search_result_header_1 = hit["entity"].get("header_1")
search_result_header_2 = hit["entity"].get("header_2")
search_result_header_3 = hit["entity"].get("header_3")
search_result_resource = hit["entity"].get("resource")
search_result_img_path = hit["entity"].get("img_path")
origin_id = hit["entity"].get("origin_id")
origin_text = hit["entity".get("origin_text")]
# search_result_img_information = hit["entity"].get("img_information")
if search_result_header_1 is None:
search_result_header_2 = ""
else:
search_result_header_1 = f"#{search_result_header_1}"
if search_result_header_2 is None:
search_result_header_2 = ""
else:
search_result_header_2 = f"##{search_result_header_2}"
if search_result_header_3 is None:
search_result_header_3 = ""
else:
search_result_header_3 = f"###{search_result_header_3}"
# search_result = f"#{search_result_header_1}\n{search_result_header_2}\n{search_result_header_3}\n\n{search_result_text}"
search_result = search_result_text
search_results.append(search_result)
search_resources.append({"resource": search_result_resource, "img_path": search_result_img_path})
# search_resources.append({"resource":search_result_resource,"img_path":search_result_img_path,"img_information":search_result_img_imformation})
# recall_results.append({"index":i+j+1,"id":origin_id,"text":origin_text,"resource":search_result_resource,"img_path":search_result_img_path,"score":distance,"page_idx":page_idx})
recall_results.append(
{"index": i + j + 1, "id": origin_id, "text": text, "resource": search_result_resource,
"img_path": search_result_img_path, "score": distance, "page_idx": page_idx})
self.release_collection(collection_name)
return search_results, search_resources, recall_results
finally:
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")
def hybrid_search(self, query_vector: List, query, top_k=3, collection_name: str = "MarkdownTest",
database_name: str = "JXTest", embedding_rate=0.3) -> List[dict | str]:
"""
"""
try:
self.client.use_database(database_name)
self.load_collection(collection_name)
search_param_1 = {
"data": [query_vector],
"anns_field": "vector",
"param": {
"metric_type": "IP", # COSINE
"params": {"ef": 64} # "params": {"nprobe": 10}
},
"limit": top_k
}
request_1 = AnnSearchRequest(**search_param_1)
search_param_2 = {
"data": [query],
"anns_field": "bm_25",
"param": {
"metric_type": "BM25",
},
"limit": top_k
}
request_2 = AnnSearchRequest(**search_param_2)
reqs = [request_1, request_2]
ranker = WeightedRanker(float(embedding_rate), 1 - float(embedding_rate))
res = self.client.hybrid_search(
collection_name=collection_name,
reqs=reqs,
ranker=ranker,
limit=top_k,
output_fields=["id", "Header_1", "Header_2", "Header_3", "text", "resource", "img_path", "page_idx",
"origin_id", "origin_text"]
# output_fields=["Header_1","Header_2","Header_3","text","resource","img_path","img_information"]
)
search_results = []
search_resources = []
recall_results = []
for i, hits in enumerate(res):
for j, hit in enumerate(hits):
id = hit["entity"].get("id")
search_result_text = hit["entity"].get("text")
distance = hit["distance"]
page_idx = hit["entity"].get("page_idx")
search_result_header_1 = hit["entity"].get("Header_1")
search_result_header_2 = hit["entity"].get("Header_2")
search_result_header_3 = hit["entity"].get("Header_3")
search_result_resource = hit["entity"].get("resource")
search_result_img_path = hit["entity"].get("img_path")
origin_id = hit["entity"].get("origin_id")
header_1_num = len(search_result_header_1) if search_result_header_1 else 0
header_2_num = len(search_result_header_2) if search_result_header_2 else 0
header_3_num = len(search_result_header_3) if search_result_header_3 else 0
origin_text = hit["entity"].get("origin_text")
# search_result_img_information = hit["entity"].get("img_information")
if search_result_header_1 is None:
search_result_header_2 = ""
else:
search_result_header_1 = f"#{search_result_header_1}"
if search_result_header_2 is None:
search_result_header_2 = ""
else:
search_result_header_2 = f"##{search_result_header_2}"
if search_result_header_3 is None:
search_result_header_3 = ""
else:
search_result_header_3 = f"###{search_result_header_3}"
# search_result = f"#{search_result_header_1}\n{search_result_header_2}\n{search_result_header_3}\n\n{search_result_text}"
search_result = search_result_text
origin_text = search_result[header_1_num + header_2_num + header_3_num + 3:]
search_results.append(search_result)
search_resources.append({"resource": search_result_resource, "img_path": search_result_img_path})
# 获取片段的下文,并且拼接
# next_chunk = self.get_chunk_by_id(collection_name=collection_name,database_name=database_name,id =id)
# next_text = next_chunk[0].get("text")
# search_result += next_text
# origin_text += next_text
# search_resources.append({"resource":search_result_resource,"img_path":search_result_img_path,"img_information":search_result_img_imformation})
recall_results.append(
{"index": i + j + 1, "id": origin_id, "text": search_result, "resource": search_result_resource,
"img_path": search_result_img_path, "score": distance, "page_idx": page_idx,
"origin_text": origin_text})
# recall_results.append({"index":i+j+1,"text":search_result,"resource":search_result_resource,"img_path":search_result_img_path,"img_information":search_result_img_imformation})
self.release_collection(collection_name)
return search_results, search_resources, recall_results
finally:
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")
# return search_results
def hybrid_search_batch(
self,
query_vectors: List[List[float]],
queries: List[str],
top_k: int = 3,
collection_name: str = "MarkdownTest",
database_name: str = "JXTest",
embedding_rate: float = 0.3
) -> Tuple[List[List[Dict]], List[List[Dict]], List[List[Dict]]]:
"""
执行批量混合检索(稠密向量 + BM25 稀疏检索),对每个查询分别进行 hybrid search。
Args:
query_vectors (List[List[float]]): 每个查询对应的稠密向量列表,形状为 [N, dim]
queries (List[str]): 原始文本查询列表,长度为 N
top_k (int): 返回每个查询的 top-k 结果
collection_name (str): 集合名称
database_name (str): 数据库名称
embedding_rate (float): 稠密向量检索结果的权重BM25 权重 = 1 - embedding_rate
Returns:
Tuple[
List[List[str]]: 每个查询对应的结果文本列表(外层 list 是 query 数量,内层是 top-k 结果),
List[List[dict]]: 每个查询对应的资源信息列表 ,
List[List[dict]]: 每个查询对应的完整召回详情列表(含 score、id、origin_text 等)
]
"""
if len(query_vectors) != len(queries):
raise ValueError("query_vectors 和 queries 的数量必须一致")
try:
self.client.use_database(database_name)
self.load_collection(collection_name)
all_search_results = [] # [[res1_q1, res2_q1], [res1_q2, ...]]
all_search_resources = []
all_recall_results = []
for i, (vec, text) in enumerate(zip(query_vectors, queries)):
try:
# 构建稠密向量请求
req1 = AnnSearchRequest(
data=[vec],
anns_field="vector",
param={"metric_type": "IP", "params": {"ef": 64}},
limit=top_k
)
# 构建稀疏向量请求BM25
req2 = AnnSearchRequest(
data=[text],
anns_field="bm_25",
param={"metric_type": "BM25"},
limit=top_k
)
# 混合搜索
ranker = WeightedRanker(embedding_rate, 1 - embedding_rate)
res = self.client.hybrid_search(
collection_name=collection_name,
reqs=[req1, req2],
ranker=ranker,
limit=top_k,
output_fields=[
"id", "Header_1", "Header_2", "Header_3", "text", "resource",
"img_path", "page_idx", "origin_id", "origin_text"
],
timeout=30
)
# 解析单个查询的结果
search_results = []
search_resources = []
recall_results = []
for j, hit in enumerate(res[0]): # res 是 List[List[Hit]], res[0] 对应第一个 request 的 hits
entity = hit["entity"]
id_val = entity.get("id")
origin_id = entity.get("origin_id")
text_val = entity.get("text") or ""
header_1 = entity.get("Header_1") or ""
header_2 = entity.get("Header_2") or ""
header_3 = entity.get("Header_3") or ""
resource = entity.get("resource")
img_path = entity.get("img_path")
page_idx = entity.get("page_idx")
distance = hit["distance"]
origin_text = entity.get("origin_text") or text_val
# 格式化标题
header_1 = f"#{header_1}" if header_1 else ""
header_2 = f"##{header_2}" if header_2 else ""
header_3 = f"###{header_3}" if header_3 else ""
# 组合内容(可根据需要调整)
combined_text = "\n".join(filter(None, [header_1, header_2, header_3, text_val])).strip()
search_results.append(combined_text)
search_resources.append({"resource": resource, "img_path": img_path})
recall_results.append({
"keyword": text,
"index": j + 1,
"id": origin_id,
"text": combined_text,
"resource": resource,
"img_path": img_path,
"score": distance,
"page_idx": page_idx,
"origin_text": origin_text
})
all_search_results.append(search_results)
all_search_resources.append(search_resources)
all_recall_results.append(recall_results)
except Exception as e:
print(f"{i + 1} 个查询 '{text}' 检索失败: {e}")
# 失败时返回空列表占位
all_search_results.append([])
all_search_resources.append([])
all_recall_results.append([])
return all_search_results, all_search_resources, all_recall_results
except Exception as e:
print(f"批量混合检索整体失败: {e}")
raise
finally:
try:
self.client.release_collection(collection_name)
except Exception as e:
print(f"释放集合失败: {e}")
# return search_results
def close(self):
self.client.close()
def sanitize_collection_name(name: str) -> str:
"""
中文转换为拼音
"""
def is_valid_milvus_collection_name(name: str) -> bool:
"""检查名称是否符合 Milvus 集合命名规则"""
import re
pattern = r'^[a-zA-Z_][a-zA-Z0-9_]*$' # 首字符必须是字母或_后续可以是字母、数字、_
return bool(re.fullmatch(pattern, name))
if is_valid_milvus_collection_name(name):
return name
pinyin = "_".join(lazy_pinyin(name))
# 确保首字符是字母或下划线
if not pinyin[0].isalpha() and pinyin[0] != "_":
pinyin = f"_{pinyin}"
return pinyin
##应用示例
# if __name__ == '__main__':
# #创建collection
# database = Milvus_Database("XIAN")
# result = database.create_connection("MarkdownTest",drop_=False)
# print(result)
# print(database.list_partitions("MarkdownTest"))
# 创建db
# from pymilvus import connections, db
# #
# conn = Milvus_Database(user='root',
# password='Milvus',
# uri="http://192.168.0.46:19530")
# print(conn.list_collections("XIAN"))
# conn = connections.connect(user='root',
# password='Milvus',
# uri="http://172.18.30.165:19530")
# database = db.create_database("JXTest")
# db.using_database("JXTest")