sales-assistant-py-new/app/repository/milvus/intent_classification_repository.py

355 lines
14 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.

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 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:
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:
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]]:
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]]:
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]:
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]]:
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", "")
}
for item in result if item.get("content")
]
async def delete_today_data(self) -> int:
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:
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()
# 1. 首次部署:迁移新增 review_content 字段
# success = repo.migrate_add_review_content()
# print(f"迁移结果: {success}")
# 2. 将 content 同步写入 review_content
# count = repo.sync_content_to_review_content()
# print(f"同步完成,共 {count} 条")
style_categories = await repo.get_style_intent_categories()
print(style_categories)
if __name__ == "__main__":
import asyncio
asyncio.run(main())