Files
ai-g3sb-backman2.0/scripts/md_workbuddy_to_jsonl.py

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