erp-agent¶
第5章 · Coding Agent 与代码生成 · 配套项目
chapter5/erp-agent
项目说明¶
实验 5-10:自然语言交互的 ERP Agent(NL → SQL,artifact 模式)¶
把中文自然语言查询自动转成 SQL,由系统执行并直接呈现结果表。核心是 artifact(制品)模式: Agent 只负责「生成 SQL」这个制品,真正的数据查询交给数据库执行,LLM 不亲自搬运数据—— 既省 token、又避免大模型手算出错,几万行结果也能秒回。
数据模型(两张表)¶
employees:员工ID、姓名、部门、级别(数字越大越高)、入职日期、离职日期(NULL = 在职)salaries:员工ID、发薪日期(每月一条,YYYY-MM-01)、工资
数据由 seed.py 用固定随机种子(42)生成、以「今天」为基准相对生成,完全可复现:
约 40 名员工跨 5 个部门/多级别,含若干已离职者;工资按「入职基准 + 每年固定涨薪额」逐月生成,
每人涨薪额互不相同(保证问题 9 排名唯一);并刻意为一名在职员工删掉某月工资(制造问题 10 的「拖欠」)。
10 个自动回答的问题¶
- 平均每个员工在职多久 2. 每个部门有多少在职员工 3. 哪个部门平均级别最高
- 每个部门今年/去年各新入职多少人 5. 前年3月到去年5月 A 部门平均工资
- 去年 A/B 部门平均工资哪个高 7. 今年每个级别平均工资
- 入职一年内 / 一到两年 / 两到三年员工的最近一月平均工资
- 去年到今年涨薪最大的 10 位员工 10. 有没有拖欠工资(某月在职却没发薪)
其中 A 部门 = 研发部,B 部门 = 销售部(在 prompt 中约定)。
运行¶
pip install -r requirements.txt
cp env.example .env # 填入 OPENAI_API_KEY
python demo.py # 等价于 python demo.py run
通用 OpenRouter 兜底:未配置 OPENAI_API_KEY 时,设置 OPENROUTER_API_KEY 即自动
改走 OpenRouter(gpt-* → openai/*,其它 → openai/gpt-5.6-luna)。默认模型
gpt-5.6-luna 属 gpt-5.x,直连 OpenAI 需组织实名认证,故设置了 OPENROUTER_API_KEY
时会优先走 OpenRouter。
demo.py 提供 4 个子命令(不带子命令时等价于 run):
| 子命令 | 是否需要 API | 作用 |
|---|---|---|
run |
需要 | 在线:Agent 生成 SQL → 执行 → 与参考实现比对,逐题打印并给出总通过率 |
gold |
不需要 | 离线自检:执行内置「标准 SQL」(gold.py)跑 10 题并比对,证明数据模型自洽 |
ask |
需要 | 单条自然语言查询 → 生成 SQL → 执行并打印结果表 |
initdb |
不需要 | 建表并把种子数据灌入一个 SQLite 文件,便于用 sqlite3 手工查看 |
常用参数:--only 1,5,10(只跑指定题号)、--db erp.db(用文件库而非内存库)、
--model gpt-5.6-luna(覆盖模型)、--output result.json(导出逐题明细)。示例:
python demo.py gold # 离线跑通 10 题,无需 API
python demo.py run --only 2,3,6 # 只让 Agent 生成这 3 题的 SQL 并校验
python demo.py ask "研发部现在有多少在职员工?"
在线模式会:建 SQLite 内存库 → 灌种子数据 → 逐题让 Agent 生成 SQL → 执行 → 打印 「问题 / 生成的 SQL / 查询结果 / 是否通过」,最后给出总通过率。
正确性校验¶
reference.py 是独立的 Python 参考实现:不走 SQL,直接在种子数据上把每题答案算一遍。
demo.py 把 SQL 的执行结果与参考答案按「多重集合 + 数值容差」比对,逐题打印 通过/不通过。
gold.py 是人工编写的 10 条「标准 SQL」,python demo.py gold 离线执行它们即可自检,无需 API。
最近一次真实运行:离线 gold 通过率 10/10;在线 run(gpt-5.6-luna)
全部 10 题稳定通过,总通过率 10/10。
文件¶
| 文件 | 作用 |
|---|---|
demo.py |
命令行入口(run/gold/ask/initdb):建库、灌数据、跑题、执行 SQL、比对、打印通过率 |
seed.py |
可复现的种子数据生成 + 建表灌数 |
reference.py |
10 题的独立 Python 参考实现(校验基准) |
gold.py |
10 题人工编写的「标准 SQL」(SQLite 方言),供 gold 离线自检 |
questions.py |
10 个自然语言问题 + 给 Agent 的「返回列/业务口径」提示 |
agent.py |
NL→SQL Agent(OpenAI SDK,读 OPENAI_API_KEY 或 OPENROUTER_API_KEY 兜底,默认 gpt-5.6-luna) |
schema_postgres.sql |
书中 PostgreSQL 版建表 DDL(迁移到真实 Postgres 时参考) |
关于数据库¶
本项目用 SQLite(零依赖、可直接复现)。书中示例用 PostgreSQL,SQL 大体通用,
差异主要在日期函数:本项目用 SQLite 的 strftime('%Y','now')、julianday()、date('now','-1 year') 等;
迁到 PostgreSQL 时对应换成 EXTRACT(YEAR FROM now())、AGE()/日期相减、now() - interval '1 year' 等即可。
说明与注意事项¶
- Agent 的 prompt 里补充了 schema 级提示:期望返回哪些列/顺序、业务口径(在职=leave_date 为空、
今年/去年如何用
strftime(...,'now',...)推导、A/B 部门映射),以及禁止硬编码年份 (否则模型不知道「今天」是哪年,会把「前年/去年」猜错)。这些是合理的 schema 提示,不泄露具体答案。 - 问题 8、10 较复杂(工龄分档取最近一月工资、递归生成在职月份找空缺),在提示里给了推荐的 SQL 结构模板,帮助较小模型稳定产出正确 SQL。
temperature=0让输出尽量稳定,但 LLM 仍非严格确定性;若个别题偶发偏差,重跑即可, 也可换更强的模型(设OPENAI_MODEL)。
源代码¶
agent.py¶
"""
NL -> SQL Agent(artifact 模式)。
Agent 只负责「生成 SQL 制品」,不亲自搬运数据:
真正的数据查询由系统(demo.py)用生成的 SQL 在 SQLite 上执行,结果表直接呈现。
"""
import os
import re
from datetime import date
from openai import OpenAI
MODEL = os.environ.get("OPENAI_MODEL", "gpt-5.6-luna")
# --- 通用 OpenRouter 兜底 ---
OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1"
def _map_to_openrouter_model(model: str) -> str:
"""把直连模型名映射为 OpenRouter 上的 id(非可映射 id 统一兜底到当前廉价旗舰)。"""
if not model or "/" in model:
return model or "openai/gpt-5.6-luna"
m = model.lower()
if m.startswith(("gpt-", "o1", "o3", "o4")):
return "openai/" + model
if m.startswith("claude"):
if "haiku" in m:
return "anthropic/claude-haiku-4.5"
if "sonnet" in m:
return "anthropic/claude-sonnet-4.6"
return "anthropic/claude-opus-4.8"
if m.startswith("gemini"):
return "google/" + model
return "openai/gpt-5.6-luna"
def _make_client_and_model(model: str):
"""构造客户端并解析模型名,含通用 OpenRouter 兜底。返回 (client, resolved_model)。
- 有 OPENAI_API_KEY:直连;但 model 为 gpt-5.x 且同时设置了 OPENROUTER_API_KEY
时优先走 OpenRouter(直连 gpt-5.6 需组织实名认证)。
- 无 OPENAI_API_KEY 但有 OPENROUTER_API_KEY:改走 OpenRouter(模型名自动映射)。
"""
api_key = os.environ.get("OPENAI_API_KEY")
base_url = os.environ.get("OPENAI_BASE_URL")
orkey = os.environ.get("OPENROUTER_API_KEY")
prefer_or = bool(orkey) and (model or "").lower().startswith("gpt-5")
if prefer_or or (not api_key and orkey):
api_key, base_url, model = orkey, OPENROUTER_BASE_URL, _map_to_openrouter_model(model)
kw = {}
if api_key:
kw["api_key"] = api_key
if base_url:
kw["base_url"] = base_url
return OpenAI(**kw), model
SYSTEM_PROMPT = """你是一个「自然语言转 SQL」的 ERP 数据助手。
用户给你一个中文问题,你只输出一条可直接执行的 **SQLite** SQL 查询,不要任何解释、不要 markdown 代码块。
今天的日期是 {today}。但**严禁在 SQL 里硬编码年份数字**(如 '2024'、'2022-01-01'),
一律用 strftime(...,'now',...) 从数据库当前日期推导,避免年份猜错。
数据库 schema(SQLite):
employees(emp_id INTEGER 主键, name 姓名, department 部门, level 级别[数字越大越高],
hire_date 入职日期'YYYY-MM-DD', leave_date 离职日期'YYYY-MM-DD',NULL 表示在职)
salaries(emp_id, pay_date 发薪日期'YYYY-MM-01'[每月一条], salary 当月工资)
salaries.emp_id 关联 employees.emp_id。
业务与方言约定:
- 「今年」= strftime('%Y','now'),「去年」= strftime('%Y','now','-1 year'),
「前年」= strftime('%Y','now','-2 years')。
- 计算「今天」请用 date('now')(不要带时间部分);两个日期相差天数用
julianday(date('now')) - julianday(hire_date)。
- 「A部门」= 研发部,「B部门」= 销售部。
- 「在职」指 leave_date IS NULL。
- 发薪月份可用 strftime('%Y-%m', pay_date) 得到 'YYYY-MM'。
- 只输出一条 SELECT(可含 WITH/CTE),不要写多条语句或 DDL/DML。
严格按用户附带的「返回列」要求组织 SELECT 的列与顺序。
"""
class SQLAgent:
def __init__(self, model: str = MODEL):
self.client, self.model = _make_client_and_model(model)
def generate_sql(self, nl_question: str, hint: str) -> str:
user = f"问题:{nl_question}\n要求:{hint}\n请只输出一条 SQLite SQL。"
# 推理模型(gpt-5 / o 系列等)不接受 temperature=0。
_reasoning = any(k in (self.model or "").lower()
for k in ("gpt-5", "o1", "o3", "o4", "thinking", "reasoner", "kimi-k3"))
resp = self.client.chat.completions.create(
model=self.model,
temperature=1 if _reasoning else 0,
messages=[
{"role": "system",
"content": SYSTEM_PROMPT.format(today=date.today().isoformat())},
{"role": "user", "content": user},
],
)
return _clean_sql(resp.choices[0].message.content)
def _clean_sql(text: str) -> str:
"""去掉 markdown 代码块围栏等杂质,只留 SQL。"""
text = text.strip()
# 去掉 ```sql ... ``` 或 ``` ... ```
fence = re.match(r"^```(?:sql)?\s*(.*?)\s*```$", text, re.DOTALL | re.IGNORECASE)
if fence:
text = fence.group(1).strip()
# 去掉可能残留的前缀反引号
text = text.strip("`").strip()
return text
demo.py¶
"""
实验 5-10:自然语言交互的 ERP Agent(NL -> SQL,artifact 模式)命令行入口。
核心思想(artifact 模式):Agent 只负责「生成 SQL 制品」,不亲自搬运数据;
真正的查询由系统用生成的 SQL 在数据库上执行,结果表直达用户界面。
子命令:
run 在线:Agent 生成 SQL -> 执行 -> 与参考实现比对(需 OPENAI_API_KEY,默认子命令)
gold 离线:执行内置「标准 SQL」跑 10 题 -> 与参考实现比对(无需 API,用于自检/演示)
ask 在线:单条自然语言查询 -> 生成 SQL -> 执行并打印结果表(需 OPENAI_API_KEY)
initdb 建表并把可复现的种子数据灌入一个 SQLite 文件(离线,便于用 sqlite3 手工查看)
不带子命令时等价于 `run`,保持与旧版 `python demo.py` 相同的默认行为。
完整用法见 `python demo.py --help`,或某个子命令的 `python demo.py <子命令> --help`。
"""
import argparse
import json
import os
import sqlite3
import sys
from datetime import date
try:
from dotenv import load_dotenv
load_dotenv()
except Exception:
pass
import seed
import reference
import gold
from questions import QUESTIONS
from agent import SQLAgent, MODEL
# ---------------- 结果比对 ----------------
def _norm(v):
"""把单个值归一化为 ('n', 数值) 或 ('s', 字符串),便于容差比对。"""
if isinstance(v, bool):
return ("n", float(v))
if isinstance(v, (int, float)):
return ("n", round(float(v), 2))
return ("s", str(v).strip())
def _row_match(a, b, tol):
if len(a) != len(b):
return False
for x, y in zip(a, b):
if x[0] != y[0]:
return False
if x[0] == "n":
if abs(x[1] - y[1]) > tol:
return False
else:
if x[1] != y[1]:
return False
return True
def compare(expected, actual, tol=0.1):
"""按多重集合(忽略行顺序)比对期望与实际结果,数值带容差。"""
exp = [tuple(_norm(v) for v in r) for r in expected]
act = [tuple(_norm(v) for v in r) for r in actual]
if len(exp) != len(act):
return False, f"行数不一致:期望 {len(exp)} 行,实际 {len(act)} 行"
remaining = list(act)
for er in exp:
for i, ar in enumerate(remaining):
if _row_match(er, ar, tol):
remaining.pop(i)
break
else:
return False, f"缺少匹配行:{_readable(er)}"
return True, "结果一致"
def _readable(norm_row):
return tuple(v[1] for v in norm_row)
# ---------------- 结果表打印 ----------------
def print_table(rows, max_rows=12):
if not rows:
print(" (空结果)")
return
for r in rows[:max_rows]:
cells = []
for v in r:
if isinstance(v, float):
cells.append(f"{v:.2f}")
else:
cells.append(str(v))
print(" | " + " | ".join(cells) + " |")
if len(rows) > max_rows:
print(f" ... 共 {len(rows)} 行")
# ---------------- 逐题执行主循环(在线/离线共用) ----------------
def run_questions(conn, employees, salaries, today, sql_provider,
qids=None, print_sql=True, max_rows=12):
"""对每个问题:取 SQL -> 执行 -> 与 Python 参考实现比对,逐题打印。
sql_provider(q) -> str:给出该题的 SQL;可能抛异常(如在线调用 LLM 失败)。
在线模式传入 `lambda q: agent.generate_sql(q["nl"], q["hint"])`,
离线模式传入 `lambda q: gold.GOLD[q["id"]]`。
qids:只跑这些题号(None 表示全部)。
返回 (passed, total, results),results 为逐题明细 dict,便于 --output 导出。
"""
results = []
passed = 0
total = 0
for q in QUESTIONS:
if qids and q["id"] not in qids:
continue
total += 1
qid, nl, hint = q["id"], q["nl"], q["hint"]
print(f"\n【问题 {qid}】{nl}")
rec = {"id": qid, "nl": nl, "sql": None, "rows": None,
"passed": False, "error": None}
# 1) 取 SQL 制品(在线由 Agent 生成,离线取内置 gold SQL)
try:
sql = sql_provider(q)
except Exception as e:
print(f" [生成 SQL 失败] {e}")
rec["error"] = f"生成 SQL 失败:{e}"
results.append(rec)
continue
rec["sql"] = sql
if print_sql:
print(" 生成的 SQL:")
for line in sql.splitlines():
print(" " + line)
# 2) 系统执行 SQL
try:
cur = conn.cursor()
cur.execute(sql)
actual = cur.fetchall()
except Exception as e:
print(f" [SQL 执行出错] {e}")
print(" 结果:不通过 ✗")
rec["error"] = f"SQL 执行出错:{e}"
results.append(rec)
continue
rec["rows"] = [list(r) for r in actual]
print(" 查询结果:")
print_table(actual, max_rows=max_rows)
# 3) 与参考实现比对
expected = reference.REFERENCE[qid](employees, salaries, today)
ok, msg = compare(expected, actual)
rec["passed"] = ok
if ok:
passed += 1
print(f" 校验:通过 ✓({msg})")
else:
print(f" 校验:不通过 ✗({msg})")
print(f" 参考期望:{[tuple(r) for r in expected][:12]}")
results.append(rec)
return passed, total, results
# ---------------- 公用:建库、题号过滤、导出、页眉页脚 ----------------
def _build_db(db_path, today):
"""按固定种子生成数据并灌入指定的 SQLite 库(':memory:' 或文件路径)。
每次都重新灌入,保证与 reference.py 的期望答案严格对齐、结果可复现。
"""
employees, salaries = seed.generate(today)
conn = sqlite3.connect(db_path)
seed.create_db(conn, employees, salaries)
return conn, employees, salaries
def _parse_only(only):
"""把 '1,5,10' 解析成 {1,5,10};空/None 表示全部题目。"""
if not only:
return None
ids = set()
for part in only.split(","):
part = part.strip()
if part:
try:
ids.add(int(part))
except ValueError:
raise SystemExit(f"题号必须是整数:{part!r}(--only 形如 1,5,10)")
unknown = ids - {q["id"] for q in QUESTIONS}
if unknown:
raise SystemExit(f"未知题号:{sorted(unknown)}(有效题号 1~{len(QUESTIONS)})")
return ids
def _header(mode, today, employees, salaries, model=None):
print("=" * 70)
tail = f" | 模型:{model}" if model else " | 离线(不调用 API)"
print(f"ERP Agent 实验 5-10 | {mode}{tail}")
print(f"今天:{today.isoformat()} | 员工 {len(employees)} 人,"
f"工资记录 {len(salaries)} 条")
print("=" * 70)
def _footer(passed, total):
print("\n" + "=" * 70)
rate = (passed / total * 100) if total else 0
print(f"总通过率:{passed}/{total} ({rate:.0f}%)")
print("=" * 70)
def _write_output(path, mode, today, passed, total, results):
payload = {
"experiment": "5-10 ERP Agent NL->SQL",
"mode": mode,
"date": today.isoformat(),
"passed": passed,
"total": total,
"results": results,
}
with open(path, "w", encoding="utf-8") as f:
json.dump(payload, f, ensure_ascii=False, indent=2)
print(f"\n已写出结果 JSON:{path}")
def _require_api():
if not (os.environ.get("OPENAI_API_KEY") or os.environ.get("OPENROUTER_API_KEY")):
print("请先设置 OPENAI_API_KEY(或 OPENROUTER_API_KEY 兜底)环境变量(可复制 env.example 为 .env)。")
print("若只想离线跑通、不调用 API,请改用:python demo.py gold")
sys.exit(1)
# ---------------- 子命令 ----------------
def cmd_run(args):
"""在线:Agent 生成 SQL -> 执行 -> 比对。"""
_require_api()
today = date.today()
conn, emps, sals = _build_db(args.db, today)
model = args.model or os.environ.get("OPENAI_MODEL", MODEL)
_header("在线(Agent 生成 SQL)", today, emps, sals, model=model)
agent = SQLAgent(model=model)
qids = _parse_only(args.only)
passed, total, results = run_questions(
conn, emps, sals, today,
sql_provider=lambda q: agent.generate_sql(q["nl"], q["hint"]),
qids=qids, max_rows=args.max_rows,
)
_footer(passed, total)
if args.output:
_write_output(args.output, "run", today, passed, total, results)
def cmd_gold(args):
"""离线:执行内置标准 SQL -> 比对(无需 API)。"""
today = date.today()
conn, emps, sals = _build_db(args.db, today)
_header("离线自检(内置 gold SQL)", today, emps, sals, model=None)
qids = _parse_only(args.only)
passed, total, results = run_questions(
conn, emps, sals, today,
sql_provider=lambda q: gold.GOLD[q["id"]],
qids=qids, max_rows=args.max_rows,
)
_footer(passed, total)
if args.output:
_write_output(args.output, "gold", today, passed, total, results)
def cmd_ask(args):
"""在线:单条自然语言查询 -> 生成 SQL -> 执行并打印结果表。"""
_require_api()
today = date.today()
conn, emps, sals = _build_db(args.db, today)
model = args.model or os.environ.get("OPENAI_MODEL", MODEL)
agent = SQLAgent(model=model)
hint = args.hint or "自行判断需要返回的列;只输出一条 SELECT。"
print(f"【问题】{args.query}")
try:
sql = agent.generate_sql(args.query, hint)
except Exception as e:
print(f"[Agent 生成 SQL 失败] {e}")
sys.exit(1)
print("生成的 SQL:")
for line in sql.splitlines():
print(" " + line)
try:
cur = conn.cursor()
cur.execute(sql)
rows = cur.fetchall()
except Exception as e:
print(f"[SQL 执行出错] {e}")
sys.exit(1)
print("查询结果:")
print_table(rows, max_rows=args.max_rows)
def cmd_initdb(args):
"""建表并把种子数据灌入一个 SQLite 文件,便于手工用 sqlite3 查看。"""
today = date.today()
if args.db == ":memory:":
raise SystemExit("initdb 需要一个文件路径,例如:python demo.py initdb --db erp.db")
if os.path.exists(args.db):
os.remove(args.db)
conn, emps, sals = _build_db(args.db, today)
conn.close()
print(f"已写入 SQLite 库:{args.db}")
print(f" 员工 {len(emps)} 人,工资记录 {len(sals)} 条,基准日期 {today.isoformat()}")
print(f" 手工查看: sqlite3 {args.db} \"SELECT * FROM employees LIMIT 5;\"")
print(f" 离线复跑: python demo.py gold --db {args.db}")
# ---------------- argparse CLI ----------------
def build_parser():
p = argparse.ArgumentParser(
prog="demo.py",
description="实验 5-10:自然语言交互的 ERP Agent(NL -> SQL,artifact 模式)",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="不带子命令时等价于 run(保持旧版默认行为)。"
"离线自检不需要 API:python demo.py gold",
)
sub = p.add_subparsers(dest="cmd", metavar="子命令")
def add_common(sp, with_model=False):
sp.add_argument("--db", default=":memory:",
help="SQLite 库:':memory:'(默认,内存库)或文件路径")
sp.add_argument("--only", default=None, metavar="题号列表",
help="只跑指定题号,逗号分隔,如 1,5,10(默认全部)")
sp.add_argument("--max-rows", type=int, default=12, dest="max_rows",
help="每题结果表最多打印多少行(默认 12)")
sp.add_argument("--output", default=None, metavar="路径",
help="把逐题结果写成 JSON 文件")
if with_model:
sp.add_argument("--model", default=None,
help=f"覆盖模型(默认读 OPENAI_MODEL,否则 {MODEL})")
sp_run = sub.add_parser("run", help="在线:Agent 生成 SQL 跑 10 题并校验(需 API)")
add_common(sp_run, with_model=True)
sp_run.set_defaults(func=cmd_run)
sp_gold = sub.add_parser("gold", help="离线:执行内置标准 SQL 跑 10 题并校验(无需 API)")
add_common(sp_gold, with_model=False)
sp_gold.set_defaults(func=cmd_gold)
sp_ask = sub.add_parser("ask", help="在线:单条自然语言查询 -> SQL -> 结果表(需 API)")
sp_ask.add_argument("query", help="要查询的自然语言问题,如“研发部现在有多少在职员工?”")
sp_ask.add_argument("--hint", default=None, help="可选:补充业务口径/期望返回列")
sp_ask.add_argument("--db", default=":memory:",
help="SQLite 库:':memory:'(默认)或文件路径")
sp_ask.add_argument("--max-rows", type=int, default=20, dest="max_rows",
help="结果表最多打印多少行(默认 20)")
sp_ask.add_argument("--model", default=None,
help=f"覆盖模型(默认读 OPENAI_MODEL,否则 {MODEL})")
sp_ask.set_defaults(func=cmd_ask)
sp_init = sub.add_parser("initdb", help="建表并把种子数据灌入 SQLite 文件(离线)")
sp_init.add_argument("--db", default="erp.db",
help="目标 SQLite 文件路径(默认 erp.db)")
sp_init.set_defaults(func=cmd_initdb)
return p
def main(argv=None):
parser = build_parser()
args = parser.parse_args(argv)
if args.cmd is None:
# 不带子命令 -> 沿用旧版默认行为:在线跑全部 10 题
args = parser.parse_args((argv or []) + ["run"])
args.func(args)
if __name__ == "__main__":
main()
gold.py¶
"""
10 道题的「标准 SQL」(gold SQL),SQLite 方言,人工编写并逐题核对过。
用途:
- 离线演示(`python demo.py gold`):不调用任何 API,直接执行这些 SQL,
证明 schema + 种子数据这套数据模型本身是自洽、可查询的;
- 作为 Agent 生成 SQL 的「参考写法」:与 reference.py(纯 Python 参考实现)
语义一致,`demo.py` 会把执行结果与 reference.py 比对,逐题打印 通过/不通过。
约定:
- 日期一律用 date('now','localtime') / strftime(...,'now','localtime') 取「今天」,
与 seed.py 里以本地 date.today() 生成的数据对齐(避免 UTC 与本地相差一天);
- **不硬编码年份**,一律从数据库当前日期用修饰符推导('-1 year' / 'start of year' 等);
- 「A部门」= 研发部,「B部门」= 销售部;「在职」= leave_date IS NULL。
"""
GOLD = {
# 1. 平均每个员工在职多久(天)。离职用 leave_date,在职用今天。
1: """
SELECT ROUND(AVG(
julianday(COALESCE(leave_date, date('now','localtime')))
- julianday(hire_date)
), 2) AS avg_tenure_days
FROM employees;
""".strip(),
# 2. 每个部门有多少在职员工。
2: """
SELECT department, COUNT(*) AS active_count
FROM employees
WHERE leave_date IS NULL
GROUP BY department;
""".strip(),
# 3. 哪个部门(含离职)平均级别最高,只返回部门名。
3: """
SELECT department
FROM employees
GROUP BY department
ORDER BY AVG(level) DESC
LIMIT 1;
""".strip(),
# 4. 每个部门今年 / 去年各新入职多少人(按 hire_date 年份)。
4: """
SELECT department,
SUM(CASE WHEN strftime('%Y', hire_date)
= strftime('%Y','now','localtime') THEN 1 ELSE 0 END) AS this_year,
SUM(CASE WHEN strftime('%Y', hire_date)
= strftime('%Y','now','localtime','-1 year') THEN 1 ELSE 0 END) AS last_year
FROM employees
GROUP BY department
HAVING this_year > 0 OR last_year > 0;
""".strip(),
# 5. 前年3月 ~ 去年5月(含两端),研发部(A部门)平均工资。
5: """
SELECT ROUND(AVG(s.salary), 2) AS avg_salary
FROM salaries s
JOIN employees e ON e.emp_id = s.emp_id
WHERE e.department = '研发部'
AND strftime('%Y-%m', s.pay_date) BETWEEN
strftime('%Y-%m','now','localtime','start of year','-2 years','+2 months')
AND strftime('%Y-%m','now','localtime','start of year','-1 year','+4 months');
""".strip(),
# 6. 去年研发部(A)与销售部(B)平均工资,两行(含已离职员工)。
6: """
SELECT e.department, ROUND(AVG(s.salary), 2) AS avg_salary
FROM salaries s
JOIN employees e ON e.emp_id = s.emp_id
WHERE e.department IN ('研发部','销售部')
AND strftime('%Y', s.pay_date) = strftime('%Y','now','localtime','-1 year')
GROUP BY e.department;
""".strip(),
# 7. 今年每个级别的员工平均工资。
7: """
SELECT e.level, ROUND(AVG(s.salary), 2) AS avg_salary
FROM salaries s
JOIN employees e ON e.emp_id = s.emp_id
WHERE strftime('%Y', s.pay_date) = strftime('%Y','now','localtime')
GROUP BY e.level;
""".strip(),
# 8. 工龄分档(入职一年内 / 一到两年 / 两到三年,三年以上不计),各档最近一月工资的平均。
8: """
WITH latest AS ( -- 每位员工「最近一个月」的工资
SELECT s.emp_id, s.salary
FROM salaries s
JOIN (SELECT emp_id, MAX(pay_date) AS mp FROM salaries GROUP BY emp_id) m
ON m.emp_id = s.emp_id AND m.mp = s.pay_date
),
bucketed AS ( -- 给每位员工打上工龄档位
SELECT e.emp_id,
CASE
WHEN julianday(date('now','localtime')) - julianday(e.hire_date) < 365 THEN '入职一年内'
WHEN julianday(date('now','localtime')) - julianday(e.hire_date) < 730 THEN '一到两年'
WHEN julianday(date('now','localtime')) - julianday(e.hire_date) < 1095 THEN '两到三年'
ELSE NULL
END AS bucket
FROM employees e
)
SELECT b.bucket, ROUND(AVG(l.salary), 2) AS avg_salary
FROM bucketed b
JOIN latest l ON l.emp_id = b.emp_id
WHERE b.bucket IS NOT NULL
GROUP BY b.bucket;
""".strip(),
# 9. 去年到今年涨薪额(今年均薪 - 去年均薪)最大的 10 人,只算两年都有工资的。
9: """
WITH ty AS (
SELECT emp_id, AVG(salary) AS a FROM salaries
WHERE strftime('%Y', pay_date) = strftime('%Y','now','localtime') GROUP BY emp_id),
ly AS (
SELECT emp_id, AVG(salary) AS a FROM salaries
WHERE strftime('%Y', pay_date) = strftime('%Y','now','localtime','-1 year') GROUP BY emp_id)
SELECT e.name, ROUND(ty.a - ly.a, 2) AS raise_amt
FROM ty
JOIN ly ON ty.emp_id = ly.emp_id
JOIN employees e ON e.emp_id = ty.emp_id
ORDER BY raise_amt DESC
LIMIT 10;
""".strip(),
# 10. 拖欠工资:某月在职却没有发薪记录。递归展开每人的在职月份再左连接工资表。
10: """
WITH RECURSIVE em(emp_id, m, end_m) AS (
SELECT emp_id,
strftime('%Y-%m', hire_date),
COALESCE(strftime('%Y-%m', leave_date), strftime('%Y-%m','now','localtime'))
FROM employees
UNION ALL
SELECT emp_id, strftime('%Y-%m', date(m || '-01', '+1 month')), end_m
FROM em WHERE m < end_m)
SELECT em.emp_id, em.m
FROM em
LEFT JOIN salaries s
ON s.emp_id = em.emp_id AND strftime('%Y-%m', s.pay_date) = em.m
WHERE s.emp_id IS NULL;
""".strip(),
}
questions.py¶
"""
10 个自然语言问题,以及给 Agent 的「输出列」提示。
hint 里只补充「业务口径 + 期望返回哪些列、什么顺序」这类 schema 级提示,
不泄露具体数值答案。列顺序与 reference.py 的返回一致,便于逐行比对。
"""
QUESTIONS = [
{
"id": 1,
"nl": "平均每个员工在职多久?",
"hint": "在职时长按天计:离职员工用 leave_date,在职员工用今天 date('now'),"
"对全部员工求平均。只返回一列:平均在职天数。",
},
{
"id": 2,
"nl": "每个部门有多少在职员工?",
"hint": "在职指 leave_date 为空。返回两列:部门, 在职人数。",
},
{
"id": 3,
"nl": "哪个部门员工平均级别最高?",
"hint": "按所有员工(含离职)的 level 求各部门平均,取最高的那个部门。"
"只返回一列:部门名称。",
},
{
"id": 4,
"nl": "每个部门今年和去年各新入职多少人?",
"hint": "按 hire_date 的年份统计。返回三列:部门, 今年入职人数, 去年入职人数;"
"只保留今年或去年至少有一人入职的部门。",
},
{
"id": 5,
"nl": "前年3月到去年5月,A部门平均工资是多少?",
"hint": "A部门=研发部;时间范围指发薪月份从『前年3月』到『去年5月』(含两端)。"
"禁止硬编码年份,时间范围可写成:"
"strftime('%Y-%m',pay_date) BETWEEN "
"strftime('%Y-%m','now','-2 years','start of year','+2 months') AND "
"strftime('%Y-%m','now','-1 year','start of year','+4 months')。"
"只返回一列:平均工资。",
},
{
"id": 6,
"nl": "去年A部门和B部门平均工资哪个高?",
"hint": "A部门=研发部,B部门=销售部;只统计去年发薪记录"
"(strftime('%Y',pay_date)=strftime('%Y','now','-1 year'))。"
"统计部门内所有员工(含已离职),不要按 leave_date 过滤。"
"返回两列:部门, 平均工资(两行,分别对应研发部和销售部)。",
},
{
"id": 7,
"nl": "今年每个级别的员工平均工资是多少?",
"hint": "只统计今年发薪记录,按 level 分组。返回两列:级别, 平均工资。",
},
{
"id": 8,
"nl": "入职一年内、一到两年、两到三年的员工,最近一个月平均工资是多少?",
"hint": "工龄按 date('now')-hire_date 的天数分档:<365 天为『入职一年内』,"
"365~730 天为『一到两年』,730~1095 天为『两到三年』,三年以上不统计。"
"『最近一个月工资』指该员工发薪日期最大的那条工资。按档位求平均。"
"返回两列:档位(值必须正好是『入职一年内』/『一到两年』/『两到三年』), 平均工资。",
},
{
"id": 9,
"nl": "去年到今年涨薪幅度最大的10位员工是谁?",
"hint": "对每位员工,涨薪额 = 今年平均工资 - 去年平均工资,只统计去年和今年都有工资的员工,"
"按涨薪额从高到低取前 10。返回两列:姓名, 涨薪额。",
},
{
"id": 10,
"nl": "有没有拖欠工资的情况(某个月还在职却没有发薪)?",
"hint": "对每位员工,其在职月份从入职月份到(离职员工用离职月份、在职员工用当前月份),"
"逐月检查是否有对应的发薪记录,找出缺失的(员工, 月份)。"
"返回两列:emp_id, 月份(格式 YYYY-MM)。"
"推荐写法(递归 CTE 里把『结束月份』一并带进去,避免相关子查询):\n"
"WITH RECURSIVE em(emp_id, m, end_m) AS (\n"
" SELECT emp_id, strftime('%Y-%m', hire_date),\n"
" COALESCE(strftime('%Y-%m', leave_date), strftime('%Y-%m','now'))\n"
" FROM employees\n"
" UNION ALL\n"
" SELECT emp_id, strftime('%Y-%m', date(m || '-01', '+1 month')), end_m\n"
" FROM em WHERE m < end_m)\n"
"SELECT em.emp_id, em.m FROM em\n"
"LEFT JOIN salaries s ON s.emp_id = em.emp_id "
"AND strftime('%Y-%m', s.pay_date) = em.m\n"
"WHERE s.emp_id IS NULL;",
},
]
# 需要「按顺序」呈现的题目(校验时其实按集合比对内容即可,这里仅用于展示)
ORDERED = {9}
reference.py¶
"""
独立的 Python 参考实现:直接在种子数据(内存 list)上计算 10 个问题的期望答案。
这些函数刻意「不走 SQL」,用来校验 Agent 生成 SQL 的执行结果是否正确。
每个函数返回 list[tuple],元组内的列顺序与 questions.py 里给 Agent 的
「列顺序提示」保持一致,便于逐行比对。
"""
from datetime import date
from statistics import mean
DEPT_A = "研发部" # 题目里的「A 部门」
DEPT_B = "销售部" # 题目里的「B 部门」
def _end_date(e, today):
return e["leave_date"] if e["leave_date"] else today
def _ym(d: date):
return (d.year, d.month)
def _latest_salary(emp_id, salaries):
recs = [s for s in salaries if s["emp_id"] == emp_id]
if not recs:
return None
return max(recs, key=lambda s: s["pay_date"])["salary"]
def q1_avg_tenure_days(emps, sals, today):
days = [(_end_date(e, today) - e["hire_date"]).days for e in emps]
return [(round(mean(days), 2),)]
def q2_active_by_dept(emps, sals, today):
counts = {}
for e in emps:
if e["leave_date"] is None:
counts[e["department"]] = counts.get(e["department"], 0) + 1
return [(d, c) for d, c in counts.items()]
def q3_dept_highest_avg_level(emps, sals, today):
by_dept = {}
for e in emps:
by_dept.setdefault(e["department"], []).append(e["level"])
top = max(by_dept.items(), key=lambda kv: mean(kv[1]))
return [(top[0],)] # 只返回部门名称
def q4_hires_this_and_last_year(emps, sals, today):
y = today.year
agg = {}
for e in emps:
hy = e["hire_date"].year
if hy not in (y, y - 1):
continue
ty, ly = agg.get(e["department"], (0, 0))
if hy == y:
ty += 1
else:
ly += 1
agg[e["department"]] = (ty, ly)
return [(d, ty, ly) for d, (ty, ly) in agg.items()]
def q5_deptA_avg_salary_range(emps, sals, today):
y = today.year
lo, hi = (y - 2, 3), (y - 1, 5) # 前年3月 ~ 去年5月(含端点)
dept = {e["emp_id"] for e in emps if e["department"] == DEPT_A}
vals = [s["salary"] for s in sals
if s["emp_id"] in dept and lo <= _ym(s["pay_date"]) <= hi]
return [(round(mean(vals), 2),)]
def q6_deptAB_avg_salary_last_year(emps, sals, today):
y = today.year - 1
out = []
for dept in (DEPT_A, DEPT_B):
ids = {e["emp_id"] for e in emps if e["department"] == dept}
vals = [s["salary"] for s in sals
if s["emp_id"] in ids and s["pay_date"].year == y]
out.append((dept, round(mean(vals), 2)))
return out
def q7_avg_salary_by_level_this_year(emps, sals, today):
y = today.year
lvl = {e["emp_id"]: e["level"] for e in emps}
by_level = {}
for s in sals:
if s["pay_date"].year == y:
by_level.setdefault(lvl[s["emp_id"]], []).append(s["salary"])
return [(l, round(mean(v), 2)) for l, v in by_level.items()]
def q8_avg_latest_salary_by_tenure(emps, sals, today):
buckets = {"入职一年内": [], "一到两年": [], "两到三年": []}
for e in emps:
days = (today - e["hire_date"]).days
if days < 365:
b = "入职一年内"
elif days < 730:
b = "一到两年"
elif days < 1095:
b = "两到三年"
else:
continue
last = _latest_salary(e["emp_id"], sals)
if last is not None:
buckets[b].append(last)
return [(b, round(mean(v), 2)) for b, v in buckets.items() if v]
def q9_top10_raise(emps, sals, today):
y = today.year
name = {e["emp_id"]: e["name"] for e in emps}
this_year, last_year = {}, {}
for s in sals:
if s["pay_date"].year == y:
this_year.setdefault(s["emp_id"], []).append(s["salary"])
elif s["pay_date"].year == y - 1:
last_year.setdefault(s["emp_id"], []).append(s["salary"])
rows = []
for eid in set(this_year) & set(last_year):
raise_amt = mean(this_year[eid]) - mean(last_year[eid])
rows.append((name[eid], round(raise_amt, 2)))
rows.sort(key=lambda r: r[1], reverse=True)
return rows[:10]
def q10_owed_salary(emps, sals, today):
cur = (today.year, today.month)
by_emp = {}
for s in sals:
by_emp.setdefault(s["emp_id"], set()).add(_ym(s["pay_date"]))
out = []
for e in emps:
start = _ym(e["hire_date"])
end = _ym(e["leave_date"]) if e["leave_date"] else cur
paid = by_emp.get(e["emp_id"], set())
y, m = start
while (y, m) <= end:
if (y, m) not in paid:
out.append((e["emp_id"], f"{y:04d}-{m:02d}"))
m += 1
if m == 13:
y, m = y + 1, 1
return out
# 题号 -> 参考实现
REFERENCE = {
1: q1_avg_tenure_days,
2: q2_active_by_dept,
3: q3_dept_highest_avg_level,
4: q4_hires_this_and_last_year,
5: q5_deptA_avg_salary_range,
6: q6_deptAB_avg_salary_last_year,
7: q7_avg_salary_by_level_this_year,
8: q8_avg_latest_salary_by_tenure,
9: q9_top10_raise,
10: q10_owed_salary,
}
seed.py¶
"""
生成可复现的 ERP 种子数据(员工表 + 工资表)。
设计要点(保证 10 个问题都有确定答案):
- 约 40 名员工,跨 5 个部门、多个级别;
- 工龄刻意覆盖「入职一年内 / 一到两年 / 两到三年 / 三年以上」四档(供问题 8);
- 若干已离职员工(leave_date 非空,供问题 2/6 等);
- 工资按「入职当年基准 + 每年固定涨薪额」逐月生成,
每位员工的年度涨薪额互不相同,从而问题 9「涨薪最大 10 人」排名唯一;
- 刻意为一名在职员工删掉某个月的工资记录(供问题 10「拖欠工资」);
- 所有日期以「今天」为基准相对生成,固定随机种子 42,可复现。
reference.py 直接在这些内存结构上计算期望答案,
与 Agent 生成 SQL 的执行结果比对,保证语义一致。
"""
import random
from datetime import date, timedelta
# ---- 业务常量 ----
DEPARTMENTS = ["研发部", "销售部", "市场部", "财务部", "人力资源部"]
# 各部门基准工资
DEPT_BASE = {"研发部": 15000, "销售部": 12000, "市场部": 11000, "财务部": 12000, "人力资源部": 10000}
_SURNAMES = list("赵钱孙李周吴郑王冯陈褚卫蒋沈韩杨朱秦尤许何吕施张孔曹严华金魏陶姜")
_GIVEN = list("伟芳娜秀英敏静丽强磊军洋勇艳杰娟涛明超霞平刚桂香建华志强晓东春梅国栋雪松")
def _first_of_month(d: date) -> date:
return date(d.year, d.month, 1)
def _add_month(d: date) -> date:
"""返回下个月的 1 号(d 需为某月 1 号)。"""
if d.month == 12:
return date(d.year + 1, 1, 1)
return date(d.year, d.month + 1, 1)
def _month_key(d: date) -> str:
return f"{d.year:04d}-{d.month:02d}"
def generate(today: date | None = None):
"""生成并返回 (employees, salaries) 两个 list[dict]。"""
if today is None:
today = date.today()
rng = random.Random(42)
cur_month = _first_of_month(today)
# 工龄分档(相对今天的天数区间),保证问题 8 各档都有人
# bucket: (人数, 最小天数, 最大天数)
tenure_plan = [
(8, 30, 360), # 入职一年内
(8, 370, 720), # 一到两年
(8, 740, 1080), # 两到三年
(16, 1100, 1900), # 三年以上
]
employees = []
emp_id = 0
for count, dmin, dmax in tenure_plan:
for _ in range(count):
emp_id += 1
days = rng.randint(dmin, dmax)
hire_date = today - timedelta(days=days)
dept = rng.choice(DEPARTMENTS)
level = rng.randint(3, 9)
name = rng.choice(_SURNAMES) + rng.choice(_GIVEN)
employees.append({
"emp_id": emp_id,
"name": name,
"department": dept,
"level": level,
"hire_date": hire_date,
"leave_date": None, # 先全部在职,稍后挑一部分离职
})
# 挑选约 6 名「三年以上」员工设为离职,离职日期落在过去 ~2 年内
senior = [e for e in employees if (today - e["hire_date"]).days > 1100]
for e in rng.sample(senior, 6):
# 离职日期 = 今天前 60~700 天,且晚于入职至少 200 天
leave = today - timedelta(days=rng.randint(60, 700))
if (leave - e["hire_date"]).days < 200:
leave = e["hire_date"] + timedelta(days=200)
# 落到月末,便于按月发薪对齐
e["leave_date"] = leave
# 为每位员工设定「入职当年基准工资」与「每年固定涨薪额」(互不相同)
for e in employees:
e["_start_base"] = (
DEPT_BASE[e["department"]] + e["level"] * 2000 + (e["emp_id"] % 7) * 300
)
# 涨薪额随 emp_id 严格递增,保证问题 9 排名唯一(无并列)
e["_annual_raise"] = 400 + e["emp_id"] * 45
# 指定一名在职、工龄较长的员工为「明显涨薪王」(问题 9 榜首)
big_raiser = next(
e for e in employees
if e["leave_date"] is None and (today - e["hire_date"]).days > 700
)
big_raiser["_annual_raise"] = 12000 # 远高于其他人
big_raiser["_is_big_raiser"] = True
# ---- 逐月生成工资 ----
salaries = []
for e in employees:
hire_m = _first_of_month(e["hire_date"])
end_m = _first_of_month(e["leave_date"]) if e["leave_date"] else cur_month
m = hire_m
while m <= end_m:
salary = e["_start_base"] + e["_annual_raise"] * (m.year - e["hire_date"].year)
salaries.append({
"emp_id": e["emp_id"],
"pay_date": m, # 每月 1 号代表当月发薪
"salary": int(salary),
})
m = _add_month(m)
# ---- 刻意制造一条「拖欠工资」:某在职员工某个过去月份缺发薪(问题 10)----
target = next(
e for e in employees
if e["leave_date"] is None and (today - e["hire_date"]).days > 800
)
# 删除「6 个月前」那条记录(确保它存在且不是当月)
missing_month = cur_month
for _ in range(6):
# 往前推 6 个月
y, mo = missing_month.year, missing_month.month - 1
if mo == 0:
y, mo = y - 1, 12
missing_month = date(y, mo, 1)
before = len(salaries)
salaries = [
s for s in salaries
if not (s["emp_id"] == target["emp_id"] and s["pay_date"] == missing_month)
]
assert len(salaries) == before - 1, "未能删除目标发薪记录,请检查种子逻辑"
target["_owed_month"] = _month_key(missing_month)
return employees, salaries
def create_db(conn, employees, salaries):
"""在给定 sqlite 连接上建表并灌入数据。"""
cur = conn.cursor()
cur.executescript(
"""
DROP TABLE IF EXISTS employees;
DROP TABLE IF EXISTS salaries;
CREATE TABLE employees (
emp_id INTEGER PRIMARY KEY,
name TEXT NOT NULL,
department TEXT NOT NULL,
level INTEGER NOT NULL, -- 级别,数字越大越高
hire_date TEXT NOT NULL, -- 入职日期 YYYY-MM-DD
leave_date TEXT -- 离职日期,NULL = 在职
);
CREATE TABLE salaries (
emp_id INTEGER NOT NULL, -- 关联 employees.emp_id
pay_date TEXT NOT NULL, -- 发薪日期 YYYY-MM-01(每月一条)
salary INTEGER NOT NULL, -- 当月工资
PRIMARY KEY (emp_id, pay_date)
);
"""
)
cur.executemany(
"INSERT INTO employees VALUES (?,?,?,?,?,?)",
[
(
e["emp_id"], e["name"], e["department"], e["level"],
e["hire_date"].isoformat(),
e["leave_date"].isoformat() if e["leave_date"] else None,
)
for e in employees
],
)
cur.executemany(
"INSERT INTO salaries VALUES (?,?,?)",
[(s["emp_id"], s["pay_date"].isoformat(), s["salary"]) for s in salaries],
)
conn.commit()
test_parse_only.py¶
"""--only 参数解析:非法题号应干净退出(SystemExit),而非 ValueError 栈。"""
import pytest
from demo import _parse_only
def test_parse_only_valid():
assert _parse_only("1,5,10") == {1, 5, 10}
def test_parse_only_empty_means_all():
assert _parse_only("") is None
assert _parse_only(None) is None
def test_parse_only_non_integer_clean_exit():
with pytest.raises(SystemExit, match="整数"):
_parse_only("1,2x")
def test_parse_only_unknown_id_clean_exit():
with pytest.raises(SystemExit, match="未知题号"):
_parse_only("999")
schema_postgres.sql¶
-- 实验 5-10 ERP Agent —— 书中要求的 PostgreSQL schema(两张表)。
--
-- 本仓库的可运行演示用 SQLite(零依赖、可离线复现,见 seed.py / demo.py);
-- 这份 DDL 给出书中原文的 PostgreSQL 版本,方便迁移到真实 Postgres 环境。
-- 两种方言的表结构一致,差异主要在日期函数:
-- SQLite: strftime('%Y','now') julianday(a)-julianday(b) date('now','-1 year')
-- PostgreSQL: EXTRACT(YEAR FROM now()) (a::date - b::date) now() - interval '1 year'
--
-- 用法(需本机有 PostgreSQL):
-- createdb erp
-- psql erp -f schema_postgres.sql
DROP TABLE IF EXISTS salaries;
DROP TABLE IF EXISTS employees;
-- 员工表:ID、姓名、部门、级别(数字越大越高)、入职日期、离职日期(NULL = 在职)
CREATE TABLE employees (
emp_id INTEGER PRIMARY KEY,
name TEXT NOT NULL,
department TEXT NOT NULL,
level INTEGER NOT NULL,
hire_date DATE NOT NULL,
leave_date DATE -- NULL 表示在职
);
-- 工资表:员工ID、发薪日期(每月一条,取当月 1 号)、当月工资
CREATE TABLE salaries (
emp_id INTEGER NOT NULL REFERENCES employees(emp_id),
pay_date DATE NOT NULL, -- 每月一条,如 2025-03-01
salary INTEGER NOT NULL,
PRIMARY KEY (emp_id, pay_date)
);
CREATE INDEX idx_salaries_pay_date ON salaries (pay_date);