421 lines
16 KiB
Python
421 lines
16 KiB
Python
from datetime import datetime
|
||
from typing import Optional, List, Dict, Any
|
||
|
||
from pymilvus import MilvusClient, DataType
|
||
|
||
from app.client.milvus_client_manager import milvus_client
|
||
from app.conf.app_config import app_config
|
||
from app.core.log import logger
|
||
|
||
|
||
class IntentClassificationRepository:
|
||
collection_name: str = 'intent_classification'
|
||
|
||
def __init__(self):
|
||
self.dim = app_config.milvus.embedding_size
|
||
self._ensure_connection()
|
||
|
||
def _ensure_connection(self):
|
||
if milvus_client.client is None:
|
||
milvus_client.init()
|
||
|
||
def _ensure_loaded(self):
|
||
"""确保集合已加载到内存,否则 Milvus query/search 可能返回空结果"""
|
||
self._ensure_connection()
|
||
try:
|
||
milvus_client.client.load_collection(self.collection_name)
|
||
except Exception:
|
||
# Milvus 某些版本 load_collection 重复调用会报错,直接忽略
|
||
pass
|
||
|
||
def create_collection(self, dim: int = None) -> bool:
|
||
if dim is not None:
|
||
self.dim = dim
|
||
|
||
if milvus_client.has_collection(self.collection_name):
|
||
return False
|
||
|
||
schema = milvus_client.client.create_schema(
|
||
auto_id=True,
|
||
enable_dynamic_field=True,
|
||
description="意图分类向量集合"
|
||
)
|
||
|
||
schema.add_field(field_name="id", datatype=DataType.INT64, is_primary=True, auto_id=True)
|
||
schema.add_field(field_name="intent_category", datatype=DataType.VARCHAR, max_length=128)
|
||
schema.add_field(field_name="content", datatype=DataType.VARCHAR, max_length=65535)
|
||
schema.add_field(field_name="review_content", datatype=DataType.VARCHAR, max_length=65535)
|
||
schema.add_field(field_name="content_vector", datatype=DataType.FLOAT_VECTOR, dim=1024)
|
||
schema.add_field(field_name="created_at", datatype=DataType.INT64)
|
||
|
||
index_params = milvus_client.client.prepare_index_params()
|
||
index_params.add_index(
|
||
field_name="content_vector",
|
||
index_type="HNSW",
|
||
metric_type="COSINE",
|
||
params={"M": 32, "efConstruction": 300}
|
||
)
|
||
milvus_client.client.create_collection(
|
||
collection_name=self.collection_name,
|
||
schema=schema,
|
||
index_params=index_params
|
||
)
|
||
|
||
return True
|
||
# 迁移
|
||
def migrate_add_review_content(self) -> bool:
|
||
self._ensure_loaded()
|
||
if not milvus_client.has_collection(self.collection_name):
|
||
logger.error(f"集合 {self.collection_name} 不存在,无法迁移")
|
||
return False
|
||
|
||
temp_collection = f"{self.collection_name}_temp"
|
||
|
||
if milvus_client.has_collection(temp_collection):
|
||
milvus_client.client.drop_collection(temp_collection)
|
||
|
||
schema = milvus_client.client.create_schema(
|
||
auto_id=True,
|
||
enable_dynamic_field=True,
|
||
description="意图分类向量集合(迁移:新增review_content)"
|
||
)
|
||
schema.add_field(field_name="id", datatype=DataType.INT64, is_primary=True, auto_id=True)
|
||
schema.add_field(field_name="intent_category", datatype=DataType.VARCHAR, max_length=128)
|
||
schema.add_field(field_name="content", datatype=DataType.VARCHAR, max_length=65535)
|
||
schema.add_field(field_name="review_content", datatype=DataType.VARCHAR, max_length=65535)
|
||
schema.add_field(field_name="content_vector", datatype=DataType.FLOAT_VECTOR, dim=1024)
|
||
schema.add_field(field_name="created_at", datatype=DataType.INT64)
|
||
|
||
index_params = milvus_client.client.prepare_index_params()
|
||
index_params.add_index(
|
||
field_name="content_vector",
|
||
index_type="HNSW",
|
||
metric_type="COSINE",
|
||
params={"M": 32, "efConstruction": 300}
|
||
)
|
||
milvus_client.client.create_collection(
|
||
collection_name=temp_collection,
|
||
schema=schema,
|
||
index_params=index_params
|
||
)
|
||
logger.info(f"创建临时集合 {temp_collection} 成功")
|
||
|
||
all_data = []
|
||
offset = 0
|
||
batch_size = 1000
|
||
while True:
|
||
batch = milvus_client.client.query(
|
||
collection_name=self.collection_name,
|
||
output_fields=["intent_category", "content", "content_vector", "created_at"],
|
||
limit=batch_size,
|
||
offset=offset
|
||
)
|
||
if not batch:
|
||
break
|
||
all_data.extend(batch)
|
||
if len(batch) < batch_size:
|
||
break
|
||
offset += batch_size
|
||
|
||
logger.info(f"查询到旧数据 {len(all_data)} 条,开始迁移")
|
||
|
||
migrated = 0
|
||
for i in range(0, len(all_data), batch_size):
|
||
batch = all_data[i:i + batch_size]
|
||
insert_batch = []
|
||
for item in batch:
|
||
insert_batch.append({
|
||
"intent_category": item.get("intent_category", ""),
|
||
"content": item.get("content", ""),
|
||
"review_content": "",
|
||
"content_vector": item.get("content_vector", []),
|
||
"created_at": item.get("created_at", 0)
|
||
})
|
||
milvus_client.client.insert(
|
||
collection_name=temp_collection,
|
||
data=insert_batch
|
||
)
|
||
migrated += len(insert_batch)
|
||
logger.info(f"迁移进度: {migrated}/{len(all_data)}")
|
||
|
||
logger.info(f"数据迁移完成,共 {len(all_data)} 条")
|
||
|
||
milvus_client.client.drop_collection(self.collection_name)
|
||
logger.info(f"旧集合 {self.collection_name} 已删除")
|
||
|
||
milvus_client.client.rename_collection(temp_collection, self.collection_name)
|
||
logger.info(f"临时集合重命名为 {self.collection_name},迁移完成")
|
||
|
||
return True
|
||
|
||
async def save_intent_data(self, data: List[Dict[str, Any]]) -> List[int]:
|
||
try:
|
||
created_at = int(datetime.now().timestamp())
|
||
insert_data = []
|
||
for item in data:
|
||
insert_data.append({
|
||
"intent_category": item.get("intent_category", ""),
|
||
"content": item.get("content", ""),
|
||
"review_content": item.get("review_content", ""),
|
||
"content_vector": item.get("content_vector", []),
|
||
"created_at": created_at
|
||
})
|
||
|
||
result = milvus_client.client.insert(
|
||
collection_name=self.collection_name,
|
||
data=insert_data
|
||
)
|
||
|
||
logger.info(f"意图分类数据入库成功,插入 {len(data)} 条记录")
|
||
return result.get("ids", [])
|
||
|
||
except Exception as e:
|
||
logger.error(f"意图分类数据入库失败: {str(e)}", exc_info=True)
|
||
return []
|
||
|
||
async def upsert_intent_data(self, data: List[Dict[str, Any]]) -> List[int]:
|
||
try:
|
||
self._ensure_loaded()
|
||
created_at = int(datetime.now().timestamp())
|
||
all_ids = []
|
||
for item in data:
|
||
category = item.get("intent_category", "")
|
||
content = item.get("content", "")
|
||
vector = item.get("content_vector", [])
|
||
|
||
escaped_content = content.replace("\\", "\\\\").replace('"', '\\"')
|
||
expr = f'intent_category == "{category}"'
|
||
existing = milvus_client.client.query(
|
||
collection_name=self.collection_name,
|
||
filter=expr,
|
||
output_fields=["id"],
|
||
limit=1000
|
||
)
|
||
if existing:
|
||
existing_ids = [e["id"] for e in existing]
|
||
milvus_client.client.delete(
|
||
collection_name=self.collection_name,
|
||
ids=existing_ids
|
||
)
|
||
logger.info(f"删除旧记录 {len(existing_ids)} 条: category={category}")
|
||
|
||
insert_result = milvus_client.client.insert(
|
||
collection_name=self.collection_name,
|
||
data=[{
|
||
"intent_category": category,
|
||
"content": content,
|
||
"review_content": item.get("review_content", ""),
|
||
"content_vector": vector,
|
||
"created_at": created_at
|
||
}]
|
||
)
|
||
all_ids.extend(insert_result.get("ids", []))
|
||
|
||
logger.info(f"意图分类数据upsert完成,共处理 {len(data)} 条")
|
||
return all_ids
|
||
|
||
except Exception as e:
|
||
logger.error(f"意图分类数据upsert失败: {str(e)}", exc_info=True)
|
||
return []
|
||
|
||
async def search_by_intent_category(self, intent_category: List[str], top_k: int = 100) -> List[Dict[str, Any]]:
|
||
self._ensure_loaded()
|
||
if isinstance(intent_category, str):
|
||
intent_category = [intent_category]
|
||
categories_str = ", ".join(f'"{c}"' for c in intent_category)
|
||
expr = f'intent_category in [{categories_str}]'
|
||
result = milvus_client.client.query(
|
||
collection_name=self.collection_name,
|
||
filter=expr,
|
||
output_fields=["intent_category", "content", "review_content", "created_at"],
|
||
limit=top_k
|
||
)
|
||
return result
|
||
|
||
async def search_by_vector(self, query_vector: List[float], top_k: int = 5) -> List[Dict[str, Any]]:
|
||
self._ensure_loaded()
|
||
res = milvus_client.client.search(
|
||
collection_name=self.collection_name,
|
||
anns_field="content_vector",
|
||
data=[query_vector],
|
||
limit=top_k,
|
||
search_params={
|
||
"metric_type": "COSINE",
|
||
"efSearch": 300
|
||
},
|
||
output_fields=["intent_category", "content", "review_content", "created_at"]
|
||
)
|
||
|
||
results = []
|
||
for hits in res:
|
||
for hit in hits:
|
||
item = {
|
||
"intent_category": hit.entity.get("intent_category"),
|
||
"content": hit.entity.get("content"),
|
||
"review_content": hit.entity.get("review_content"),
|
||
"created_at": hit.entity.get("created_at"),
|
||
"score": hit.get("distance", 0)
|
||
}
|
||
results.append(item)
|
||
return results
|
||
|
||
async def get_all_intent_categories(self) -> List[str]:
|
||
self._ensure_loaded()
|
||
result = milvus_client.client.query(
|
||
collection_name=self.collection_name,
|
||
output_fields=["intent_category"],
|
||
limit=1000
|
||
)
|
||
categories = list(set([item.get("intent_category", "") for item in result]))
|
||
return categories
|
||
|
||
async def get_style_intent_categories(self) -> List[str]:
|
||
all_categories = await self.get_all_intent_categories()
|
||
return [c for c in all_categories]
|
||
|
||
async def get_content_by_intent_category(self, intent_category: str) -> List[Dict[str, str]]:
|
||
self._ensure_loaded()
|
||
result = milvus_client.client.query(
|
||
collection_name=self.collection_name,
|
||
filter=f'intent_category == "{intent_category}"',
|
||
output_fields=["content", "review_content"],
|
||
limit=1000
|
||
)
|
||
return [
|
||
{
|
||
"content": item.get("content", ""),
|
||
"review_content": item.get("review_content", "")
|
||
}
|
||
# content 或 review_content 任一有非空值就保留,避免全空但 review_content 有值的被误过滤
|
||
for item in result
|
||
if (item.get("content") or "").strip() or (item.get("review_content") or "").strip()
|
||
]
|
||
|
||
async def delete_today_data(self) -> int:
|
||
self._ensure_loaded()
|
||
today_start = int(datetime.now().replace(hour=0, minute=0, second=0, microsecond=0).timestamp())
|
||
query_result = milvus_client.client.query(
|
||
collection_name=self.collection_name,
|
||
filter=f'created_at >= {today_start}',
|
||
output_fields=["id"],
|
||
limit=10000
|
||
)
|
||
if not query_result:
|
||
logger.info("今日无数据可删除")
|
||
return 0
|
||
ids_to_delete = [item["id"] for item in query_result]
|
||
milvus_client.client.delete(
|
||
collection_name=self.collection_name,
|
||
ids=ids_to_delete
|
||
)
|
||
logger.info(f"删除今日数据 {len(ids_to_delete)} 条")
|
||
return len(ids_to_delete)
|
||
|
||
def sync_content_to_review_content(self) -> int:
|
||
self._ensure_loaded()
|
||
result = milvus_client.client.query(
|
||
collection_name=self.collection_name,
|
||
output_fields=["id", "intent_category", "content", "review_content", "content_vector", "created_at"],
|
||
limit=10000
|
||
)
|
||
if not result:
|
||
logger.info("无数据需要同步")
|
||
return 0
|
||
|
||
updated = 0
|
||
batch_size = 1000
|
||
for i in range(0, len(result), batch_size):
|
||
batch = result[i:i + batch_size]
|
||
ids_to_delete = []
|
||
insert_batch = []
|
||
for item in batch:
|
||
ids_to_delete.append(item["id"])
|
||
insert_batch.append({
|
||
"intent_category": item.get("intent_category", ""),
|
||
"content": item.get("content", ""),
|
||
"review_content": item.get("content", ""),
|
||
"content_vector": item.get("content_vector", []),
|
||
"created_at": item.get("created_at", 0)
|
||
})
|
||
milvus_client.client.delete(
|
||
collection_name=self.collection_name,
|
||
ids=ids_to_delete
|
||
)
|
||
milvus_client.client.insert(
|
||
collection_name=self.collection_name,
|
||
data=insert_batch
|
||
)
|
||
updated += len(insert_batch)
|
||
logger.info(f"同步进度: {updated}/{len(result)}")
|
||
|
||
logger.info(f"content → review_content 同步完成,共更新 {updated} 条")
|
||
return updated
|
||
|
||
|
||
async def main():
|
||
milvus_client.init()
|
||
repo = IntentClassificationRepository()
|
||
|
||
# ===== 诊断输出 =====
|
||
client = milvus_client.client
|
||
collection_name = repo.collection_name
|
||
|
||
# 1. 集合是否存在
|
||
has = client.has_collection(collection_name)
|
||
print(f"[1] 集合 {collection_name} 是否存在: {has}")
|
||
if not has:
|
||
print(" 集合不存在,跳过后续诊断")
|
||
return
|
||
|
||
# 2. 集合 describe
|
||
desc = client.describe_collection(collection_name)
|
||
print(f"[2] 集合信息: {desc}")
|
||
|
||
# 3. Milvus 关键:加载集合到内存(否则 query/search 可能返回空)
|
||
try:
|
||
client.load_collection(collection_name)
|
||
print("[3] load_collection 成功")
|
||
except Exception as e:
|
||
print(f"[3] load_collection 提示: {e}")
|
||
|
||
# 4. 不指定任何 filter,直接 count / query 看总条数
|
||
stats = client.get_collection_stats(collection_name)
|
||
print(f"[4] 集合统计信息: {stats}")
|
||
|
||
# 5. 无条件 query 前 20 条 intent_category
|
||
all_sample = client.query(
|
||
collection_name=collection_name,
|
||
output_fields=["id", "intent_category"],
|
||
limit=20
|
||
)
|
||
print(f"[5] 前20条 intent_category 样本(共{len(all_sample)}条):")
|
||
for it in all_sample:
|
||
print(f" id={it.get('id')} intent_category={repr(it.get('intent_category'))}")
|
||
|
||
# 6. 用优惠作为条件查询
|
||
expr = 'intent_category == "优惠"'
|
||
print(f"[6] 执行 filter: {expr}")
|
||
youhui = client.query(
|
||
collection_name=collection_name,
|
||
filter=expr,
|
||
output_fields=["id", "intent_category", "content", "review_content"],
|
||
limit=100
|
||
)
|
||
print(f" 命中 {len(youhui)} 条")
|
||
for it in youhui:
|
||
c = it.get("content", "")
|
||
rc = it.get("review_content", "")
|
||
print(f" id={it.get('id')} intent_category={repr(it.get('intent_category'))}")
|
||
print(f" content[:100]={repr(c[:100] if c else c)} len={len(c)}")
|
||
print(f" review_content[:100]={repr(rc[:100] if rc else rc)} len={len(rc)}")
|
||
|
||
# 7. 原有方法调用
|
||
style_categories = await repo.get_content_by_intent_category("优惠")
|
||
print(f"[7] repo.get_content_by_intent_category('优惠') 结果: {style_categories}")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
import asyncio
|
||
asyncio.run(main())
|
||
|