285 lines
8.4 KiB
Python
285 lines
8.4 KiB
Python
"""
|
||
查询意图识别模块
|
||
|
||
功能:识别用户查询是否为统计类查询
|
||
"""
|
||
|
||
import re
|
||
|
||
# ========== 统计类查询关键词 ==========
|
||
|
||
STAT_AGG_KEYWORDS = [
|
||
# 中文常见统计类关键词(已优化)
|
||
"多少", "几种", "几类", "几个", "多少个",
|
||
"数量", "总数", "总量", "总共", "合计", "总计", "共", "共计", "一共有",
|
||
"比例", "占比", "比率", "百分比", "占多少", "占了多少",
|
||
"平均", "均值", "平均值",
|
||
"最大值", "最小值", "最高", "最低", "最大", "最小",
|
||
"分布", "分布情况", "分布状况",
|
||
"排行", "排名", "前几名", "前几", "前十",
|
||
"频次", "累计", "累加", "累积",
|
||
"超限次数", "告警次数", "故障次数", "异常次数", "出现次数", "发生次数",
|
||
"是否存在", "有没有", # 隐含 count > 0
|
||
"增长", "下降", "变化趋势", "趋势", "增长率", "减少率"
|
||
]
|
||
|
||
EN_STAT_AGG_PATTERNS = [
|
||
r"\bhow many\b",
|
||
r"\bhow much\b",
|
||
r"\bcount\b",
|
||
r"\bnumber of\b",
|
||
r"\bamount of\b",
|
||
r"\btotal\b",
|
||
r"\baverage\b",
|
||
r"\bmean\b",
|
||
r"\bmax(imum)?\b",
|
||
r"\bmin(imum)?\b",
|
||
r"\bsum\b",
|
||
r"\bmedian\b",
|
||
r"\bvariance\b",
|
||
r"\bstd\b",
|
||
r"\bratio\b",
|
||
r"\bpercentage\b",
|
||
r"\bproportion\b",
|
||
r"\bdistribution\b",
|
||
r"\branking\b",
|
||
r"\btop\s*\d+\b",
|
||
r"\boccurrence(s)?\b",
|
||
r"\bfrequency\b",
|
||
r"\bgrowth rate\b",
|
||
r"\bexists?\b",
|
||
r"\bwhether\b.*\b(any|exist)\b"
|
||
]
|
||
|
||
|
||
def is_aggregate_query(query: str) -> bool:
|
||
"""
|
||
简单的统计类查询意图识别
|
||
|
||
规则优先,不调用LLM,尽量覆盖常见"多少 / 比例 / 总数 / 平均 / 分布 / 排名 / 频率"等问法
|
||
|
||
Args:
|
||
query: 用户查询字符串
|
||
|
||
Returns:
|
||
bool: 如果是统计类查询返回True,否则返回False
|
||
"""
|
||
if not query:
|
||
return False
|
||
q = query.strip()
|
||
|
||
# 中文关键词命中
|
||
for kw in STAT_AGG_KEYWORDS:
|
||
if kw in q:
|
||
return True
|
||
|
||
# 英文模式命中
|
||
lower_q = q.lower()
|
||
for pat in EN_STAT_AGG_PATTERNS:
|
||
if re.search(pat, lower_q):
|
||
return True
|
||
|
||
return False
|
||
|
||
|
||
# ========== 时间相关关键词 ==========
|
||
TIME_KEYWORDS = [
|
||
"年", "月", "日", "季度",
|
||
"今年", "去年", "前年", "明年",
|
||
"本月", "上月",
|
||
"今天", "昨天",
|
||
"同比", "环比", "较上年", "比去年", "相比去年",
|
||
"年初", "年末", "季初", "季末",
|
||
"近一年", "近半年", "近三个月", "近一周", "近期",
|
||
"过去一年", "过去一个月",
|
||
"截至", "截止", "期间", "时间段"
|
||
]
|
||
|
||
# 移除了易误判的:"发生", "时间", "日期"(除非配合具体时间格式)
|
||
|
||
ZH_TIME_PATTERNS = [
|
||
r"\b\d{4}年\b", # 2025年
|
||
r"\b\d{2}年\b", # 25年(关键!)
|
||
r"\b\d{1,2}月\b", # 3月、12月
|
||
r"\b\d{1,2}日\b", # 5日
|
||
r"\b近[一二三四五六七八九十\d]+[天月年]\b",
|
||
r"\b过去[一二三四五六七八九十\d]+[天月年]\b",
|
||
r"\b\d{4}[-/年]\d{1,2}([-/月]\d{1,2})?\b", # 2025-12-01, 2025/12, 2025年12月
|
||
]
|
||
|
||
EN_TIME_PATTERNS = [
|
||
r"\b\d{4}\b", # 年份(独立出现,如 "in 2025")
|
||
r"\byear\b", r"\bmonth\b", r"\bday\b",
|
||
r"\bthis year\b", r"\blast year\b", r"\bnext year\b",
|
||
r"\bthis month\b", r"\blast month\b",
|
||
r"\btoday\b", r"\byesterday\b",
|
||
r"\byear over year\b", r"\bmonth over month\b",
|
||
r"\bq[1-4]\b",
|
||
r"\bjan|feb|mar|apr|may|jun|jul|aug|sep|oct|nov|dec\b",
|
||
r"\brecent\b",
|
||
r"\bpast\s+\d+\s+(day|week|month|year)s?\b",
|
||
r"\blast\s+\d+\s+(day|week|month|year)s?\b",
|
||
r"\bas of\b",
|
||
r"\b\d{4}[-/]\d{1,2}([-/]\d{1,2})?\b", # ISO or slash date
|
||
]
|
||
|
||
|
||
def has_time_in_query(query: str) -> bool:
|
||
"""
|
||
判断查询中是否包含时间相关信息
|
||
|
||
Args:
|
||
query: 用户查询字符串
|
||
|
||
Returns:
|
||
bool: 如果包含时间信息返回True,否则返回False
|
||
"""
|
||
if not query:
|
||
return False
|
||
|
||
q = query.strip()
|
||
lower_q = q.lower()
|
||
|
||
# 中文时间关键词
|
||
for kw in TIME_KEYWORDS:
|
||
if kw in q:
|
||
return True
|
||
|
||
# 中文时间模式(如 "2025年"、"25年"、"3月"、"2025-12-01" 等)
|
||
for pat in ZH_TIME_PATTERNS:
|
||
if re.search(pat, q):
|
||
return True
|
||
|
||
# 英文时间模式
|
||
for pat in EN_TIME_PATTERNS:
|
||
if re.search(pat, lower_q):
|
||
return True
|
||
|
||
return False
|
||
|
||
|
||
def get_current_date_info() -> dict:
|
||
"""
|
||
获取当前系统时间信息(用于注入到Prompt中)
|
||
|
||
Returns:
|
||
dict: 包含当前日期信息的字典
|
||
"""
|
||
from datetime import datetime
|
||
now = datetime.now()
|
||
|
||
# 计算下个月的第一天
|
||
if now.month == 12:
|
||
# 如果是12月,下个月是下一年的1月
|
||
next_month_start = datetime(now.year + 1, 1, 1).strftime("%Y-%m-%d")
|
||
else:
|
||
# 否则是当前年的下一个月
|
||
next_month_start = datetime(now.year, now.month + 1, 1).strftime("%Y-%m-%d")
|
||
|
||
return {
|
||
"current_year": str(now.year),
|
||
"current_month": str(now.month),
|
||
"current_day": str(now.day),
|
||
"current_date": now.strftime("%Y-%m-%d"),
|
||
"current_datetime": now.strftime("%Y-%m-%d %H:%M:%S"),
|
||
"last_year": str(now.year - 1),
|
||
"next_year": str(now.year + 1),
|
||
"next_month_start": next_month_start,
|
||
}
|
||
|
||
|
||
# ========== 时间段统计类型识别 ==========
|
||
|
||
# 双时间段对比关键词
|
||
DUAL_TIME_COMPARISON_KEYWORDS = [
|
||
"同比", "环比", "较上年", "比去年", "相比去年", "相比", "对比",
|
||
"比", "较", "vs", "VS", "对比", "比较",
|
||
"多多少", "少多少", "增加", "减少", "增长", "下降",
|
||
"比...多", "比...少", "比...增加", "比...减少",
|
||
]
|
||
|
||
DUAL_TIME_COMPARISON_PATTERNS = [
|
||
r"比.*[多少]", # 比去年多、比去年少
|
||
r"较.*[多少]", # 较去年多
|
||
r"相比.*[多少]", # 相比去年多
|
||
r"vs\s+", # vs
|
||
r"对比\s+", # 对比
|
||
r"year over year", # 同比
|
||
r"month over month", # 环比
|
||
r"compared to", # 相比
|
||
r"vs\.", # vs.
|
||
]
|
||
|
||
# 多时间段关键词(3个及以上时间段)
|
||
MULTI_TIME_KEYWORDS = [
|
||
"各年", "每年", "历年", "逐年",
|
||
"各月", "每月", "各季度", "每季度",
|
||
"分别", "分别统计",
|
||
]
|
||
|
||
MULTI_TIME_PATTERNS = [
|
||
r"各[年月季度周天]", # 各年、各月、各季度
|
||
r"每[年月季度周天]", # 每年、每月、每季度
|
||
r"历年", # 历年
|
||
r"逐年", # 逐年
|
||
r"分别", # 分别
|
||
r"each\s+(year|month|quarter)", # each year, each month
|
||
r"per\s+(year|month|quarter)", # per year, per month
|
||
]
|
||
|
||
|
||
def classify_time_aggregate_type(query: str) -> str:
|
||
"""
|
||
分类时间统计查询类型
|
||
|
||
类型:
|
||
- "single": 单时间段统计(如"2024年发生的故障有多少?")
|
||
- "dual": 双时间段统计(如"今年比去年多吗?"、"25年异常同比减少多少?")
|
||
- "multi": 多时间段统计(如"2023、2024、2025年的故障数"、"各年的故障数")
|
||
|
||
Args:
|
||
query: 用户查询字符串
|
||
|
||
Returns:
|
||
str: 时间段统计类型("single"、"dual"、"multi")
|
||
"""
|
||
if not query:
|
||
return "single"
|
||
|
||
q = query.strip()
|
||
lower_q = q.lower()
|
||
|
||
# 1. 检查是否为多时间段统计
|
||
for kw in MULTI_TIME_KEYWORDS:
|
||
if kw in q:
|
||
return "multi"
|
||
|
||
for pat in MULTI_TIME_PATTERNS:
|
||
if re.search(pat, q) or re.search(pat, lower_q):
|
||
return "multi"
|
||
|
||
# 检查是否包含多个明确的时间段(如"2023、2024、2025年")
|
||
# 匹配多个年份或时间段
|
||
year_matches = len(re.findall(r"\d{4}年|\d{2}年", q))
|
||
if year_matches >= 2:
|
||
# 检查是否有对比关键词,如果没有,可能是多时间段
|
||
has_comparison = False
|
||
for kw in DUAL_TIME_COMPARISON_KEYWORDS:
|
||
if kw in q:
|
||
has_comparison = True
|
||
break
|
||
if not has_comparison and year_matches >= 2:
|
||
return "multi"
|
||
|
||
# 2. 检查是否为双时间段统计
|
||
for kw in DUAL_TIME_COMPARISON_KEYWORDS:
|
||
if kw in q:
|
||
return "dual"
|
||
|
||
for pat in DUAL_TIME_COMPARISON_PATTERNS:
|
||
if re.search(pat, q) or re.search(pat, lower_q):
|
||
return "dual"
|
||
|
||
# 3. 默认返回单时间段统计
|
||
return "single"
|