456 lines
21 KiB
Python
456 lines
21 KiB
Python
import os
|
||
|
||
import asyncio
|
||
import json
|
||
import time
|
||
from datetime import datetime, timedelta
|
||
from typing import List, Dict
|
||
|
||
from langchain_core.prompts import PromptTemplate
|
||
|
||
from app.chat_qa.classification.text_intent import TextIntent
|
||
from app.chat_qa.classification.general_text_extract import TextExtract
|
||
from app.chat_qa.classification.standard_intent import standard_intent
|
||
from app.client.embedding_client_manager import EmbeddingClientManager, embedding_client
|
||
from app.client.milvus_client_manager import milvus_client
|
||
from app.client.mysql_client_manager import db_assistant_mysql_client_manager
|
||
from app.dto.summary_dto import SummaryDto
|
||
from app.llm import llm
|
||
from app.models.mysql import ArchiveMessages
|
||
from app.repository.archive_messages_repository import ArchiveMessagesRepository
|
||
from app.repository.milvus.intent_classification_repository import IntentClassificationRepository
|
||
from app.repository.milvus.summary_repository import SummaryRepository
|
||
from app.core.log import logger
|
||
from app.reranker_utils import get_reranker_model
|
||
import re
|
||
|
||
|
||
intent_prompt = """
|
||
你是一个专业的QA提取专家。请分析以下聊天记录,提取或生成1-2个高质量的QA对。
|
||
|
||
## 核心规则
|
||
|
||
### 1. QA提取规则
|
||
- 优先提取用户明确提出的问题及其对应的回答
|
||
- 如果对话中没有明确的问题,根据讨论内容生成一个总结性问题
|
||
- 答案必须基于聊天记录中的实际内容,不能编造
|
||
- 如果某个问题在聊天记录中没有明确答案,则不生成QA
|
||
- 最多生成2个QA对,优先选择最有价值的信息
|
||
|
||
### 2. 单意图选择规则(关键)
|
||
**每个QA对只能有一个意图标签,不能以顿号、逗号或“和”分隔多个意图。**
|
||
|
||
当内容可能对应多个意图时,按以下优先级选择唯一意图:
|
||
|
||
| 优先级 | 选择规则 | 说明 |
|
||
|-------|---------|------|
|
||
| **P0** | **优先匹配已有意图标签** | 必须优先从已有意图体系(含扩展示例)中选择最匹配的标签,禁止新建与原标签语义高度重复的意图。**当销售单方面介绍课程内容(无用户明确提问)时,直接使用对应的主题标签(如“3D场景建模”、“3D角色建模”),严禁新建“XXX课程咨询”类标签。** 例:“场景线上课程介绍” → 使用“3D场景建模”,而非“场景线上课程咨询”。 |
|
||
| **P1** | **核心诉求优先** | 选择最能代表用户核心诉求的意图。例如:报名流程说明 → “报名咨询”(核心是报名咨询),而非“线上直播教学”(只是提到线上) |
|
||
| **P2** | **动作维度优先于主题维度** | 当主题和诉求冲突时,优先选择主题维度。例如:3D建模费用 → “3D建模”(主题) |
|
||
| **P3** | **直接提及优先于间接关联** | 用户明确说出的意图优先于推断出的意图 |
|
||
| **P4** | **去重合并** | 多个子诉求属于同一类别时,只输出一个意图标签 |
|
||
|
||
### 3. 意图分类体系(必须从中选择或基于其扩充)
|
||
{intent_list}
|
||
|
||
**扩充规则:** 仅当用户问题涉及的具体方向在已有标签中完全无法覆盖时,才可创建更细粒度的标签。创建时必须严格对比已有标签,确保无重复语义。例如:“原画好就业吗” → 可扩充为“原画就业”,因为已有标签中无“就业”与“原画”的结合且不冲突。
|
||
|
||
### 4. 多意图拆分规则(新增)
|
||
- 若同一段对话明确涉及多个不重叠的核心意图(如同时介绍“3D场景建模”和“就业咨询”),必须拆分为多个QA对,每个QA对聚焦一个意图。
|
||
- 每个QA对的answer仅包含与该意图直接相关的信息,不相关的信息不得混杂其中,避免信息污染。
|
||
- 若几个意图高度关联无法拆分(如介绍课程时自然提到费用,但核心是课程),则按单意图规则选择最核心的标签,其余信息可融入answer作为补充,但intents仍为单一核心意图。
|
||
|
||
## 输出格式要求
|
||
- 输出纯JSON,不要任何解释、不要markdown标记、不要注释
|
||
- 将中文引号 "" '' 替换为英文引号 "" ''
|
||
- 删除所有注释(// 和 /* */)
|
||
- 删除markdown代码块标记(```json ```)
|
||
- 删除JSON前后多余的文本说明
|
||
- 修复末尾多余的逗号(如 }},] 或 }},}})
|
||
- 确保所有字符串用双引号包裹
|
||
- 确保布尔值是小写的 true/false,null 是小写的
|
||
- 如果键名没有引号,给键名加上双引号
|
||
|
||
## 输出格式
|
||
{{
|
||
"qa_pairs": [
|
||
{{
|
||
"intents": "意图标签(只能有一个,不能并列)",
|
||
"answer": "回答内容"
|
||
}},
|
||
{{
|
||
"intents": "意图标签(只能有一个,不能并列)",
|
||
"answer": "回答内容"
|
||
}}
|
||
]
|
||
}}
|
||
|
||
|
||
# 历史对话
|
||
{message_str}
|
||
其中INTERNAL表示销售,EXTERNAL表示用户
|
||
|
||
要求:
|
||
- 一个intents对应的answer必须包含聊天记录中所有可用信息,确保信息不遗漏
|
||
- 意图标签必须从上述分类中选择或基于其合理扩充
|
||
- **intents字段必须是单个字符串,不能包含顿号、逗号或“和”**
|
||
- 优先选择最有价值的1-2个QA对
|
||
- 当销售(INTERNAL)主动进行课程/产品介绍,而非回答用户(EXTERNAL)明确提问时,意图标签应选择主题维度(如:3D角色建模),而非诉求维度(如:课程咨询)
|
||
"""
|
||
from FlagEmbedding import FlagReranker
|
||
|
||
class HistoryIntent:
|
||
|
||
def __init__(self, archive_messages_repository: ArchiveMessagesRepository,
|
||
embedding_client: EmbeddingClientManager,
|
||
summary_repository: SummaryRepository,
|
||
reranker: FlagReranker
|
||
):
|
||
self.archive_messages_repository = archive_messages_repository
|
||
self.embedding_client = embedding_client
|
||
self.summary_repository = summary_repository
|
||
self.reranker = reranker
|
||
|
||
async def get_day_ms_timestamp(self, date_str: str):
|
||
"""
|
||
输入日期字符串 如 "2025-06-10"
|
||
:return: start_ms(当日0点毫秒时间戳), end_ms(当日23:59:59毫秒时间戳)
|
||
"""
|
||
# 解析日期,生成当日 00:00:00
|
||
day_start = datetime.strptime(date_str, "%Y-%m-%d")
|
||
# 当日最后一秒 23:59:59
|
||
day_end = day_start + timedelta(days=1) - timedelta(seconds=1)
|
||
|
||
# timestamp() 返回秒,*1000 转毫秒,转整数
|
||
start_ms = int(day_start.timestamp() * 1000)
|
||
end_ms = int(day_end.timestamp() * 1000)
|
||
return start_ms, end_ms
|
||
|
||
async def get_message_summary(self, date: str = None):
|
||
stime = time.time()
|
||
intent_repo = IntentClassificationRepository()
|
||
if date:
|
||
start_ts, end_ts = await self.get_day_ms_timestamp(date)
|
||
day_ArchiveMessages = await self.archive_messages_repository.get_message_statistics_by_date(start_date=start_ts,
|
||
end_date=end_ts)
|
||
logger.info(f"{date}共产生了{len(day_ArchiveMessages)}条会话")
|
||
else:
|
||
day_ArchiveMessages = await self.archive_messages_repository.get_message_statistics_by_date()
|
||
|
||
|
||
# 获取意图分类
|
||
intent_list:list[str] = await intent_repo.get_style_intent_categories()
|
||
# 改成map
|
||
date_user_message_a_set = set()
|
||
summary_list: list[str] = []
|
||
for idx, day_message in enumerate(day_ArchiveMessages, 1): #打印计数
|
||
logger.info(f"处理第{idx}会话")
|
||
day_date = str(day_message.created_at)
|
||
user_a = day_message.from_user
|
||
# 找出发送人相关的接收人信息
|
||
date_user_message_a = await self.archive_messages_repository.get_message_statistics_by_date_user(day_date,
|
||
user_a)
|
||
|
||
# 找出发送人与接收人 相关的所有会话信息
|
||
for date_a in date_user_message_a:
|
||
|
||
date_user_b = date_a.to_user
|
||
# if date_user_b != "[\"wmI1AkDQAAdVhy2ENtY6fecaw9-pi4uQ\"]":
|
||
# continue
|
||
|
||
date_user_a = date_a.from_user
|
||
date_user_b_clean = date_user_b.replace("[\"", "").replace("\"]", "")
|
||
date_user_a_append = "[\"" + date_user_a + "\"]"
|
||
# 判断只要存在就跳过
|
||
if date_user_a + date_user_b in date_user_message_a_set or date_user_b_clean + date_user_a_append in date_user_message_a_set:
|
||
logger.info(f"已存在{date_user_a}和{date_user_b}在{day_date}对话,跳过")
|
||
continue
|
||
logger.info(f"查询{date_user_a}和{date_user_b}在{day_date}对话开始")
|
||
list1 = await self.archive_messages_repository.get_message_statistics_by_date_user(day_date,
|
||
date_user_a,
|
||
date_user_b)
|
||
|
||
list2 = await self.archive_messages_repository.get_message_statistics_by_date_user(day_date,
|
||
date_user_b_clean,
|
||
date_user_a_append)
|
||
# 查询过的聊天账号,后续不再查询
|
||
date_user_message_a_set.add(date_user_a + date_user_b)
|
||
|
||
date_user_message_a_set.add(date_user_b_clean + date_user_a_append)
|
||
coro_list = list1 + list2
|
||
coro_list.sort(key=lambda x: x.created_at, reverse=False)
|
||
logger.info(f"查询{date_user_a}和{date_user_b}对话共:{len(coro_list)}")
|
||
# 获取意图分类
|
||
summary_res = await self.intent_message(coro_list, intent_list)
|
||
summary_list.append(summary_res)
|
||
|
||
logger.info(f"总共:{summary_list}")
|
||
|
||
logger.info(f"识别到意图数量:{len(summary_list)}")
|
||
|
||
#将意图分类结果进行合并
|
||
summary_merged :Dict[str, List[str]] = await self.merge_intents_from_content(summary_list)
|
||
logger.info(f"合并后的意图数量:{len(summary_merged)},{summary_merged}")
|
||
|
||
check_text = TextIntent()
|
||
text_extract = TextExtract()
|
||
new_summary_merged: Dict[str, List[str]]={}
|
||
|
||
for key, value in summary_merged.items():
|
||
if key in standard_intent:
|
||
logger.info(f"已存在标准意图:{key}")
|
||
continue
|
||
# 重排序,取最相近的2个
|
||
reranked_value:List[str] = await self.rerank_merged_intents(value, key)
|
||
# 检查意图, 并返回最匹配的意图
|
||
intent_key: str = await check_text.intent_message(str(reranked_value), key, intent_list)
|
||
logger.info(f"老的意图:{key},新的意图:{intent_key}, value:{reranked_value}")
|
||
if intent_key in standard_intent:
|
||
logger.info(f"已存在标准意图:{intent_key}")
|
||
continue
|
||
key = intent_key
|
||
# 跟标准进行重排序
|
||
std_pairs = [[key, std] for std in standard_intent]
|
||
try:
|
||
std_scores = self.reranker.compute_score(std_pairs, normalize=True)
|
||
if isinstance(std_scores, list):
|
||
max_sim = max(std_scores) if std_scores else 0.0
|
||
else:
|
||
max_sim = std_scores or 0.0
|
||
if max_sim > 0.9:
|
||
top_idx = std_scores.index(max_sim) if isinstance(std_scores, list) else 0
|
||
logger.info(f"key[{key}]与标准意图[{standard_intent[top_idx]}]相似度{max_sim:.3f}>0.9,跳过")
|
||
continue
|
||
except Exception as e:
|
||
logger.warning(f"标准意图重排序失败: {e},继续处理")
|
||
|
||
|
||
reranked_value_tmp: List[str] = []
|
||
for answers in reranked_value:
|
||
# 文本改写
|
||
answer = await text_extract.text_extract(answers)
|
||
reranked_value_tmp.append(answer)
|
||
reranked_value = reranked_value_tmp
|
||
#从milvus中获取
|
||
milvus_content = await intent_repo.get_content_by_intent_category(key)
|
||
#合并
|
||
merged = list(dict.fromkeys(reranked_value + milvus_content))
|
||
if len(merged) > 2:
|
||
# 重新排序
|
||
merged = await self.rerank_merged_intents(merged, key, top_k=1)
|
||
|
||
for content in merged:
|
||
insert_data = []
|
||
batch_embeddings = await self.embedding_client.aembed_documents([content])
|
||
if batch_embeddings:
|
||
insert_data.append({
|
||
"intent_category": key,
|
||
"content": content,
|
||
"content_vector": batch_embeddings[0]
|
||
})
|
||
if insert_data:
|
||
await intent_repo.upsert_intent_data(insert_data)
|
||
logger.info(f"key[{key}]入库完成,upsert {len(insert_data)} 条")
|
||
|
||
logger.info(f"运行结束:{(time.time() - stime)}")
|
||
|
||
|
||
async def intent_message(self, coro_list: List[ArchiveMessages], intent_list:list[str]) -> str:
|
||
|
||
lines = [
|
||
f"{msg.created_at} {msg.from_role}:{msg.content}"
|
||
for msg in coro_list
|
||
if not (msg.msgtype == "voice" and "[语音]" not in msg.content)
|
||
]
|
||
message_str = "\n".join(lines)
|
||
|
||
prompt_template = PromptTemplate(template=intent_prompt, input_variables=["message_str","intent_list"])
|
||
chain = prompt_template | llm
|
||
result = await chain.ainvoke({"message_str": message_str,"intent_list": str(intent_list)})
|
||
res_new = str(result.content) if hasattr(
|
||
result, 'content') else str(result)
|
||
return res_new
|
||
|
||
async def merge_intents_from_content(self, summary_list:list[str]) -> Dict[str, List[str]]:
|
||
merged: Dict[str, List[str]] = {}
|
||
for content in summary_list:
|
||
|
||
json_candidates = self._extract_jsons(content)
|
||
|
||
for json_str in json_candidates:
|
||
try:
|
||
data = json.loads(json_str)
|
||
qa_pairs = data.get("qa_pairs", [])
|
||
for pair in qa_pairs:
|
||
intent = pair.get("intents", "")
|
||
answer = pair.get("answer", "")
|
||
if intent and answer and len(answer) > 6:
|
||
|
||
if intent not in merged:
|
||
merged[intent] = []
|
||
merged[intent].append(answer)
|
||
except (json.JSONDecodeError, AttributeError):
|
||
continue
|
||
|
||
#logger.info(f"合并完成,共 {len(merged)} 个intent")
|
||
return merged
|
||
|
||
def _extract_jsons(self, content: str) -> List[str]:
|
||
results = []
|
||
start = 0
|
||
while start < len(content):
|
||
brace_start = content.find("{", start)
|
||
if brace_start == -1:
|
||
break
|
||
depth = 0
|
||
end = brace_start
|
||
for i in range(brace_start, len(content)):
|
||
if content[i] == "{":
|
||
depth += 1
|
||
elif content[i] == "}":
|
||
depth -= 1
|
||
if depth == 0:
|
||
end = i + 1
|
||
break
|
||
if end > brace_start:
|
||
results.append(content[brace_start:end])
|
||
start = end
|
||
else:
|
||
start = brace_start + 1
|
||
return results
|
||
|
||
|
||
def _calc_completeness_score(self, answer: str, query: str) -> float:
|
||
if not answer:
|
||
return 0.0
|
||
|
||
length_score = min(len(answer) / 50.0, 1.0)
|
||
|
||
punctuations = "。!?!?::;;,,"
|
||
punct_score = 1.0 if any(p in answer for p in punctuations) else 0.0
|
||
|
||
modal_particles = ["哦", "呢", "啊", "吧", "嘛", "呀", "啦"]
|
||
modal_penalty = 0.3 if any(answer.rstrip().endswith(m) for m in modal_particles) else 0.0
|
||
|
||
unique_ratio = len(set(answer)) / max(len(answer), 1)
|
||
diversity_score = min(unique_ratio / 0.4, 1.0)
|
||
|
||
query_chars = set(query)
|
||
answer_chars = set(answer)
|
||
overlap = query_chars & answer_chars
|
||
keyword_score = min(len(overlap) / max(len(query_chars), 1), 1.0)
|
||
|
||
score = (
|
||
0.30 * length_score
|
||
+ 0.20 * punct_score
|
||
+ 0.20 * diversity_score
|
||
+ 0.20 * keyword_score
|
||
- 0.10 * modal_penalty
|
||
)
|
||
return max(0.0, min(score, 1.0))
|
||
|
||
def _filter_invalid_answers(self, answers: List[str]) -> List[str]:
|
||
modal_particles = ["哦", "呢", "啊", "吧", "嘛", "呀", "啦"]
|
||
filtered = []
|
||
for a in answers:
|
||
if len(a) < 6:
|
||
continue
|
||
stripped = a.strip()
|
||
if not stripped:
|
||
continue
|
||
if len(stripped) <= 3 and any(stripped.endswith(m) for m in modal_particles):
|
||
continue
|
||
if len(set(stripped)) < max(2, len(stripped) // 3):
|
||
continue
|
||
filtered.append(a)
|
||
return filtered
|
||
|
||
async def rerank_merged_intents(self, intent_value: List[str], query: str, top_k: int = 2, max_rerank: int = 20) -> List[str]:
|
||
if not intent_value:
|
||
return []
|
||
|
||
filtered = self._filter_invalid_answers(intent_value)
|
||
if not filtered:
|
||
logger.warning(f"所有答案均被规则过滤,返回原数据前{top_k}条")
|
||
return intent_value[:top_k]
|
||
|
||
sorted_by_len = sorted(filtered, key=lambda x: len(x), reverse=True)
|
||
rerank_candidates = sorted_by_len[:max_rerank]
|
||
|
||
if self.reranker is None:
|
||
logger.warning("reranker不可用,直接返回前top_k条")
|
||
return rerank_candidates[:top_k]
|
||
|
||
pairs = [[query, answer] for answer in rerank_candidates]
|
||
try:
|
||
scores = self.reranker.compute_score(pairs, normalize=True)
|
||
except Exception as e:
|
||
logger.error(f"rerank失败: {e},直接返回前top_k条")
|
||
return rerank_candidates[:top_k]
|
||
|
||
if not isinstance(scores, list):
|
||
scores = [scores]
|
||
|
||
final_scores = []
|
||
for i, answer in enumerate(rerank_candidates):
|
||
rerank_score = scores[i] if i < len(scores) else 0.0
|
||
completeness_score = self._calc_completeness_score(answer, query)
|
||
fused = 0.7 * rerank_score + 0.3 * completeness_score
|
||
final_scores.append((answer, fused, rerank_score, completeness_score))
|
||
|
||
final_scores.sort(key=lambda x: x[1], reverse=True)
|
||
print(final_scores)
|
||
result = [item[0] for item in final_scores[:top_k]]
|
||
|
||
details = ", ".join(
|
||
f"'{a[:20]}...'(融合={f:.3f}, rerank={r:.3f},完整={c:.3f})"
|
||
for a, f, r, c in final_scores[:top_k]
|
||
)
|
||
logger.info(f"rerank完成: 过滤前{len(intent_value)}条→过滤后{len(filtered)}条→rerank{len(rerank_candidates)}条→top{top_k}: {details}")
|
||
return result
|
||
|
||
async def build(day:str):
|
||
logger.info(f"执行日期:{day}")
|
||
db_assistant_mysql_client_manager.init()
|
||
embedding_client.init()
|
||
milvus_client.init()
|
||
try:
|
||
async with db_assistant_mysql_client_manager.session_factory() as db_assistant:
|
||
archive_messages_repository = ArchiveMessagesRepository(db_assistant)
|
||
summary_repository = SummaryRepository()
|
||
reranker = get_reranker_model()
|
||
service = HistoryIntent(archive_messages_repository, embedding_client, summary_repository, reranker)
|
||
await service.get_message_summary(day)
|
||
|
||
# intent_value: List[str] = ['公司介绍','课程咨询','3D建模','3D场景建模','3D角色建模','2D原画','二次元风格','写实风格','动作/分镜','游戏动作/分镜','开发','UE引擎开发','3D大师课','GGAC绘殓课程','大学生扶持计划','兴趣/提升班','学习模式','线下面授教学','线上直播教学','寒暑假教学','提供服务','优势/竞争力','校区分布','退费咨询','住宿咨询','学不会怎么办','就业咨询','学历咨询']
|
||
# pairs = [["场景建模", answer] for answer in intent_value]
|
||
#
|
||
# scores = reranker.compute_score(pairs, normalize=True)
|
||
# print(scores)
|
||
|
||
|
||
|
||
finally:
|
||
await embedding_client.close()
|
||
await db_assistant_mysql_client_manager.close()
|
||
milvus_client.close()
|
||
|
||
if __name__ == '__main__':
|
||
|
||
start = datetime(2026, 7, 21)
|
||
end = datetime(2026, 7, 21)
|
||
|
||
result = []
|
||
temp = start
|
||
while temp <= end:
|
||
# 按 年-月-日 无补零格式输出
|
||
date_format = f"{temp.year}-{temp.month}-{temp.day}"
|
||
|
||
result.append(date_format)
|
||
temp += timedelta(days=1)
|
||
for date in result:
|
||
asyncio.run(build(date))
|
||
|
||
|
||
|