Enhance dialog context handling for Text2SQL queries by integrating session history. Introduce new methods for summarizing previous assistant messages and determining if the last interaction was a data query. Update environment configuration for embedding options and improve error handling in SQL generation. Add user-facing delivery messages for successful SQL execution. This update supports more coherent follow-up questions and improves user experience in conversational interactions.

This commit is contained in:
陈辅元
2026-04-14 18:02:12 +08:00
parent 026d3bf88c
commit ca8bc5e7de
87 changed files with 5808 additions and 6351 deletions
+100
View File
@@ -0,0 +1,100 @@
#!/usr/bin/env python3
"""
将 Few-shot 经验数据集(JSONL)写入 Chroma 持久化向量库。
用法(仓库根目录):
python scripts/build_fewshot_chroma_index.py
python scripts/build_fewshot_chroma_index.py --samples data/experiences/all_samples.jsonl --force
构建完成后在 .env 中设置:
FEWSHOT_USE_CHROMA=true
FEWSHOT_CHROMA_PATH=./data/embeddings/chroma_fewshot
Embedding 与 Schema 向量一致,由 USE_LOCAL_EMBEDDING / OPENAI_* / MODELSCOPE_* 等决定。
"""
from __future__ import annotations
import argparse
import json
import logging
import os
import sys
from pathlib import Path
_REPO_ROOT = Path(__file__).resolve().parents[1]
_BACKEND = _REPO_ROOT / "backend"
if str(_BACKEND) not in sys.path:
sys.path.insert(0, str(_BACKEND))
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
logger = logging.getLogger("build_fewshot_chroma")
def main() -> int:
parser = argparse.ArgumentParser(description="Few-shot JSONL → Chroma 向量索引")
parser.add_argument(
"--samples",
default=os.getenv("FEWSHOT_DATA_PATH", "data/experiences/all_samples.jsonl"),
help="JSONL 路径(相对仓库根或绝对路径)",
)
parser.add_argument(
"--persist-dir",
default=os.getenv("FEWSHOT_CHROMA_PATH", "data/embeddings/chroma_fewshot"),
help="Chroma 持久化目录",
)
parser.add_argument(
"--force",
action="store_true",
help="清空已有集合并全量重建",
)
args = parser.parse_args()
samples_path = Path(args.samples)
if not samples_path.is_absolute():
samples_path = _REPO_ROOT / samples_path
if not samples_path.is_file():
logger.error("样本文件不存在: %s", samples_path)
return 1
from dotenv import load_dotenv
load_dotenv(_REPO_ROOT / ".env")
from utils.embedding import get_embedder
from utils.fewshot_chroma_store import FewShotChromaStore
from utils.fewshot_selector import ExperienceSample
rows: list[ExperienceSample] = []
with samples_path.open("r", encoding="utf-8") as f:
for line_no, line in enumerate(f, 1):
line = line.strip()
if not line:
continue
try:
rows.append(ExperienceSample.from_dict(json.loads(line)))
except json.JSONDecodeError as e:
logger.warning("跳过第 %s 行 JSON 错误: %s", line_no, e)
if not rows:
logger.error("未解析到任何样本")
return 1
embed_path = os.getenv("EMBEDDING_MODEL_PATH", "").strip() or None
embedder = get_embedder(embed_path)
persist = Path(args.persist_dir)
if not persist.is_absolute():
persist = _REPO_ROOT / persist
store = FewShotChromaStore(embedder, persist_dir=str(persist))
n = store.build_from_samples(rows, force_rebuild=args.force)
logger.info("完成: 写入 %s 条 → %s", n, persist)
return 0
if __name__ == "__main__":
raise SystemExit(main())
+182
View File
@@ -0,0 +1,182 @@
"""
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()