first commit

This commit is contained in:
陈辅元
2026-04-10 16:52:07 +08:00
commit 84fe545640
87 changed files with 20842 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
# utils 包初始化
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+530
View File
@@ -0,0 +1,530 @@
"""
Embedding 封装:本地 Qwen3-Embedding,或兼容 OpenAI /v1/embeddings 的远程 API
(ModelScope 推理、阿里云 DashScope 等)。
"""
import os
from typing import Union, List, Optional, Any
import numpy as np
from pathlib import Path
try:
from transformers import AutoModel, AutoTokenizer
import torch
_TRANSFORMERS_AVAILABLE = True
except ImportError:
_TRANSFORMERS_AVAILABLE = False
import logging
from dataclasses import dataclass
logger = logging.getLogger(__name__)
@dataclass
class _RemoteEmbeddingEnv:
"""从环境变量解析出的远程 OpenAI-Compatible Embedding 配置。"""
api_key: str
base_url: str
model: str
max_batch: int
label: str
def _remote_embedding_from_env() -> _RemoteEmbeddingEnv:
"""
优先 ModelScope(MODELSCOPE_*);未配置时回退 DashScope(DASHSCOPE_*)。
"""
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
ms_base = os.getenv("MODELSCOPE_BASE_URL", "").strip()
ds_key = os.getenv("DASHSCOPE_API_KEY", "").strip()
ds_base = (
os.getenv("DASHSCOPE_BASE_URL") or os.getenv("DASHSCOPE_base_url", "")
).strip()
# 仅当配置了 API Key 时走 ModelScope(避免仅有 BASE_URL 时误判、阻断 DashScope)
if ms_key:
base_url = ms_base or "https://api-inference.modelscope.cn/v1"
model = (
os.getenv("MODELSCOPE_EMBEDDING_MODEL")
or os.getenv("MODELSCOPE_MODEL", "Qwen/Qwen3-Embedding-8B")
).strip()
mb = os.getenv("MODELSCOPE_EMBEDDING_MAX_BATCH", "32").strip()
max_batch = max(1, int(mb)) if mb.isdigit() else 32
if not model:
raise ValueError("未配置 MODELSCOPE_EMBEDDING_MODEL(或 MODELSCOPE_MODEL)")
return _RemoteEmbeddingEnv(
api_key=ms_key,
base_url=base_url.rstrip("/"),
model=model,
max_batch=max_batch,
label="ModelScope",
)
# OpenAI 官方或兼容网关:OPENAI_API_KEY、OPENAI_EMBEDDING_MODEL、可选 OPENAI_BASE_URL
oa_key = os.getenv("OPENAI_API_KEY", "").strip()
if oa_key:
if oa_key.startswith("http://") or oa_key.startswith("https://"):
raise ValueError(
"OPENAI_API_KEY 不能填写为 URL:请将网关地址写到 OPENAI_BASE_URL"
"(例如 http://host:9080/v1),密钥单独写在 OPENAI_API_KEY"
)
oa_base = (
os.getenv("OPENAI_BASE_URL", "").strip() or "https://api.openai.com/v1"
)
oa_model = os.getenv("OPENAI_EMBEDDING_MODEL", "").strip()
if not oa_model:
raise ValueError(
"使用 OpenAI 兼容 Embedding 时请设置 OPENAI_EMBEDDING_MODEL"
)
mb = os.getenv("OPENAI_EMBEDDING_MAX_BATCH", "100").strip()
max_batch = max(1, int(mb)) if mb.isdigit() else 100
return _RemoteEmbeddingEnv(
api_key=oa_key,
base_url=oa_base.rstrip("/"),
model=oa_model,
max_batch=max_batch,
label="OpenAI",
)
if not ds_key:
raise ValueError(
"远程 Embedding 未配置:请设置 MODELSCOPE_API_KEY(及可选 BASE_URL),"
"或 OPENAI_API_KEY / OPENAI_EMBEDDING_MODEL(及可选 OPENAI_BASE_URL),"
"或 DASHSCOPE_API_KEY / DASHSCOPE_BASE_URL / DASHSCOPE_MODEL"
)
if not ds_base:
raise ValueError("未配置 DASHSCOPE_BASE_URL(或 DASHSCOPE_base_url)")
model = os.getenv("DASHSCOPE_MODEL", "").strip()
if not model:
raise ValueError("未配置 DASHSCOPE_MODEL")
mb = os.getenv("DASHSCOPE_EMBEDDING_MAX_BATCH", "10").strip()
max_batch = max(1, int(mb)) if mb.isdigit() else 10
return _RemoteEmbeddingEnv(
api_key=ds_key,
base_url=ds_base.rstrip("/"),
model=model,
max_batch=max_batch,
label="DashScope",
)
class Qwen3Embedding:
"""
Qwen3-Embedding-0.6B 向量化封装
使用 Mean Pooling 将token embeddings聚合为句子向量,
并进行L2归一化以支持余弦相似度计算。
"""
def __init__(
self,
model_path: Optional[str] = None,
device: Optional[str] = None,
use_fp16: bool = False
):
"""
初始化 embedding 模型
Args:
model_path: 本地模型路径,若为None则从环境变量或默认路径加载
device: 推理设备('cpu', 'cuda', 'cuda:0'等),None则自动选择
use_fp16: 是否使用FP16混合精度(GPU可用时建议开启,速度更快)
"""
if not _TRANSFORMERS_AVAILABLE:
raise ImportError(
"transformers 和 torch 未安装。请运行:\n"
"pip install transformers torch sentencepiece accelerate"
)
# 确定模型路径
if model_path is None:
model_path = os.getenv(
"EMBEDDING_MODEL_PATH",
"./data/models/Qwen3-Embedding-0.6B"
)
model_path = Path(model_path)
if not model_path.exists():
raise FileNotFoundError(
f"模型目录不存在:{model_path}\n"
"请先下载模型:\n"
" modelscope download --model 'Qwen/Qwen3-Embedding-0.6B' "
f"--local_dir '{model_path}'\n"
"或从Hugging Face下载:git lfs install && git clone "
f"https://huggingface.co/Qwen/Qwen3-Embedding-0.6B {model_path}"
)
# 确定设备
if device is None:
device = "cuda" if torch.cuda.is_available() else "cpu"
self.device = device
logger.info(f"加载Qwen3-Embedding模型:{model_path},设备:{device}")
# 加载 tokenizer:fast(Rust) 解析 tokenizer.json 需较新 tokenizers;
# 旧版本会报 ModelWrapper / untagged enum,回退到慢速 tokenizer 可恢复。
try:
self.tokenizer = AutoTokenizer.from_pretrained(
str(model_path), trust_remote_code=True
)
except Exception as e:
err = str(e).lower()
if "modelwrapper" in err or "untagged enum" in err:
logger.warning(
"快速 tokenizer 解析 tokenizer.json 失败(多为 tokenizers 过旧),"
"改用 use_fast=False:%s",
e,
)
self.tokenizer = AutoTokenizer.from_pretrained(
str(model_path), use_fast=False, trust_remote_code=True
)
else:
raise
try:
self.model = AutoModel.from_pretrained(
str(model_path), trust_remote_code=True
)
except ValueError as e:
msg = str(e)
if "qwen3" in msg.lower() or "does not recognize this architecture" in msg:
raise RuntimeError(
"当前 transformers 版本不支持 Qwen3(model_type=qwen3)。"
"请升级:pip install \"transformers>=4.51.0\" \"tokenizers>=0.21\""
) from e
raise
# 设置为评估模式并移动设备
self.model.eval()
self.model.to(device)
# 混合精度(仅GPU)
self.use_fp16 = use_fp16 and device != "cpu"
if self.use_fp16:
self.model.half()
# 嵌入维度
self.embedding_dim = self.model.config.hidden_size
logger.info(f"[OK] 模型加载完成,嵌入维度:{self.embedding_dim}")
def encode(
self,
texts: Union[str, List[str]],
batch_size: int = 32,
normalize: bool = True,
max_length: int = 8192,
show_progress: bool = False
) -> np.ndarray:
"""
编码文本为向量
Args:
texts: 单个文本或文本列表
batch_size: 批处理大小(根据显存调整)
normalize: 是否L2归一化(余弦相似度必需)
max_length: 最大序列长度(模型支持8192,建议512-1024平衡速度与精度)
show_progress: 是否显示进度条(需安装tqdm)
Returns:
numpy数组,shape=(len(texts), embedding_dim)
"""
if isinstance(texts, str):
texts = [texts]
if not texts:
return np.empty((0, self.embedding_dim), dtype=np.float32)
all_embeddings = []
# 可选进度条
iterator = range(0, len(texts), batch_size)
if show_progress:
try:
from tqdm import tqdm
iterator = tqdm(iterator, desc="Embedding")
except ImportError:
pass
for i in iterator:
batch = texts[i:i + batch_size]
# Tokenize
inputs = self.tokenizer(
batch,
padding=True,
truncation=True,
max_length=max_length,
return_tensors="pt"
).to(self.device)
# Inference
with torch.no_grad():
outputs = self.model(**inputs)
# Mean Pooling: 取序列维度的平均值
# outputs.last_hidden_state shape: (batch, seq_len, hidden_size)
embeddings = outputs.last_hidden_state.mean(dim=1)
# 转换为numpy(保持在CPU)
if self.device != "cpu":
embeddings = embeddings.cpu()
embeddings = embeddings.numpy()
if normalize:
# L2归一化(余弦相似度必需)
norms = np.linalg.norm(embeddings, axis=1, keepdims=True)
embeddings = embeddings / (norms + 1e-10)
all_embeddings.append(embeddings)
return np.vstack(all_embeddings).astype(np.float32)
def similarity(
self,
emb1: np.ndarray,
emb2: np.ndarray
) -> np.ndarray:
"""
计算两组embedding的余弦相似度
Args:
emb1: 第一组向量 (n, dim)
emb2: 第二组向量 (m, dim)
Returns:
相似度矩阵 (n, m),值域[-1, 1](若已归一化则为[0, 1])
"""
# 确保已归一化
return np.dot(emb1, emb2.T)
def encode_and_search(
self,
query: str,
documents: List[str],
top_k: int = 5
) -> List[dict]:
"""
便捷方法:编码查询并检索最相似的文档
Args:
query: 查询文本
documents: 候选文档列表
top_k: 返回前K个结果
Returns:
[{"score": float, "document": str, "index": int}, ...]
"""
query_emb = self.encode([query], normalize=True)
doc_embs = self.encode(documents, normalize=True)
scores = self.similarity(query_emb, doc_embs)[0]
# 获取top_k
top_indices = np.argsort(scores)[::-1][:top_k]
results = []
for idx in top_indices:
results.append({
"score": float(scores[idx]),
"document": documents[idx],
"index": int(idx)
})
return results
def _env_flag(name: str, default: str = "true") -> bool:
return os.getenv(name, default).strip().lower() in ("1", "true", "yes", "on")
class OpenAICompatibleRemoteEmbedding:
"""
通过 OpenAI 兼容接口获取文本向量(POST /v1/embeddings)。
环境变量(优先级:ModelScope → OpenAI 兼容 → DashScope):
- ModelScope:MODELSCOPE_API_KEY、可选 MODELSCOPE_BASE_URL(默认
https://api-inference.modelscope.cn/v1)、MODELSCOPE_EMBEDDING_MODEL
或 MODELSCOPE_MODEL、可选 MODELSCOPE_EMBEDDING_MAX_BATCH
- OpenAI 兼容:OPENAI_API_KEY、OPENAI_EMBEDDING_MODEL、可选 OPENAI_BASE_URL
(默认 https://api.openai.com/v1)、可选 OPENAI_EMBEDDING_MAX_BATCH
- DashScope:DASHSCOPE_API_KEY、DASHSCOPE_BASE_URL、DASHSCOPE_MODEL、
可选 DASHSCOPE_EMBEDDING_MAX_BATCH
可选 VECTOR_DIM:在首次请求前确定空列表返回的维度。
"""
def __init__(
self,
api_key: Optional[str] = None,
base_url: Optional[str] = None,
model: Optional[str] = None,
max_batch: Optional[int] = None,
provider_label: Optional[str] = None,
):
try:
from openai import OpenAI
except ImportError as e:
raise ImportError(
"使用远程 Embedding 需要安装 openai:pip install openai"
) from e
if api_key is not None and base_url is not None and model is not None:
cfg = _RemoteEmbeddingEnv(
api_key=api_key.strip(),
base_url=base_url.strip().rstrip("/"),
model=model.strip(),
max_batch=max(1, int(max_batch)) if max_batch is not None else 32,
label=provider_label or "custom",
)
else:
cfg = _remote_embedding_from_env()
self.api_key = cfg.api_key
self.base_url = cfg.base_url
self.model = cfg.model
self._api_max_batch = cfg.max_batch
self._provider_label = cfg.label
self._client = OpenAI(api_key=self.api_key, base_url=self.base_url)
vd = os.getenv("VECTOR_DIM", "").strip()
self._embedding_dim: Optional[int] = int(vd) if vd.isdigit() else None
logger.info(
"使用 %s Embedding API:model=%s,base_url=%s,max_batch=%s",
self._provider_label,
self.model,
self.base_url,
self._api_max_batch,
)
@property
def embedding_dim(self) -> int:
if self._embedding_dim is None:
raise RuntimeError(
"尚未获知向量维度:请先执行一次 encode,或在 .env 中设置 VECTOR_DIM"
)
return self._embedding_dim
def _set_dim_from_vector(self, vec: List[float]) -> None:
if self._embedding_dim is None:
self._embedding_dim = len(vec)
logger.info("[OK] Embedding 向量维度:%s", self._embedding_dim)
def encode(
self,
texts: Union[str, List[str]],
batch_size: int = 10,
normalize: bool = True,
max_length: int = 8192,
show_progress: bool = False,
) -> np.ndarray:
del max_length # API 侧截断,此处仅保持签名与本地实现一致
if isinstance(texts, str):
texts = [texts]
if not texts:
dim = self._embedding_dim
if dim is None:
vd = os.getenv("VECTOR_DIM", "").strip()
dim = int(vd) if vd.isdigit() else 1024
return np.empty((0, dim), dtype=np.float32)
# 无论调用方传多大,不能超过远端接口单次条数上限
step = max(1, min(int(batch_size), self._api_max_batch))
all_embeddings: List[np.ndarray] = []
iterator = range(0, len(texts), step)
if show_progress:
try:
from tqdm import tqdm
iterator = tqdm(iterator, desc="Embedding (API)")
except ImportError:
pass
for i in iterator:
batch = texts[i : i + step]
resp = self._client.embeddings.create(
model=self.model,
input=batch,
encoding_format="float",
)
rows = sorted(
[(d.index, d.embedding) for d in resp.data],
key=lambda x: x[0],
)
batch_embs = np.array([e for _, e in rows], dtype=np.float32)
if batch_embs.size > 0:
self._set_dim_from_vector(batch_embs[0].tolist())
if normalize:
norms = np.linalg.norm(batch_embs, axis=1, keepdims=True)
batch_embs = batch_embs / (norms + 1e-10)
all_embeddings.append(batch_embs)
return np.vstack(all_embeddings).astype(np.float32)
def similarity(self, emb1: np.ndarray, emb2: np.ndarray) -> np.ndarray:
return np.dot(emb1, emb2.T)
def encode_and_search(
self,
query: str,
documents: List[str],
top_k: int = 5,
) -> List[dict]:
query_emb = self.encode([query], normalize=True)
doc_embs = self.encode(documents, normalize=True)
scores = self.similarity(query_emb, doc_embs)[0]
top_indices = np.argsort(scores)[::-1][:top_k]
return [
{"score": float(scores[idx]), "document": documents[idx], "index": int(idx)}
for idx in top_indices
]
# 向后兼容旧名称
DashScopeOpenAIEmbedding = OpenAICompatibleRemoteEmbedding
# 全局单例(避免重复加载模型,节省显存/内存)
_embedding_instance: Optional[Any] = None
def get_embedder(
model_path: Optional[str] = None,
device: Optional[str] = None,
force_reload: bool = False,
) -> Any:
"""
获取 Embedding 单例:USE_LOCAL_EMBEDDING=true 时用本地 Qwen3,否则用远程
OpenAI 兼容 API(优先级见 OpenAICompatibleRemoteEmbedding)。
"""
global _embedding_instance
if force_reload or _embedding_instance is None:
if _env_flag("USE_LOCAL_EMBEDDING", "true"):
_embedding_instance = Qwen3Embedding(
model_path=model_path,
device=device,
)
else:
_embedding_instance = OpenAICompatibleRemoteEmbedding()
return _embedding_instance
def clear_embedder():
"""清空单例(用于测试或切换模型)"""
global _embedding_instance
_embedding_instance = None
import gc
gc.collect()
if _TRANSFORMERS_AVAILABLE:
import torch
if torch.cuda.is_available():
torch.cuda.empty_cache()
+356
View File
@@ -0,0 +1,356 @@
"""
Few-shot示例选择器 - 基于经验数据集动态选择相关示例
与 Schema 向量检索一致,使用 utils.embedding.get_embedder()(本地 Qwen3 或远程 OpenAI 兼容 API,
由 USE_LOCAL_EMBEDDING 及 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 等环境变量决定)。
用法:
from utils.fewshot_selector import FewShotSelector
selector = FewShotSelector("data/experiences/all_samples.jsonl")
examples = selector.select(question="查询2024年1月的销售额", top_k=3)
# 在Prompt中使用
prompt = f"{examples}\n当前问题:{question}\nSchema:{schema}"
"""
import json
import os
from pathlib import Path
from typing import List, Dict, Optional
from dataclasses import dataclass
import numpy as np
import logging
logger = logging.getLogger(__name__)
_DEFAULT_LOCAL_EMBED_PATH = "./data/models/Qwen3-Embedding-0.6B"
@dataclass
class ExperienceSample:
"""经验数据样本"""
qid: str
question_zh: str
question_en: Optional[str]
sql: str
explanation: str
rating: Optional[int]
tags: List[str]
difficulty: str
@classmethod
def from_dict(cls, data: dict) -> "ExperienceSample":
return cls(
qid=data.get("qid", ""),
question_zh=data.get("question_zh", ""),
question_en=data.get("question_en"),
sql=data.get("sql", ""),
explanation=data.get("explanation", ""),
rating=data.get("rating"),
tags=data.get("tags", []),
difficulty=data.get("difficulty", "medium")
)
def to_fewshot_format(self, include_explanation: bool = True) -> str:
"""转换为few-shot格式"""
result = f"问题:{self.question_zh}\nSQL:\n{self.sql}"
if include_explanation and self.explanation:
result += f"\n说明:{self.explanation[:200]}"
return result
def to_dict(self) -> dict:
return {
"qid": self.qid,
"question": self.question_zh,
"sql": self.sql,
"rating": self.rating,
"tags": self.tags,
"difficulty": self.difficulty
}
class FewShotSelector:
"""
Few-shot示例选择器
根据用户问题,从经验数据集中检索最相似的示例,
用于增强Prompt,提升LLM生成质量。
"""
def __init__(
self,
samples_path: str,
embedding_model_path: Optional[str] = None,
use_cache: bool = True,
):
"""
Args:
samples_path: 样本JSONL文件路径
embedding_model_path: 本地 Embedding 模型目录;None 时用环境变量
EMBEDDING_MODEL_PATH(仅 USE_LOCAL_EMBEDDING=true 时有效)
use_cache: 是否缓存样本向量(按向量维度分文件,换模型会自动重建)
"""
self.samples_path = Path(samples_path)
self.samples: List[ExperienceSample] = []
self._embedder = None
self.embeddings: Optional[np.ndarray] = None
self.use_cache = use_cache
self.cache_path: Optional[Path] = None
self._embedding_model_path = (
embedding_model_path
if embedding_model_path is not None
else os.getenv("EMBEDDING_MODEL_PATH", _DEFAULT_LOCAL_EMBED_PATH).strip()
)
self._load_samples()
self._build_index()
def _load_samples(self):
"""加载样本数据"""
if not self.samples_path.exists():
raise FileNotFoundError(f"样本文件不存在: {self.samples_path}")
logger.info(f"加载样本: {self.samples_path}")
with self.samples_path.open('r', encoding='utf-8') as f:
for line in f:
data = json.loads(line.strip())
self.samples.append(ExperienceSample.from_dict(data))
logger.info(f"[OK] 加载 {len(self.samples)} 个样本")
def _build_index(self):
"""用项目统一 Embedder 构建语义索引"""
from utils.embedding import get_embedder
self._embedder = get_embedder(self._embedding_model_path)
probe = self._embedder.encode(
[" "],
batch_size=1,
normalize=True,
show_progress=False,
)
dim = int(probe.shape[1])
self.cache_path = (
self.samples_path.parent / f"{self.samples_path.stem}.fewshot_dim{dim}.npy"
)
if self.use_cache and self.cache_path.exists():
try:
self.embeddings = np.load(self.cache_path)
if (
self.embeddings.shape[0] == len(self.samples)
and self.embeddings.shape[1] == dim
):
logger.info(f"[OK] 加载 Few-shot 向量缓存: {self.cache_path}")
return
except Exception as e:
logger.warning(f"Few-shot 缓存加载失败: {e},将重新计算")
# 远程 API 通常不接受空字符串作 input
questions = [
(s.question_zh or "").strip() or " "
for s in self.samples
]
if not questions:
self.embeddings = np.empty((0, dim), dtype=np.float32)
return
logger.info(f"计算 {len(questions)} 个 Few-shot 样本向量...")
self.embeddings = self._embedder.encode(
questions,
batch_size=min(32, len(questions)),
normalize=True,
show_progress=True,
)
if self.use_cache and self.cache_path is not None:
np.save(self.cache_path, self.embeddings)
logger.info(f"[OK] Few-shot 向量已缓存: {self.cache_path}")
def select(
self,
question: str,
top_k: int = 3,
min_rating: Optional[int] = None,
required_tags: Optional[List[str]] = None,
max_difficulty: str = "hard",
exclude_qids: Optional[List[str]] = None
) -> List[ExperienceSample]:
"""
选择最相关的few-shot示例
Args:
question: 用户问题
top_k: 返回示例数量
min_rating: 最低评分(None表示不限制)
required_tags: 必须包含的标签(如["aggregation", "join"])
max_difficulty: 最大难度(过滤更难的示例)
exclude_qids: 排除的QID(避免与当前问题相同)
Returns:
排序后的示例列表(最相关优先)
"""
if self._embedder is None or self.embeddings is None:
logger.error("Few-shot 索引未初始化")
return []
if len(self.samples) == 0:
return []
q_emb = self._embedder.encode(
[question],
batch_size=1,
normalize=True,
show_progress=False,
)[0]
scores = np.dot(self.embeddings, q_emb)
candidates = []
for idx, (score, sample) in enumerate(zip(scores, self.samples)):
if exclude_qids and sample.qid in exclude_qids:
continue
if min_rating and sample.rating and sample.rating < min_rating:
continue
if max_difficulty == "easy" and sample.difficulty != "easy":
continue
if max_difficulty == "medium" and sample.difficulty == "hard":
continue
if required_tags and not all(tag in sample.tags for tag in required_tags):
continue
candidates.append((idx, score, sample))
candidates.sort(key=lambda x: -x[1])
selected = [sample for _, _, sample in candidates[:top_k]]
logger.info(
f"Few-shot选择: 问题='{question[:30]}...' "
f"→ 选中{len(selected)}个示例 (top_k={top_k}, min_rating={min_rating})"
)
for s in selected:
logger.debug(
f" [{s.qid}] {s.question_zh[:50]}... (rating={s.rating}, tags={s.tags[:3]})"
)
return selected
def get_examples_prompt(
self,
question: str,
top_k: int = 3,
min_rating: int = 7,
**kwargs
) -> str:
"""
生成few-shot prompt片段
Returns:
格式化的示例字符串,可直接插入Prompt
"""
examples = self.select(question, top_k=top_k, min_rating=min_rating, **kwargs)
if not examples:
return ""
lines = ["以下为相似问题的参考SQL示例:\n"]
for i, ex in enumerate(examples, 1):
lines.append(f"示例{i}:")
lines.append(f"问题:{ex.question_zh}")
lines.append(f"SQL:\n{ex.sql}")
if ex.explanation:
lines.append(f"说明:{ex.explanation[:150]}...")
lines.append("") # 空行分隔
return "\n".join(lines)
def get_tagged_examples(self, tags: List[str], top_k_per_tag: int = 2) -> str:
"""获取特定标签的示例"""
tagged_samples = []
for sample in self.samples:
if any(tag in sample.tags for tag in tags):
tagged_samples.append(sample)
tagged_samples.sort(key=lambda s: -(s.rating or 0))
selected = tagged_samples[:top_k_per_tag * len(tags)]
lines = [f"# {tags} 相关示例\n"]
for ex in selected:
lines.append(f"## {ex.qid}. {ex.question_zh[:50]}")
lines.append(f"评分: {ex.rating}/10")
lines.append(f"标签: {', '.join(ex.tags)}")
lines.append(f"```sql\n{ex.sql}\n```\n")
return "\n".join(lines)
def get_stats(self) -> Dict:
"""获取数据集统计"""
stats = {
"total": len(self.samples),
"by_rating": {},
"by_difficulty": {},
"by_tag": {},
"avg_rating": 0.0
}
ratings = [s.rating for s in self.samples if s.rating]
if ratings:
stats["avg_rating"] = sum(ratings) / len(ratings)
for r in range(1, 11):
stats["by_rating"][r] = sum(1 for s in self.samples if s.rating == r)
for diff in ["easy", "medium", "hard"]:
stats["by_difficulty"][diff] = sum(
1 for s in self.samples if s.difficulty == diff
)
tag_counts = {}
for s in self.samples:
for tag in s.tags:
tag_counts[tag] = tag_counts.get(tag, 0) + 1
stats["by_tag"] = tag_counts
return stats
def load_fewshot_selector() -> FewShotSelector:
"""加载默认的few-shot选择器"""
default_path = Path(__file__).resolve().parent.parent / "data" / "experiences" / "all_samples.jsonl"
return FewShotSelector(str(default_path))
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="Few-shot示例选择器")
parser.add_argument("--samples", default="data/experiences/all_samples.jsonl")
parser.add_argument("--question", help="测试问题")
parser.add_argument("--top-k", type=int, default=3)
parser.add_argument("--min-rating", type=int, default=7)
parser.add_argument("--stats", action="store_true", help="显示数据集统计")
args = parser.parse_args()
selector = FewShotSelector(args.samples)
if args.stats:
stats = selector.get_stats()
print("📊 数据集统计:")
print(f" 总样本: {stats['total']}")
print(f" 平均评分: {stats['avg_rating']:.1f}")
print(f" 难度分布: {stats['by_difficulty']}")
print(f"\n Top 10 标签:")
sorted_tags = sorted(stats["by_tag"].items(), key=lambda x: -x[1])[:10]
for tag, count in sorted_tags:
print(f" {tag}: {count}")
elif args.question:
examples = selector.select(args.question, top_k=args.top_k, min_rating=args.min_rating)
print(f"\n为问题 '{args.question}' 选择的示例:\n")
for ex in examples:
print(f"[{ex.qid}] 评分:{ex.rating} 难度:{ex.difficulty}")
print(f"问题: {ex.question_zh}")
print(f"SQL:\n{ex.sql}\n")
else:
print("请指定 --question 或 --stats")
+364
View File
@@ -0,0 +1,364 @@
"""
SQL 解析与验证工具(基于 sqlglot)
"""
import logging
import re
from typing import List, Tuple, Optional, Dict
import sqlglot
from sqlglot import exp, parse_one, ParseError
from schema.manager import SchemaManager
from schema.models import Table, Column
logger = logging.getLogger(__name__)
def validate_sql_syntax(sql: str, dialect: str = "tsql") -> Tuple[bool, List[str]]:
"""
验证SQL语法是否正确
Args:
sql: SQL语句
dialect: SQL方言
Returns:
(是否有效, 错误信息列表)
"""
errors = []
try:
# 尝试解析
parsed = parse_one(sql, dialect=dialect)
if parsed is None:
errors.append("SQL解析返回空结果")
return False, errors
# 检查是否为只读查询(SELECT/CTE/SHOW等)
# 动态获取可用表达式类型(兼容不同sqlglot版本)
readable_ops = [exp.Select, exp.Union, exp.Intersect, exp.Except, exp.With]
# 可选:添加 Show, Describe, Explain(如果存在)
for op_name in ['Show', 'Describe', 'Explain']:
if hasattr(exp, op_name):
readable_ops.append(getattr(exp, op_name))
if not isinstance(parsed, tuple(readable_ops)):
op_type = type(parsed).__name__
errors.append(f"非查询操作({op_type}),只允许SELECT等只读语句")
return True, []
except ParseError as e:
errors.append(f"SQL语法错误: {str(e)}")
return False, errors
except Exception as e:
errors.append(f"解析异常: {str(e)}")
return False, errors
def extract_tables_from_sql(sql: str, dialect: str = "tsql") -> List[str]:
"""
从SQL中提取所有表名
Args:
sql: SQL语句
dialect: SQL方言
Returns:
表名列表(去重)
"""
try:
parsed = parse_one(sql, dialect=dialect)
tables = []
# 遍历AST查找所有表名
for node in parsed.walk():
if isinstance(node, exp.Table):
table_name = node.name
if table_name and table_name not in tables:
tables.append(table_name)
return tables
except Exception as e:
logger.warning(f"提取表名失败: {e}")
return []
def build_table_alias_map(parsed: exp.Expression) -> Dict[str, str]:
"""
从已解析的 AST 构建「别名/表名 -> 物理表名」映射。
FROM T a 时 a -> T,且 T -> T,便于将 a.col 解析到表 T 的列。
"""
alias_map: Dict[str, str] = {}
for node in parsed.walk():
if not isinstance(node, exp.Table):
continue
physical = node.name
if not physical:
continue
alias_map[physical] = physical
talias = node.args.get("alias")
if talias is not None:
aname = talias.name
if aname:
alias_map[aname] = physical
return alias_map
def extract_columns_from_sql(sql: str, dialect: str = "tsql") -> List[Tuple[str, str]]:
"""
从SQL中提取所有字段引用(表.字段)
Args:
sql: SQL语句
dialect: SQL方言
Returns:
[(表名, 字段名), ...] 列表
"""
columns = []
try:
parsed = parse_one(sql, dialect=dialect)
for node in parsed.walk():
if isinstance(node, exp.Column):
table_name = node.table
col_name = node.name
if table_name and col_name:
columns.append((table_name, col_name))
return columns
except Exception as e:
logger.warning(f"提取字段失败: {e}")
return []
def validate_schema_consistency(
sql: str,
schema_manager: SchemaManager,
dialect: str = "tsql"
) -> Tuple[bool, List[str]]:
"""
验证SQL与Schema的一致性
检查:
1. 所有表名存在于Schema
2. 所有字段名属于对应的表
3. JOIN条件字段存在且类型兼容
Args:
sql: SQL语句
schema_manager: Schema管理器
dialect: SQL方言
Returns:
(是否一致, 错误信息列表)
"""
errors = []
try:
parsed = parse_one(sql, dialect=dialect)
except Exception as e:
logger.debug(f"Schema一致性检查跳过(解析失败): {e}")
return True, []
alias_map = build_table_alias_map(parsed)
tables_used: List[str] = []
for node in parsed.walk():
if isinstance(node, exp.Table):
tname = node.name
if tname and tname not in tables_used:
tables_used.append(tname)
columns_used: List[Tuple[str, str]] = []
for node in parsed.walk():
if isinstance(node, exp.Column):
tref, cname = node.table, node.name
if tref and cname:
columns_used.append((tref, cname))
# 检查表存在性
for tbl in tables_used:
if not schema_manager.get_table(tbl):
errors.append(f"表不存在: '{tbl}'")
# 检查字段存在性(表引用可为物理表名或别名)
for tbl_name, col_name in columns_used:
physical = alias_map.get(tbl_name, tbl_name)
table = schema_manager.get_table(physical)
if not table:
errors.append(
f"字段不存在: '{tbl_name}.{col_name}'"
f"(无法将表引用解析到已加载Schema中的表)"
)
continue
col_names = [c.name for c in table.columns]
if col_name not in col_names:
errors.append(
f"字段不存在: '{tbl_name}.{col_name}'"
f"(表 '{physical}' 可用字段: {col_names[:5]}...)"
)
# 检查JOIN条件(外键匹配)
try:
for join in parsed.find_all(exp.Join):
# 解析ON条件
on_condition = join.args.get("on")
if on_condition:
# 检查ON条件中涉及的字段
for eq in on_condition.find_all(exp.EQ):
left = eq.left
right = eq.right
# 提取左右两边的表.字段
for side in [left, right]:
if isinstance(side, exp.Column):
tbl = side.table
col = side.name
physical = alias_map.get(tbl, tbl)
table = schema_manager.get_table(physical)
if table and col not in [c.name for c in table.columns]:
errors.append(f"JOIN条件字段不存在: {tbl}.{col}")
except Exception as e:
logger.debug(f"JOIN条件检查异常: {e}")
return len(errors) == 0, errors
def rewrite_mysql_builtins_for_tsql(sql: str) -> str:
"""
模型在 T-SQL 目标下仍常输出 MySQL 函数;sqlglot 转写也可能遗漏。
SQL Server 无 CURDATE()/NOW(),需替换为 GETDATE 族。
"""
if not sql:
return sql
out = sql
out = re.sub(
r"\bCURDATE\s*\(\s*\)",
"CAST(GETDATE() AS DATE)",
out,
flags=re.IGNORECASE,
)
out = re.sub(
r"\bCURRENT_DATE\b",
"CAST(GETDATE() AS DATE)",
out,
flags=re.IGNORECASE,
)
out = re.sub(r"\bNOW\s*\(\s*\)", "GETDATE()", out, flags=re.IGNORECASE)
return out
def normalize_sql_for_dialect(sql: str, dialect: str) -> str:
"""
将模型输出的 SQL 规范为目标方言。
对于 T-SQL,主要进行 MySQL 函数替换(因为模型仍可能输出 CURDATE() 等)。
"""
sql = (sql or "").strip()
if not sql:
return sql
# 如果目标是 T-SQL,只做函数名替换,不再用 sqlglot 转写
if dialect == "tsql":
return rewrite_mysql_builtins_for_tsql(sql)
return sql
def format_sql(sql: str, dialect: str = "tsql", indent: int = 2) -> str:
"""
格式化SQL(可读性)
Args:
sql: SQL语句
dialect: SQL方言
indent: 缩进空格数
Returns:
格式化后的SQL
"""
try:
parsed = parse_one(sql, dialect=dialect)
return parsed.sql(dialect=dialect, pretty=True, indent=indent)
except Exception as e:
logger.warning(f"SQL格式化失败: {e}")
return sql
def normalize_sql(sql: str, dialect: str = "tsql") -> str:
"""
标准化SQL(用于比较去重)
去除多余空格、统一引号、移除注释等
Args:
sql: SQL语句
dialect: SQL方言
Returns:
标准化后的SQL
"""
try:
# 解析后重新生成(会规范化格式)
parsed = parse_one(sql, dialect=dialect)
normalized = parsed.sql(dialect=dialect, pretty=False)
# 转换为大写关键词
return normalized.upper()
except Exception:
# 降级:简单处理
import re
# 移除多余空格
sql = re.sub(r'\s+', ' ', sql.strip())
# 移除注释
sql = re.sub(r'--.*?$', '', sql, flags=re.MULTILINE)
sql = re.sub(r'/\*.*?\*/', '', sql, flags=re.DOTALL)
return sql.upper()
def count_joins(sql: str, dialect: str = "tsql") -> int:
"""统计JOIN数量"""
try:
parsed = parse_one(sql, dialect=dialect)
joins = list(parsed.find_all(exp.Join))
return len(joins)
except Exception:
return 0
def has_subquery(sql: str, dialect: str = "tsql") -> bool:
"""检查是否包含子查询"""
try:
parsed = parse_one(sql, dialect=dialect)
# 检查嵌套的SELECT
for select in parsed.find_all(exp.Select):
if select is not parsed: # 不是最外层的SELECT
return True
return False
except Exception:
return False
def get_query_complexity(sql: str, dialect: str = "tsql") -> Dict[str, int]:
"""
评估查询复杂度
Returns:
复杂度指标字典
"""
try:
parsed = parse_one(sql, dialect=dialect)
return {
"join_count": len(list(parsed.find_all(exp.Join))),
"subquery_count": len([s for s in parsed.find_all(exp.Select) if s is not parsed]),
"where_conditions": len(list(parsed.find_all(exp.Predicate))),
"aggregation_functions": len(list(parsed.find_all(exp.AggFunc))),
"column_count": len(list(parsed.find_all(exp.Column))),
}
except Exception as e:
logger.warning(f"复杂度评估失败: {e}")
return {}
+346
View File
@@ -0,0 +1,346 @@
"""
SQL 验证工具集
"""
import re
import logging
from typing import Tuple, List, Dict
logger = logging.getLogger(__name__)
# 危险操作关键词(除非明确允许)
DANGEROUS_KEYWORDS = [
"DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE",
"CREATE", "DROP DATABASE", "DROP TABLE", "DROP INDEX",
"GRANT", "REVOKE", "PURGE", "FLUSH", "KILL"
]
# 允许的操作(仅查询)
ALLOWED_KEYWORDS = [
"SELECT", "WITH", "FROM", "WHERE", "JOIN", "LEFT JOIN", "RIGHT JOIN",
"INNER JOIN", "OUTER JOIN", "ON", "USING", "GROUP BY", "HAVING",
"ORDER BY", "LIMIT", "OFFSET", "UNION", "UNION ALL", "EXCEPT", "INTERSECT",
"AS", "CASE", "WHEN", "THEN", "ELSE", "END",
"COUNT", "SUM", "AVG", "MIN", "MAX", "DISTINCT",
"AND", "OR", "NOT", "IN", "EXISTS", "BETWEEN", "LIKE", "IS NULL", "IS NOT NULL",
"CAST", "COALESCE", "NULLIF", "IFNULL",
"DATE", "TIME", "TIMESTAMP", "EXTRACT", "DATE_FORMAT", "STR_TO_DATE",
"CURRENT_DATE", "CURRENT_TIMESTAMP",
]
def check_dangerous_operations(sql: str) -> Tuple[bool, List[str]]:
"""
检查SQL是否包含危险操作
Args:
sql: SQL语句(大小写不敏感)
Returns:
(是否安全, 危险关键词列表)
"""
sql_upper = sql.upper()
found_dangers = []
for keyword in DANGEROUS_KEYWORDS:
# 使用正则避免部分匹配(如"DROP"不应匹配"DROPOUT")
pattern = r'\b' + re.escape(keyword) + r'\b'
if re.search(pattern, sql_upper):
found_dangers.append(keyword)
is_safe = len(found_dangers) == 0
if not is_safe:
logger.warning(f"检测到危险操作: {found_dangers}")
return is_safe, found_dangers
def validate_no_dml(sql: str) -> Tuple[bool, str]:
"""
验证SQL不是DML/DDL操作(仅允许SELECT等查询)
Returns:
(是否通过, 错误消息)
"""
sql_upper = sql.strip().upper()
# 检查是否以危险关键词开头
first_word = sql_upper.split()[0] if sql_upper.split() else ""
if first_word in ["INSERT", "UPDATE", "DELETE", "DROP", "ALTER", "CREATE", "TRUNCATE"]:
return False, f"禁止的操作: {first_word}"
is_safe, dangers = check_dangerous_operations(sql)
if not is_safe:
return False, f"SQL包含危险操作: {', '.join(dangers)}"
return True, ""
def check_sql_injection_patterns(sql: str) -> List[str]:
"""
检查明显的SQL注入模式
Args:
sql: SQL语句
Returns:
发现的注入模式列表
"""
patterns = {
"union_all_injection": r"UNION\s+ALL\s+SELECT",
"union_injection": r"UNION\s+SELECT",
"comment_injection": r"(--|\#|/\*).*SELECT",
"semicolon_injection": r";\s*(DROP|DELETE|UPDATE|INSERT)",
"or_true_condition": r"OR\s+['\"]?\s*1\s*['\"]?\s*=\s*1",
"always_true": r"1\s*=\s*1",
}
findings = []
sql_lower = sql.lower()
for name, pattern in patterns.items():
if re.search(pattern, sql, re.IGNORECASE):
findings.append(name)
return findings
def validate_aggregation_groupby(sql: str, dialect: str = "tsql") -> Tuple[bool, List[str]]:
"""
验证聚合查询的GROUP BY正确性
检查:SELECT中的非聚合字段是否都在GROUP BY中
Args:
sql: SQL语句
dialect: SQL方言
Returns:
(是否有效, 错误列表)
"""
from utils.sql_parser import parse_one, exp
errors = []
try:
parsed = parse_one(sql, dialect=dialect)
# 只检查SELECT语句
if not isinstance(parsed, exp.Select):
return True, []
# 获取SELECT列表中的表达式
select_exprs = parsed.expressions
# 获取GROUP BY字段
group_by = parsed.args.get("group")
if not group_by:
# 没有GROUP BY但有聚合函数,通常是错误的
has_agg = any(
expr.find(exp.AggFunc) is not None
for expr in select_exprs
)
if has_agg:
errors.append("包含聚合函数但缺少GROUP BY子句")
return len(errors) == 0, errors
group_by_exprs = group_by.expressions
# 提取GROUP BY的字段名(简单处理)
group_by_cols = set()
for expr in group_by_exprs:
if isinstance(expr, exp.Column):
group_by_cols.add(expr.name)
elif isinstance(expr, exp.Ordered):
# GROUP BY x ASC/DESC
this = expr.this
if isinstance(this, exp.Column):
group_by_cols.add(this.name)
# 检查每个SELECT表达式
for expr in select_exprs:
# 如果是聚合函数,跳过
if expr.find(exp.AggFunc):
continue
# 如果是字面量或表达式,跳过
if isinstance(expr, exp.Literal):
continue
# 如果是列引用,检查是否在GROUP BY中
if isinstance(expr, exp.Column):
col_name = expr.name
if col_name not in group_by_cols:
errors.append(
f"字段 '{col_name}' 在SELECT中但不在GROUP BY中"
)
elif isinstance(expr, exp.Alias):
# 别名: column AS alias
this = expr.this
if isinstance(this, exp.Column):
col_name = this.name
if col_name not in group_by_cols:
errors.append(
f"字段 '{col_name}' (别名为'{expr.alias}') 在SELECT中但不在GROUP BY中"
)
except Exception as e:
logger.debug(f"GROUP BY验证异常: {e}")
return len(errors) == 0, errors
def check_join_conditions(sql: str, dialect: str = "tsql") -> List[str]:
"""
检查JOIN条件是否完整
Args:
sql: SQL语句
dialect: SQL方言
Returns:
问题列表(空表示无问题)
"""
from utils.sql_parser import parse_one, exp
issues = []
try:
parsed = parse_one(sql, dialect=dialect)
# 遍历所有JOIN
for join in parsed.find_all(exp.Join):
# 检查是否有ON条件
on_condition = join.args.get("on")
if on_condition is None:
# 检查是否使用USING
using = join.args.get("using")
if using is None:
issues.append("JOIN缺少ON条件")
else:
# ON条件为空表达式
if isinstance(on_condition, exp.Empty):
issues.append("JOIN的ON条件为空")
except Exception as e:
logger.debug(f"JOIN条件检查异常: {e}")
return issues
def validate_order_by_fields(
sql: str,
schema_manager,
dialect: str = "tsql"
) -> List[str]:
"""
验证ORDER BY字段是否存在于对应表中
Args:
sql: SQL语句
schema_manager: Schema管理器
dialect: SQL方言
Returns:
问题列表
"""
from utils.sql_parser import parse_one, exp
issues = []
try:
parsed = parse_one(sql, dialect=dialect)
order = parsed.args.get("order")
if order:
for ordered in order.expressions:
expr = ordered.this
# 提取字段和表
if isinstance(expr, exp.Column):
tbl_name = expr.table
col_name = expr.name
if tbl_name:
table = schema_manager.get_table(tbl_name)
if table:
col_names = [c.name for c in table.columns]
if col_name not in col_names:
issues.append(
f"ORDER BY字段不存在: {tbl_name}.{col_name}"
)
except Exception as e:
logger.debug(f"ORDER BY验证异常: {e}")
return issues
def full_validation_pipeline(
sql: str,
schema_manager,
dialect: str = "tsql",
check_dangerous: bool = True
) -> Dict:
"""
完整验证流水线
Args:
sql: SQL语句
schema_manager: Schema管理器
dialect: SQL方言
check_dangerous: 是否检查危险操作
Returns:
验证结果字典
"""
result = {
"valid": True,
"errors": [],
"warnings": [],
"suggestions": []
}
# 1. 语法验证
syntax_ok, syntax_errors = validate_sql_syntax(sql, dialect)
if not syntax_ok:
result["valid"] = False
result["errors"].extend(syntax_errors)
# 2. 危险操作检查
if check_dangerous:
safe, dangers = check_dangerous_operations(sql)
if not safe:
result["valid"] = False
result["errors"].append(f"包含危险操作: {', '.join(dangers)}")
# 3. Schema一致性验证
schema_ok, schema_errors = validate_schema_consistency(sql, schema_manager, dialect)
if not schema_ok:
result["valid"] = False
result["errors"].extend(schema_errors)
# 4. GROUP BY验证
groupby_ok, groupby_errors = validate_aggregation_groupby(sql, dialect)
if not groupby_ok:
result["valid"] = False
result["errors"].extend(groupby_errors)
# 5. JOIN条件验证
join_issues = check_join_conditions(sql, dialect)
if join_issues:
result["valid"] = False
result["errors"].extend(join_issues)
# 6. ORDER BY验证
order_issues = validate_order_by_fields(sql, schema_manager, dialect)
if order_issues:
result["warnings"].extend(order_issues)
# 7. SQL注入模式检查(警告)
injection_patterns = check_sql_injection_patterns(sql)
if injection_patterns:
result["warnings"].append(f"检测到可疑模式: {', '.join(injection_patterns)}")
return result