0.1.1 暂存
This commit is contained in:
@@ -0,0 +1 @@
|
||||
# utils 包初始化
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,135 @@
|
||||
"""
|
||||
用户输入意图分类:区分「自然语言查数 / Text2SQL」与「寒暄、致谢、元问题」等不适合直接生成 SQL 的对话。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
import unicodedata
|
||||
from enum import Enum
|
||||
from typing import NamedTuple, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DialogIntent(str, Enum):
|
||||
TEXT2SQL = "text2sql"
|
||||
CONVERSATION = "conversation"
|
||||
|
||||
|
||||
class DialogClassifyResult(NamedTuple):
|
||||
intent: DialogIntent
|
||||
"""若为 CONVERSATION,可展示给用户的引导文案;TEXT2SQL 时为 None。"""
|
||||
reply_suggestion: Optional[str] = None
|
||||
|
||||
|
||||
DEFAULT_CONVERSATION_REPLY = (
|
||||
"您好,我是业务库 Text2SQL 助手。\n"
|
||||
"请用自然语言描述要查询或统计的内容(例如:查询某账户可用余额、按经纪商汇总未结算交易笔数)。\n"
|
||||
"输入 quit 或 exit 可退出。"
|
||||
)
|
||||
|
||||
_EMPTY_INPUT_REPLY = "请输入具体的业务查询问题,或输入 quit 退出。"
|
||||
|
||||
# 一旦出现,倾向于按「要查数据」处理(含常见业务词,避免误判)
|
||||
_SQL_OR_QUERY_HINT_RE = re.compile(
|
||||
r"(查|查询|查出|检索|统计|列出|汇总|求和|平均|分组|排序|排名|显示|导出|筛选|过滤|"
|
||||
r"多少|几个|几张|哪些|占比|同比|环比|"
|
||||
r"余额|交易|账户|持仓|报表|结算|合约|订单|流水|经纪商|对手方|证券|资金|"
|
||||
r"query|select|list|show|count|sum|avg|how\s+many|statistics|\bfrom\b|\bwhere\b|\btable\b)",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
_CHITCHAT_PHRASES = frozenset(
|
||||
{
|
||||
"你好",
|
||||
"您好",
|
||||
"嗨",
|
||||
"哈喽",
|
||||
"hello",
|
||||
"hi",
|
||||
"hey",
|
||||
"早上好",
|
||||
"下午好",
|
||||
"晚上好",
|
||||
"在吗",
|
||||
"在不在",
|
||||
"谢谢",
|
||||
"多谢",
|
||||
"感谢",
|
||||
"thanks",
|
||||
"thank you",
|
||||
"thx",
|
||||
"再见",
|
||||
"拜拜",
|
||||
"bye",
|
||||
"goodbye",
|
||||
"哈哈",
|
||||
"哈哈哈",
|
||||
"嗯",
|
||||
"嗯嗯",
|
||||
"好的",
|
||||
"好",
|
||||
"ok",
|
||||
"okay",
|
||||
"行",
|
||||
"收到",
|
||||
"👋",
|
||||
"😀",
|
||||
"哈哈谢谢",
|
||||
}
|
||||
)
|
||||
|
||||
_CHITCHAT_KEYS = frozenset(p.casefold() for p in _CHITCHAT_PHRASES)
|
||||
|
||||
_META_QUESTION_RE = re.compile(
|
||||
r"(你是谁|你是什么|你能(做|干)什么|你会什么|怎么用|如何使用|使用说明|帮助|help\b|"
|
||||
r"什么功能|干啥的)",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _normalize(text: str) -> str:
|
||||
t = unicodedata.normalize("NFKC", text or "").strip()
|
||||
t = re.sub(r"\s+", " ", t)
|
||||
return t
|
||||
|
||||
|
||||
def _strip_trailing_punct(t: str) -> str:
|
||||
return re.sub(r"[!!。.??,,;;:~~…、]+$", "", t).strip()
|
||||
|
||||
|
||||
def classify_dialog(user_text: str) -> DialogClassifyResult:
|
||||
"""
|
||||
对用户一轮输入做粗分类。
|
||||
|
||||
策略:优先用「查询/业务」关键词锁定 TEXT2SQL;否则对短寒暄、致谢、元问题判为 CONVERSATION;
|
||||
其余默认 TEXT2SQL,避免漏判真实查询。
|
||||
"""
|
||||
t = _normalize(user_text)
|
||||
if not t:
|
||||
return DialogClassifyResult(
|
||||
DialogIntent.CONVERSATION, reply_suggestion=_EMPTY_INPUT_REPLY
|
||||
)
|
||||
|
||||
if _SQL_OR_QUERY_HINT_RE.search(t):
|
||||
logger.debug("[dialog] intent=text2sql (query/business hint)")
|
||||
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
|
||||
|
||||
core = _strip_trailing_punct(t)
|
||||
if core.casefold() in _CHITCHAT_KEYS:
|
||||
logger.debug("[dialog] intent=conversation (chitchat phrase)")
|
||||
return DialogClassifyResult(
|
||||
DialogIntent.CONVERSATION, reply_suggestion=DEFAULT_CONVERSATION_REPLY
|
||||
)
|
||||
|
||||
if _META_QUESTION_RE.search(t):
|
||||
logger.debug("[dialog] intent=conversation (meta question)")
|
||||
return DialogClassifyResult(
|
||||
DialogIntent.CONVERSATION,
|
||||
reply_suggestion=DEFAULT_CONVERSATION_REPLY,
|
||||
)
|
||||
|
||||
logger.debug("[dialog] intent=text2sql (default)")
|
||||
return DialogClassifyResult(DialogIntent.TEXT2SQL, None)
|
||||
@@ -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()
|
||||
@@ -0,0 +1,357 @@
|
||||
"""
|
||||
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选择器"""
|
||||
# __file__ = backend/utils/fewshot_selector.py → 仓库根为 parents[2]
|
||||
default_path = Path(__file__).resolve().parents[2] / "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")
|
||||
@@ -0,0 +1,23 @@
|
||||
"""自然语言问题语种启发式(用于是否走「英译中」再 Text2SQL)。"""
|
||||
|
||||
import re
|
||||
|
||||
|
||||
_CJK_RE = re.compile(r"[\u4e00-\u9fff]")
|
||||
_HANGUL_RE = re.compile(r"[\uac00-\ud7af]")
|
||||
_KANA_RE = re.compile(r"[\u3040-\u30ff]")
|
||||
|
||||
|
||||
def looks_like_english_only(text: str) -> bool:
|
||||
"""
|
||||
判断问题是否主要为英文(无中日韩表意文字),适合先译成中文再走检索/选表。
|
||||
|
||||
含中文、日文假名、韩文时不翻译,避免破坏中英混合问句。
|
||||
"""
|
||||
s = (text or "").strip()
|
||||
if not s or len(s) < 2:
|
||||
return False
|
||||
if _CJK_RE.search(s) or _HANGUL_RE.search(s) or _KANA_RE.search(s):
|
||||
return False
|
||||
latin = sum(1 for c in s if ("a" <= c <= "z") or ("A" <= c <= "Z"))
|
||||
return latin >= 3
|
||||
@@ -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 {}
|
||||
@@ -0,0 +1,430 @@
|
||||
"""
|
||||
SQL 验证工具集
|
||||
"""
|
||||
|
||||
import re
|
||||
import logging
|
||||
from typing import Tuple, List, Dict
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# CJK Unified Ideographs + 兼容扩展(用于禁止中文业务词出现在 SQL 字符串字面量中)
|
||||
_CJK_IN_STRING_RE = re.compile(
|
||||
r"[\u3000-\u303f\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff\uff00-\uffef]"
|
||||
)
|
||||
|
||||
|
||||
def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
|
||||
"""
|
||||
扫描 SQL 中单引号字符串(含 T-SQL N'…'),若字面量内出现 CJK 则判失败。
|
||||
|
||||
跳过 ``--`` 行注释与 ``/* */`` 块注释内的文本,避免误报。
|
||||
"""
|
||||
errors: List[str] = []
|
||||
i = 0
|
||||
n = len(sql)
|
||||
in_line_comment = False
|
||||
in_block_comment = False
|
||||
|
||||
def _read_single_quoted_string(start: int) -> Tuple[str, int]:
|
||||
"""从 start 指向的 opening `'` 之后开始读,返回 (内容, 闭合引号后下标)。"""
|
||||
j = start
|
||||
parts: List[str] = []
|
||||
while j < n:
|
||||
ch = sql[j]
|
||||
if ch == "'":
|
||||
if j + 1 < n and sql[j + 1] == "'":
|
||||
parts.append("'")
|
||||
j += 2
|
||||
continue
|
||||
return "".join(parts), j + 1
|
||||
parts.append(ch)
|
||||
j += 1
|
||||
return "".join(parts), j
|
||||
|
||||
while i < n:
|
||||
if in_line_comment:
|
||||
if sql[i] == "\n":
|
||||
in_line_comment = False
|
||||
i += 1
|
||||
continue
|
||||
if in_block_comment:
|
||||
if i + 1 < n and sql[i : i + 2] == "*/":
|
||||
in_block_comment = False
|
||||
i += 2
|
||||
else:
|
||||
i += 1
|
||||
continue
|
||||
|
||||
two = sql[i : i + 2]
|
||||
if two == "--":
|
||||
in_line_comment = True
|
||||
i += 2
|
||||
continue
|
||||
if two == "/*":
|
||||
in_block_comment = True
|
||||
i += 2
|
||||
continue
|
||||
|
||||
# N' 或 n' 前缀的 Unicode 字面量
|
||||
if i + 1 < n and sql[i] in "Nn" and sql[i + 1] == "'":
|
||||
body, i = _read_single_quoted_string(i + 2)
|
||||
if _CJK_IN_STRING_RE.search(body):
|
||||
prev = body[:48] + ("…" if len(body) > 48 else "")
|
||||
errors.append(
|
||||
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
|
||||
f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。"
|
||||
)
|
||||
continue
|
||||
|
||||
if sql[i] == "'":
|
||||
body, i = _read_single_quoted_string(i + 1)
|
||||
if _CJK_IN_STRING_RE.search(body):
|
||||
prev = body[:48] + ("…" if len(body) > 48 else "")
|
||||
errors.append(
|
||||
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
|
||||
f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。"
|
||||
)
|
||||
continue
|
||||
|
||||
i += 1
|
||||
|
||||
return len(errors) == 0, errors
|
||||
|
||||
|
||||
# 危险操作关键词(除非明确允许)
|
||||
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
|
||||
Reference in New Issue
Block a user