107 lines
3.8 KiB
Python
107 lines
3.8 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""NL→SQL/API 查询翻译器测试(issue #76)。
|
|
|
|
覆盖:
|
|
1. 配置资产加载(指标字典/意图关键词/表/默认窗口);
|
|
2. 意图识别(trend/latest/kpi/alarm + 未识别降级 unsupported);
|
|
3. 指标映射(自然语言名 → point_id,长名优先);
|
|
4. 时间范围抽取("最近 1 小时" → 1h);
|
|
5. TDengine SQL 生成(latest/kpi/alarm/trend 模板);
|
|
6. to_api_params(驾驶舱 API 调用参数)。
|
|
"""
|
|
import os
|
|
import sys
|
|
import unittest
|
|
|
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
|
import _bootstrap # noqa: F401
|
|
|
|
from ti_scenarios.nl_query import NLQueryTranslator # noqa: E402
|
|
|
|
CONFIG = os.path.join(
|
|
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
|
|
"config", "nl_query.template.yaml",
|
|
)
|
|
|
|
|
|
class TestConfigLoad(unittest.TestCase):
|
|
"""配置资产加载。"""
|
|
|
|
def setUp(self):
|
|
self.t = NLQueryTranslator.from_template_config(CONFIG)
|
|
|
|
def test_metrics_and_table(self):
|
|
self.assertEqual(self.t.metrics["氯气流量"], "CLF-01.FLOW")
|
|
self.assertEqual(self.t.table, "tpl_ti_cl4.points")
|
|
self.assertEqual(self.t.default_range, "1h")
|
|
|
|
def test_intent_keywords_loaded(self):
|
|
self.assertIn("趋势", self.t._intent["trend"])
|
|
|
|
|
|
class TestTranslate(unittest.TestCase):
|
|
"""翻译:意图/指标/时间范围 → SQL/API 参数。"""
|
|
|
|
def setUp(self):
|
|
self.t = NLQueryTranslator.from_template_config(CONFIG)
|
|
|
|
def test_trend_query(self):
|
|
q = self.t.translate("氯气流量最近1小时趋势")
|
|
self.assertEqual(q.intent, "trend")
|
|
self.assertEqual(q.metric, "氯气流量")
|
|
self.assertEqual(q.point_id, "CLF-01.FLOW")
|
|
self.assertEqual(q.time_range, "1h")
|
|
self.assertIn("INTERVAL(1m)", q.sql)
|
|
self.assertIn("avg(value)", q.sql)
|
|
self.assertIn("CLF-01.FLOW", q.sql)
|
|
|
|
def test_latest_query(self):
|
|
q = self.t.translate("炉温最新数值是多少")
|
|
self.assertEqual(q.intent, "latest")
|
|
self.assertIn("last_row", q.sql)
|
|
self.assertEqual(q.point_id, "CLF-01.TEMP")
|
|
|
|
def test_kpi_query_with_default_range(self):
|
|
q = self.t.translate("氯气流量平均")
|
|
self.assertEqual(q.intent, "kpi")
|
|
self.assertEqual(q.time_range, "") # 未识别 → 默认窗口
|
|
self.assertIn("now - 1h", q.sql) # 默认 1h
|
|
|
|
def test_alarm_query(self):
|
|
q = self.t.translate("炉温报警")
|
|
self.assertEqual(q.intent, "alarm")
|
|
self.assertIn("value > threshold", q.sql)
|
|
|
|
def test_unsupported_intent(self):
|
|
q = self.t.translate("今天天气怎么样")
|
|
self.assertEqual(q.intent, "unsupported")
|
|
self.assertEqual(q.sql, "")
|
|
|
|
def test_unknown_metric_keeps_intent(self):
|
|
# 未识别指标不阻断查询(SQL 用全表过滤,交由上层 LLM 兜底)
|
|
q = self.t.translate("进料泵转速趋势")
|
|
self.assertEqual(q.intent, "trend")
|
|
self.assertEqual(q.point_id, "")
|
|
|
|
def test_api_params(self):
|
|
q = self.t.translate("氯气流量最近1小时趋势")
|
|
params = q.to_api_params()
|
|
self.assertEqual(params["intent"], "trend")
|
|
self.assertEqual(params["metric"], "氯气流量")
|
|
self.assertEqual(params["point_id"], "CLF-01.FLOW")
|
|
self.assertEqual(params["time_range"], "1h")
|
|
|
|
|
|
class TestTimeRange(unittest.TestCase):
|
|
"""时间范围抽取。"""
|
|
|
|
def test_hour_minute_day(self):
|
|
t = NLQueryTranslator(metrics={})
|
|
self.assertEqual(t.translate("最近 2 小时趋势").time_range, "2h")
|
|
self.assertEqual(t.translate("最近30分钟走势").time_range, "30m")
|
|
self.assertEqual(t.translate("最近 3 天曲线").time_range, "3d")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|