183 lines
5.7 KiB
Python
183 lines
5.7 KiB
Python
"""
|
|
Parse 常问问题对应SQL_50个问题知识库_Workbuddy生成.md into JSONL aligned with all_samples.jsonl.
|
|
Merges schema_info, tags, difficulty from all_samples when qid matches.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import re
|
|
from pathlib import Path
|
|
|
|
ROOT = Path(__file__).resolve().parents[1]
|
|
MD_PATH = ROOT / "data" / "常问问题对应SQL_50个问题知识库_Workbuddy生成.md"
|
|
REF_JSONL = ROOT / "data" / "experiences" / "all_samples.jsonl"
|
|
OUT_PATH = ROOT / "data" / "experiences" / "workbuddy_50_questions.jsonl"
|
|
|
|
# `question_en` right after `question_zh` so long lines are easy to spot in editors.
|
|
JSONL_KEY_ORDER = [
|
|
"qid",
|
|
"question_zh",
|
|
"question_en",
|
|
"sql",
|
|
"schema_info",
|
|
"explanation",
|
|
"rating",
|
|
"execution_result",
|
|
"result_note",
|
|
"tags",
|
|
"difficulty",
|
|
]
|
|
|
|
CJK_RE = re.compile(r"[\u4e00-\u9fff\u3400-\u4dbf]")
|
|
SECTION_RE = re.compile(r"^## (Q\d+)\.\s*.*$", re.MULTILINE)
|
|
SQL_BLOCK_RE = re.compile(r"```sql\s*\n(.*?)```", re.DOTALL)
|
|
RATING_RE = re.compile(r"\*\*评分:\s*(\d+)/10")
|
|
|
|
|
|
def has_cjk(s: str) -> bool:
|
|
return bool(CJK_RE.search(s))
|
|
|
|
|
|
def parse_questions(text: str) -> tuple[str, str]:
|
|
"""Split body after '### 问题' into English vs Chinese lines."""
|
|
lines = [ln.strip() for ln in text.splitlines() if ln.strip()]
|
|
en_parts: list[str] = []
|
|
zh_parts: list[str] = []
|
|
for ln in lines:
|
|
if has_cjk(ln):
|
|
zh_parts.append(ln)
|
|
else:
|
|
en_parts.append(ln)
|
|
return (" ".join(en_parts).strip(), " ".join(zh_parts).strip())
|
|
|
|
|
|
def dedupe_sql_blocks(blocks: list[str]) -> list[str]:
|
|
out: list[str] = []
|
|
seen: set[str] = set()
|
|
for b in blocks:
|
|
key = b.strip()
|
|
if key in seen:
|
|
continue
|
|
seen.add(key)
|
|
out.append(b.rstrip())
|
|
return out
|
|
|
|
|
|
def extract_sql_blocks(region: str) -> list[str]:
|
|
blocks = [m.group(1).strip() for m in SQL_BLOCK_RE.finditer(region)]
|
|
return dedupe_sql_blocks(blocks)
|
|
|
|
|
|
def reorder_row(d: dict) -> dict:
|
|
return {k: d[k] for k in JSONL_KEY_ORDER}
|
|
|
|
|
|
def parse_block(qid: str, body: str, ref: dict) -> dict:
|
|
ap = body.find("\n## 附录")
|
|
if ap != -1:
|
|
body = body[:ap].rstrip()
|
|
|
|
# --- 问题 ---
|
|
m_q = re.search(r"### 问题\s*\n(.*?)(?=\n### 查询SQL\s)", body, re.DOTALL)
|
|
if not m_q:
|
|
raise ValueError(f"{qid}: missing ### 问题")
|
|
question_en, question_zh = parse_questions(m_q.group(1))
|
|
if not question_en and question_zh:
|
|
if qid == "Q50":
|
|
question_en = "List all client trade records for today."
|
|
else:
|
|
question_en = ref.get("question_en") or ""
|
|
|
|
# --- SQL + result_note ---
|
|
m_sql_hdr = re.search(r"### 查询SQL\s*\n", body)
|
|
if not m_sql_hdr:
|
|
raise ValueError(f"{qid}: missing ### 查询SQL")
|
|
tail = body[m_sql_hdr.end() :]
|
|
idx_exec = tail.find("### SQL执行结果")
|
|
idx_reason = tail.find("### SQL理由")
|
|
|
|
if idx_exec != -1 and (idx_reason == -1 or idx_exec < idx_reason):
|
|
sql_region = tail[:idx_exec]
|
|
chunk = tail[idx_exec:idx_reason] if idx_reason != -1 else tail[idx_exec:]
|
|
result_note = re.sub(
|
|
r"^### SQL执行结果\s*",
|
|
"",
|
|
chunk,
|
|
count=1,
|
|
flags=re.MULTILINE,
|
|
).strip()
|
|
elif idx_reason != -1:
|
|
sql_region = tail[:idx_reason]
|
|
last_end = 0
|
|
for m in SQL_BLOCK_RE.finditer(sql_region):
|
|
last_end = m.end()
|
|
result_note = sql_region[last_end:].strip() if last_end else ""
|
|
else:
|
|
raise ValueError(f"{qid}: no ### SQL理由 / ### SQL执行结果")
|
|
|
|
sql_parts = extract_sql_blocks(sql_region)
|
|
sql = "\n\n".join(sql_parts) if sql_parts else ref.get("sql", "")
|
|
|
|
# --- explanation & rating ---
|
|
m_reason = re.search(r"### SQL理由\s*\n(.*?)(?=\n### SQL准确度评分\s|\Z)", body, re.DOTALL)
|
|
explanation = ""
|
|
if m_reason:
|
|
explanation = "\n".join(
|
|
ln.rstrip() for ln in m_reason.group(1).splitlines() if ln.strip().startswith("-")
|
|
).strip()
|
|
|
|
m_rating = RATING_RE.search(body)
|
|
rating = int(m_rating.group(1)) if m_rating else ref.get("rating")
|
|
|
|
base = dict(ref)
|
|
base.update(
|
|
{
|
|
"qid": qid,
|
|
"question_zh": question_zh or ref.get("question_zh") or "",
|
|
"question_en": question_en,
|
|
"sql": sql,
|
|
"explanation": explanation if explanation else (ref.get("explanation") or ""),
|
|
"rating": rating,
|
|
"execution_result": None,
|
|
"result_note": result_note.strip() if result_note else (ref.get("result_note") or ""),
|
|
}
|
|
)
|
|
return reorder_row(base)
|
|
|
|
|
|
def main() -> None:
|
|
md = MD_PATH.read_text(encoding="utf-8")
|
|
refs: dict[str, dict] = {}
|
|
with REF_JSONL.open(encoding="utf-8") as f:
|
|
for line in f:
|
|
line = line.strip()
|
|
if not line:
|
|
continue
|
|
obj = json.loads(line)
|
|
refs[obj["qid"]] = obj
|
|
|
|
matches = list(SECTION_RE.finditer(md))
|
|
rows: list[dict] = []
|
|
for i, m in enumerate(matches):
|
|
qid = m.group(1)
|
|
start = m.end()
|
|
end = matches[i + 1].start() if i + 1 < len(matches) else len(md)
|
|
body = md[start:end]
|
|
ref = refs.get(qid, {})
|
|
row = parse_block(qid, body, ref)
|
|
rows.append(row)
|
|
|
|
if len(rows) != 50:
|
|
raise SystemExit(f"expected 50 sections, got {len(rows)}")
|
|
|
|
OUT_PATH.parent.mkdir(parents=True, exist_ok=True)
|
|
with OUT_PATH.open("w", encoding="utf-8", newline="\n") as out:
|
|
for row in rows:
|
|
out.write(json.dumps(row, ensure_ascii=False) + "\n")
|
|
|
|
print(f"Wrote {OUT_PATH} ({len(rows)} lines)")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|