kgrag/graph_search/query_intent.py
2026-07-29 18:10:19 +08:00

285 lines
8.4 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 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"