| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300 |
- """Doris 自然语言查询助手 —— 基于 openai-agents-python + glm-5 (百炼) + Gradio WebUI。
- 用法:
- python app.py --build-schema # 重新抓取 ods 库的表/列/注释到 schema.json
- python app.py --selftest # 不连 Doris,只验证 LLM 接线
- python app.py # 启动 WebUI (默认 http://0.0.0.0:7860)
- """
- import argparse
- import asyncio
- import json
- import os
- import re
- from functools import lru_cache
- import pymysql
- from dotenv import load_dotenv
- load_dotenv()
- # --- 配置 -----------------------------------------------------------------
- DB = dict(
- host=os.getenv("DB_HOST"),
- port=int(os.getenv("DB_PORT", "9030")),
- user=os.getenv("DB_USER"),
- password=os.getenv("DB_PASSWORD"),
- database=os.getenv("DB_NAME", "ods"),
- charset="utf8mb4",
- connect_timeout=10,
- read_timeout=30,
- write_timeout=30,
- autocommit=True,
- )
- SCHEMA_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), "schema.json")
- def _conn(retries: int = 4, delay: float = 1.5):
- """连 Doris。经 VPN 的链路会在 MySQL 握手阶段偶发掉线(2013),这里重试。"""
- import time
- last = None
- for i in range(retries):
- try:
- return pymysql.connect(**DB)
- except pymysql.err.OperationalError as e:
- last = e
- if i < retries - 1:
- time.sleep(delay)
- raise last
- # --- schema 缓存 ----------------------------------------------------------
- def build_schema():
- """从 Doris INFORMATION_SCHEMA 抓取所有表/列/注释,写成 schema.json。"""
- conn = _conn()
- try:
- with conn.cursor(pymysql.cursors.DictCursor) as cur:
- cur.execute(
- "SELECT TABLE_NAME, TABLE_COMMENT FROM INFORMATION_SCHEMA.TABLES "
- "WHERE TABLE_SCHEMA=%s",
- (DB["database"],),
- )
- tables = cur.fetchall()
- cur.execute(
- "SELECT TABLE_NAME, COLUMN_NAME, COLUMN_TYPE, COLUMN_COMMENT, IS_NULLABLE "
- "FROM INFORMATION_SCHEMA.COLUMNS WHERE TABLE_SCHEMA=%s "
- "ORDER BY TABLE_NAME, ORDINAL_POSITION",
- (DB["database"],),
- )
- cols = cur.fetchall()
- finally:
- conn.close()
- schema = {t["TABLE_NAME"]: {"comment": t["TABLE_COMMENT"] or "", "columns": []}
- for t in tables}
- for c in cols:
- schema.setdefault(c["TABLE_NAME"], {"comment": "", "columns": []})["columns"].append(
- {
- "name": c["COLUMN_NAME"],
- "type": c["COLUMN_TYPE"],
- "comment": c["COLUMN_COMMENT"] or "",
- "nullable": c["IS_NULLABLE"],
- }
- )
- with open(SCHEMA_PATH, "w", encoding="utf-8") as f:
- json.dump(schema, f, ensure_ascii=False, indent=2)
- print(f"schema.json 已生成,共 {len(schema)} 张表。")
- return schema
- @lru_cache(maxsize=1)
- def load_schema():
- if not os.path.exists(SCHEMA_PATH):
- try:
- return build_schema()
- except Exception as e: # selftest 等场景下 Doris 不可达时也能继续
- print(f"[警告] 无法连接 Doris 生成 schema,使用空 schema: {e}")
- return {}
- with open(SCHEMA_PATH, encoding="utf-8") as f:
- return json.load(f)
- # --- Agent 设置 -----------------------------------------------------------
- from agents import ( # noqa: E402
- Agent,
- OpenAIChatCompletionsModel,
- Runner,
- function_tool,
- set_tracing_disabled,
- )
- from openai import AsyncOpenAI # noqa: E402
- # 框架默认会把 trace 上报到 OpenAI,用国产模型时必须关掉,否则会报错/泄露。
- set_tracing_disabled(True)
- _client = AsyncOpenAI(
- base_url=os.getenv("LLM_BASE_URL"),
- api_key=os.getenv("LLM_API_KEY"),
- )
- _model = OpenAIChatCompletionsModel(
- model=os.getenv("LLM_MODEL", "glm-5"),
- openai_client=_client,
- )
- SCHEMA = load_schema()
- # 关键表的人工说明(平台字典为空,这些规则是探查数据得出的,agent 易出错的地方)。
- TABLE_NOTES = {
- "ods_sa_device_cbdphoto_b": (
- "测报灯图片识别【主表,约2000万行,查虫量默认用这个】。"
- "indentify_result 格式 '虫码,数量#虫码,数量'(虫码=ods_base_pest_code.pest_yfkj_code,数量=头数),"
- "例 '3,1#53,3#136,1'。device_id = ods_sa_device.id。时间字段 uptime(datetime)。"
- "取某虫数量: regexp_extract(indentify_result,'(^|#)<虫码>,([0-9]+)',2),配 RLIKE '(^|#)<虫码>,' 过滤。"
- ),
- "ods_sa_device_lpsphoto": "性诱识别(约2.6万行)。indentify_result='虫码,数量'。device_id=ods_sa_device.id。时间 uptime。",
- "ods_sa_device_lpsphoto_count": "性诱识别计数(约2.6万行),格式同 lpsphoto。",
- "ods_sa_device": (
- "设备主表。province/city/district 值带后缀(如'河南省''新乡市')。"
- "device_id 是 IMEI;id 是数字主键,被各识别表的 device_id 引用。"
- ),
- "ods_base_pest_code": "害虫字典(1226种)。pest_yfkj_code=识别结果里的虫码;pest_name=虫名;pest_level=等级;pest_code 是另一套编码,别混用。",
- "ods_sa_device_cbd_data": "测报灯设备遥测(JSON),只有温湿度/电池/灯状态等,【没有虫种和数量】,别用来查虫量。",
- }
- # --- tools ---------------------------------------------------------------
- _DML = re.compile(
- r"\b(insert|update|delete|drop|alter|create|truncate|grant|revoke|load|merge)\b"
- )
- @function_tool
- def list_tables(keyword: str = "") -> str:
- """列出 ods 库的表。可传 keyword 按表名或注释过滤。返回 '表名 - 注释'。"""
- items = [(n, d["comment"]) for n, d in SCHEMA.items()]
- if keyword:
- k = keyword.lower()
- items = [(n, c) for n, c in items if k in n.lower() or k in (c or "").lower()]
- items = items[:50]
- return "\n".join(f"{n} - {c}" for n, c in items) or "没有匹配的表。"
- @function_tool
- def get_table_schema(table_name: str) -> str:
- """获取某张表的列定义、类型和注释。"""
- t = SCHEMA.get(table_name) or SCHEMA.get(table_name.lower())
- if not t:
- cands = [n for n in SCHEMA if table_name.lower() in n.lower()][:10]
- return f"表 {table_name} 不存在。相近的表: {cands}"
- lines = [f"表 {table_name}: {t['comment']}"]
- if TABLE_NOTES.get(table_name):
- lines.append(f" ⚑ 数据说明: {TABLE_NOTES[table_name]}")
- for col in t["columns"]:
- lines.append(
- f" - {col['name']} {col['type']} "
- f"{'NULL' if col['nullable'] == 'YES' else 'NOT NULL'} -- {col['comment']}"
- )
- return "\n".join(lines)
- @function_tool
- def execute_sql(sql: str) -> str:
- """对 Doris 执行只读 SELECT 查询,返回最多 200 行结果。只允许 SELECT。"""
- s = sql.strip().rstrip(";").strip()
- low = re.sub(r"/\*.*?\*/", " ", s, flags=re.S) # 去掉块注释
- low = re.sub(r"--[^\n]*", " ", low) # 去掉行注释
- if not low.strip().lower().startswith("select") or _DML.search(low.lower()):
- return "错误: 只允许只读 SELECT 查询。"
- conn = _conn()
- try:
- with conn.cursor() as cur:
- # ponytail: 把用户 SQL 包进子查询再 LIMIT,是只读/行数上限的主要保障
- # (只读账号 + 子查询无法 DML/DDL),上面的正则只是第二道防线。
- cur.execute("SET query_timeout = 30")
- cur.execute(f"SELECT * FROM ({s}) AS _q LIMIT 200")
- cols = [d[0] for d in cur.description]
- rows = cur.fetchall()
- if not rows:
- return "查询成功,但没有数据。"
- body = "\n".join(" | ".join(str(v) for v in r) for r in rows)
- return f"列: {cols}\n行数(已截断到200): {len(rows)}\n{body}"
- except Exception as e:
- return f"SQL 执行错误: {e}\n请根据错误修正 SQL 后重试。"
- finally:
- conn.close()
- @function_tool
- def search_pest(keyword: str) -> str:
- """按虫名(或虫码)查害虫编码。返回 pest_yfkj_code | 虫名 | 等级。
- 查询某虫的数量前【必须】先调用本工具拿到 pest_yfkj_code(共1226种,不要靠记忆猜)。"""
- conn = _conn()
- try:
- with conn.cursor() as cur:
- cur.execute(
- "SELECT pest_yfkj_code, pest_name, pest_level FROM ods_base_pest_code "
- "WHERE pest_name LIKE %s OR pest_yfkj_code = %s LIMIT 20",
- (f"%{keyword}%", keyword),
- )
- rows = cur.fetchall()
- finally:
- conn.close()
- if not rows:
- return f"没找到含 '{keyword}' 的害虫,请换关键词(如'夜蛾''飞虱')。"
- return "\n".join(f"{c} | {n} | {lv}" for c, n, lv in rows)
- INSTRUCTIONS = """你是农业病虫害监测数据平台的自然语言查询助手。底层是 Apache Doris,库 ods。用户用自然语言问"某地某时某虫监测到多少头",你转成 SQL 查询并用中文回答。
- ==== 数据模型(务必遵守,这是查对虫量的关键)====
- "虫量(头数)"来自设备的**图片识别结果**,不是设备遥测。
- 1. 主表 ods_sa_device_cbdphoto_b(测报灯,约2000万行,查虫量默认用它):每行=一张虫情照片。
- - indentify_result 格式 `虫码,数量#虫码,数量`,例 "3,1#53,3#136,1" = 虫码3计1头、53计3头、136计1头。
- - 虫码 = ods_base_pest_code.pest_yfkj_code。
- - 时间字段 uptime(datetime)。
- - 取某虫数量: regexp_extract(indentify_result, '(^|#)<虫码>,([0-9]+)', 2) ,并用 indentify_result RLIKE '(^|#)<虫码>,' 过滤后再 SUM。
- 2. 地理: 识别表 device_id = ods_sa_device.id(数字主键,不是IMEI)。JOIN ods_sa_device d ON d.id = <识别表>.device_id,取 d.province/d.city/district(值带后缀:"河南省""新乡市")。
- 3. 虫名→虫码: 先 search_pest(虫名) 拿 pest_yfkj_code。
- 4. device_data 等 JSON 遥测字段没有虫,别用。
- ==== 工作流程 ====
- 1. search_pest 把虫名转 pest_yfkj_code(多个结果挑最匹配并说明)。
- 2. 必要时 get_table_schema 看表结构(默认查虫量用 ods_sa_device_cbdphoto_b)。
- 3. 写 SQL:识别表 JOIN ods_sa_device 过滤省市 + uptime 时间段 + regexp 取该虫数量求 SUM。
- 4. 报错就修正重试(最多3次)。
- 5. 中文回答:给总头数 + 覆盖地区/时间/虫名 + 末尾附 SQL 代码块。
- ==== 注意 ====
- - "某月"按自然月;时间范围拿不准先 SELECT MIN/MAX(uptime) 确认。
- - 大表查询【必须】带时间范围过滤,避免全表扫描。
- - 若该虫查出来是 0 或没有:如实说明——这类设备(测报灯/性诱)主要监测趋光/趋性诱的蛾、飞虱、叶蝉等,不监测蝗虫等非趋光昆虫,不要编造数字。
- - Doris 与 MySQL 有差异,日期用 'YYYY-MM-DD'。
- """
- agent = Agent(
- name="Doris查询助手",
- instructions=INSTRUCTIONS,
- model=_model,
- tools=[list_tables, get_table_schema, search_pest, execute_sql],
- )
- # --- WebUI ---------------------------------------------------------------
- def _run(user_message: str, history: list[dict]):
- # 历史由 Gradio 管理(每用户每标签页独立),直接拼成 input 交给 Runner。
- result = asyncio.run(
- Runner.run(agent, history + [{"role": "user", "content": user_message}])
- )
- return result.final_output
- def build_ui():
- import gradio as gr
- # gradio 6 默认 history 即为 OpenAI 消息字典格式 {"role","content"}。
- return gr.ChatInterface(
- fn=_run,
- title="Doris 数据查询助手",
- description="用自然语言查询 ods 库。例:『上个月每天的订单量』",
- textbox=gr.Textbox(placeholder="问点什么,比如:每个地区近 7 天的销售额"),
- )
- def main():
- p = argparse.ArgumentParser()
- p.add_argument("--build-schema", action="store_true")
- p.add_argument("--selftest", action="store_true")
- args = p.parse_args()
- if args.build_schema:
- build_schema()
- return
- if args.selftest:
- r = asyncio.run(Runner.run(agent, "只回复 OK"))
- print("模型自检返回:", repr(r.final_output))
- return
- build_ui().launch(server_name="0.0.0.0", server_port=7860)
- if __name__ == "__main__":
- main()
|