Files
iAOP/templates/ti-cl4/llm-scenarios/tests/test_nl_query.py
T

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()