app.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300
  1. """Doris 自然语言查询助手 —— 基于 openai-agents-python + glm-5 (百炼) + Gradio WebUI。
  2. 用法:
  3. python app.py --build-schema # 重新抓取 ods 库的表/列/注释到 schema.json
  4. python app.py --selftest # 不连 Doris,只验证 LLM 接线
  5. python app.py # 启动 WebUI (默认 http://0.0.0.0:7860)
  6. """
  7. import argparse
  8. import asyncio
  9. import json
  10. import os
  11. import re
  12. from functools import lru_cache
  13. import pymysql
  14. from dotenv import load_dotenv
  15. load_dotenv()
  16. # --- 配置 -----------------------------------------------------------------
  17. DB = dict(
  18. host=os.getenv("DB_HOST"),
  19. port=int(os.getenv("DB_PORT", "9030")),
  20. user=os.getenv("DB_USER"),
  21. password=os.getenv("DB_PASSWORD"),
  22. database=os.getenv("DB_NAME", "ods"),
  23. charset="utf8mb4",
  24. connect_timeout=10,
  25. read_timeout=30,
  26. write_timeout=30,
  27. autocommit=True,
  28. )
  29. SCHEMA_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), "schema.json")
  30. def _conn(retries: int = 4, delay: float = 1.5):
  31. """连 Doris。经 VPN 的链路会在 MySQL 握手阶段偶发掉线(2013),这里重试。"""
  32. import time
  33. last = None
  34. for i in range(retries):
  35. try:
  36. return pymysql.connect(**DB)
  37. except pymysql.err.OperationalError as e:
  38. last = e
  39. if i < retries - 1:
  40. time.sleep(delay)
  41. raise last
  42. # --- schema 缓存 ----------------------------------------------------------
  43. def build_schema():
  44. """从 Doris INFORMATION_SCHEMA 抓取所有表/列/注释,写成 schema.json。"""
  45. conn = _conn()
  46. try:
  47. with conn.cursor(pymysql.cursors.DictCursor) as cur:
  48. cur.execute(
  49. "SELECT TABLE_NAME, TABLE_COMMENT FROM INFORMATION_SCHEMA.TABLES "
  50. "WHERE TABLE_SCHEMA=%s",
  51. (DB["database"],),
  52. )
  53. tables = cur.fetchall()
  54. cur.execute(
  55. "SELECT TABLE_NAME, COLUMN_NAME, COLUMN_TYPE, COLUMN_COMMENT, IS_NULLABLE "
  56. "FROM INFORMATION_SCHEMA.COLUMNS WHERE TABLE_SCHEMA=%s "
  57. "ORDER BY TABLE_NAME, ORDINAL_POSITION",
  58. (DB["database"],),
  59. )
  60. cols = cur.fetchall()
  61. finally:
  62. conn.close()
  63. schema = {t["TABLE_NAME"]: {"comment": t["TABLE_COMMENT"] or "", "columns": []}
  64. for t in tables}
  65. for c in cols:
  66. schema.setdefault(c["TABLE_NAME"], {"comment": "", "columns": []})["columns"].append(
  67. {
  68. "name": c["COLUMN_NAME"],
  69. "type": c["COLUMN_TYPE"],
  70. "comment": c["COLUMN_COMMENT"] or "",
  71. "nullable": c["IS_NULLABLE"],
  72. }
  73. )
  74. with open(SCHEMA_PATH, "w", encoding="utf-8") as f:
  75. json.dump(schema, f, ensure_ascii=False, indent=2)
  76. print(f"schema.json 已生成,共 {len(schema)} 张表。")
  77. return schema
  78. @lru_cache(maxsize=1)
  79. def load_schema():
  80. if not os.path.exists(SCHEMA_PATH):
  81. try:
  82. return build_schema()
  83. except Exception as e: # selftest 等场景下 Doris 不可达时也能继续
  84. print(f"[警告] 无法连接 Doris 生成 schema,使用空 schema: {e}")
  85. return {}
  86. with open(SCHEMA_PATH, encoding="utf-8") as f:
  87. return json.load(f)
  88. # --- Agent 设置 -----------------------------------------------------------
  89. from agents import ( # noqa: E402
  90. Agent,
  91. OpenAIChatCompletionsModel,
  92. Runner,
  93. function_tool,
  94. set_tracing_disabled,
  95. )
  96. from openai import AsyncOpenAI # noqa: E402
  97. # 框架默认会把 trace 上报到 OpenAI,用国产模型时必须关掉,否则会报错/泄露。
  98. set_tracing_disabled(True)
  99. _client = AsyncOpenAI(
  100. base_url=os.getenv("LLM_BASE_URL"),
  101. api_key=os.getenv("LLM_API_KEY"),
  102. )
  103. _model = OpenAIChatCompletionsModel(
  104. model=os.getenv("LLM_MODEL", "glm-5"),
  105. openai_client=_client,
  106. )
  107. SCHEMA = load_schema()
  108. # 关键表的人工说明(平台字典为空,这些规则是探查数据得出的,agent 易出错的地方)。
  109. TABLE_NOTES = {
  110. "ods_sa_device_cbdphoto_b": (
  111. "测报灯图片识别【主表,约2000万行,查虫量默认用这个】。"
  112. "indentify_result 格式 '虫码,数量#虫码,数量'(虫码=ods_base_pest_code.pest_yfkj_code,数量=头数),"
  113. "例 '3,1#53,3#136,1'。device_id = ods_sa_device.id。时间字段 uptime(datetime)。"
  114. "取某虫数量: regexp_extract(indentify_result,'(^|#)<虫码>,([0-9]+)',2),配 RLIKE '(^|#)<虫码>,' 过滤。"
  115. ),
  116. "ods_sa_device_lpsphoto": "性诱识别(约2.6万行)。indentify_result='虫码,数量'。device_id=ods_sa_device.id。时间 uptime。",
  117. "ods_sa_device_lpsphoto_count": "性诱识别计数(约2.6万行),格式同 lpsphoto。",
  118. "ods_sa_device": (
  119. "设备主表。province/city/district 值带后缀(如'河南省''新乡市')。"
  120. "device_id 是 IMEI;id 是数字主键,被各识别表的 device_id 引用。"
  121. ),
  122. "ods_base_pest_code": "害虫字典(1226种)。pest_yfkj_code=识别结果里的虫码;pest_name=虫名;pest_level=等级;pest_code 是另一套编码,别混用。",
  123. "ods_sa_device_cbd_data": "测报灯设备遥测(JSON),只有温湿度/电池/灯状态等,【没有虫种和数量】,别用来查虫量。",
  124. }
  125. # --- tools ---------------------------------------------------------------
  126. _DML = re.compile(
  127. r"\b(insert|update|delete|drop|alter|create|truncate|grant|revoke|load|merge)\b"
  128. )
  129. @function_tool
  130. def list_tables(keyword: str = "") -> str:
  131. """列出 ods 库的表。可传 keyword 按表名或注释过滤。返回 '表名 - 注释'。"""
  132. items = [(n, d["comment"]) for n, d in SCHEMA.items()]
  133. if keyword:
  134. k = keyword.lower()
  135. items = [(n, c) for n, c in items if k in n.lower() or k in (c or "").lower()]
  136. items = items[:50]
  137. return "\n".join(f"{n} - {c}" for n, c in items) or "没有匹配的表。"
  138. @function_tool
  139. def get_table_schema(table_name: str) -> str:
  140. """获取某张表的列定义、类型和注释。"""
  141. t = SCHEMA.get(table_name) or SCHEMA.get(table_name.lower())
  142. if not t:
  143. cands = [n for n in SCHEMA if table_name.lower() in n.lower()][:10]
  144. return f"表 {table_name} 不存在。相近的表: {cands}"
  145. lines = [f"表 {table_name}: {t['comment']}"]
  146. if TABLE_NOTES.get(table_name):
  147. lines.append(f" ⚑ 数据说明: {TABLE_NOTES[table_name]}")
  148. for col in t["columns"]:
  149. lines.append(
  150. f" - {col['name']} {col['type']} "
  151. f"{'NULL' if col['nullable'] == 'YES' else 'NOT NULL'} -- {col['comment']}"
  152. )
  153. return "\n".join(lines)
  154. @function_tool
  155. def execute_sql(sql: str) -> str:
  156. """对 Doris 执行只读 SELECT 查询,返回最多 200 行结果。只允许 SELECT。"""
  157. s = sql.strip().rstrip(";").strip()
  158. low = re.sub(r"/\*.*?\*/", " ", s, flags=re.S) # 去掉块注释
  159. low = re.sub(r"--[^\n]*", " ", low) # 去掉行注释
  160. if not low.strip().lower().startswith("select") or _DML.search(low.lower()):
  161. return "错误: 只允许只读 SELECT 查询。"
  162. conn = _conn()
  163. try:
  164. with conn.cursor() as cur:
  165. # ponytail: 把用户 SQL 包进子查询再 LIMIT,是只读/行数上限的主要保障
  166. # (只读账号 + 子查询无法 DML/DDL),上面的正则只是第二道防线。
  167. cur.execute("SET query_timeout = 30")
  168. cur.execute(f"SELECT * FROM ({s}) AS _q LIMIT 200")
  169. cols = [d[0] for d in cur.description]
  170. rows = cur.fetchall()
  171. if not rows:
  172. return "查询成功,但没有数据。"
  173. body = "\n".join(" | ".join(str(v) for v in r) for r in rows)
  174. return f"列: {cols}\n行数(已截断到200): {len(rows)}\n{body}"
  175. except Exception as e:
  176. return f"SQL 执行错误: {e}\n请根据错误修正 SQL 后重试。"
  177. finally:
  178. conn.close()
  179. @function_tool
  180. def search_pest(keyword: str) -> str:
  181. """按虫名(或虫码)查害虫编码。返回 pest_yfkj_code | 虫名 | 等级。
  182. 查询某虫的数量前【必须】先调用本工具拿到 pest_yfkj_code(共1226种,不要靠记忆猜)。"""
  183. conn = _conn()
  184. try:
  185. with conn.cursor() as cur:
  186. cur.execute(
  187. "SELECT pest_yfkj_code, pest_name, pest_level FROM ods_base_pest_code "
  188. "WHERE pest_name LIKE %s OR pest_yfkj_code = %s LIMIT 20",
  189. (f"%{keyword}%", keyword),
  190. )
  191. rows = cur.fetchall()
  192. finally:
  193. conn.close()
  194. if not rows:
  195. return f"没找到含 '{keyword}' 的害虫,请换关键词(如'夜蛾''飞虱')。"
  196. return "\n".join(f"{c} | {n} | {lv}" for c, n, lv in rows)
  197. INSTRUCTIONS = """你是农业病虫害监测数据平台的自然语言查询助手。底层是 Apache Doris,库 ods。用户用自然语言问"某地某时某虫监测到多少头",你转成 SQL 查询并用中文回答。
  198. ==== 数据模型(务必遵守,这是查对虫量的关键)====
  199. "虫量(头数)"来自设备的**图片识别结果**,不是设备遥测。
  200. 1. 主表 ods_sa_device_cbdphoto_b(测报灯,约2000万行,查虫量默认用它):每行=一张虫情照片。
  201. - indentify_result 格式 `虫码,数量#虫码,数量`,例 "3,1#53,3#136,1" = 虫码3计1头、53计3头、136计1头。
  202. - 虫码 = ods_base_pest_code.pest_yfkj_code。
  203. - 时间字段 uptime(datetime)。
  204. - 取某虫数量: regexp_extract(indentify_result, '(^|#)<虫码>,([0-9]+)', 2) ,并用 indentify_result RLIKE '(^|#)<虫码>,' 过滤后再 SUM。
  205. 2. 地理: 识别表 device_id = ods_sa_device.id(数字主键,不是IMEI)。JOIN ods_sa_device d ON d.id = <识别表>.device_id,取 d.province/d.city/district(值带后缀:"河南省""新乡市")。
  206. 3. 虫名→虫码: 先 search_pest(虫名) 拿 pest_yfkj_code。
  207. 4. device_data 等 JSON 遥测字段没有虫,别用。
  208. ==== 工作流程 ====
  209. 1. search_pest 把虫名转 pest_yfkj_code(多个结果挑最匹配并说明)。
  210. 2. 必要时 get_table_schema 看表结构(默认查虫量用 ods_sa_device_cbdphoto_b)。
  211. 3. 写 SQL:识别表 JOIN ods_sa_device 过滤省市 + uptime 时间段 + regexp 取该虫数量求 SUM。
  212. 4. 报错就修正重试(最多3次)。
  213. 5. 中文回答:给总头数 + 覆盖地区/时间/虫名 + 末尾附 SQL 代码块。
  214. ==== 注意 ====
  215. - "某月"按自然月;时间范围拿不准先 SELECT MIN/MAX(uptime) 确认。
  216. - 大表查询【必须】带时间范围过滤,避免全表扫描。
  217. - 若该虫查出来是 0 或没有:如实说明——这类设备(测报灯/性诱)主要监测趋光/趋性诱的蛾、飞虱、叶蝉等,不监测蝗虫等非趋光昆虫,不要编造数字。
  218. - Doris 与 MySQL 有差异,日期用 'YYYY-MM-DD'。
  219. """
  220. agent = Agent(
  221. name="Doris查询助手",
  222. instructions=INSTRUCTIONS,
  223. model=_model,
  224. tools=[list_tables, get_table_schema, search_pest, execute_sql],
  225. )
  226. # --- WebUI ---------------------------------------------------------------
  227. def _run(user_message: str, history: list[dict]):
  228. # 历史由 Gradio 管理(每用户每标签页独立),直接拼成 input 交给 Runner。
  229. result = asyncio.run(
  230. Runner.run(agent, history + [{"role": "user", "content": user_message}])
  231. )
  232. return result.final_output
  233. def build_ui():
  234. import gradio as gr
  235. # gradio 6 默认 history 即为 OpenAI 消息字典格式 {"role","content"}。
  236. return gr.ChatInterface(
  237. fn=_run,
  238. title="Doris 数据查询助手",
  239. description="用自然语言查询 ods 库。例:『上个月每天的订单量』",
  240. textbox=gr.Textbox(placeholder="问点什么,比如:每个地区近 7 天的销售额"),
  241. )
  242. def main():
  243. p = argparse.ArgumentParser()
  244. p.add_argument("--build-schema", action="store_true")
  245. p.add_argument("--selftest", action="store_true")
  246. args = p.parse_args()
  247. if args.build_schema:
  248. build_schema()
  249. return
  250. if args.selftest:
  251. r = asyncio.run(Runner.run(agent, "只回复 OK"))
  252. print("模型自检返回:", repr(r.final_output))
  253. return
  254. build_ui().launch(server_name="0.0.0.0", server_port=7860)
  255. if __name__ == "__main__":
  256. main()