2026-04-10 16:52:07 +08:00
|
|
|
|
#!/usr/bin/env python3
|
|
|
|
|
|
"""
|
2026-04-14 10:28:22 +08:00
|
|
|
|
Text2SQL 多智能体系统 - CLI 入口
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
用法:在项目根目录执行 python backend/main.py
|
2026-04-10 16:52:07 +08:00
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
import os
|
|
|
|
|
|
import sys
|
|
|
|
|
|
import argparse
|
|
|
|
|
|
import logging
|
|
|
|
|
|
from pathlib import Path
|
|
|
|
|
|
from typing import Optional
|
|
|
|
|
|
|
|
|
|
|
|
from dotenv import load_dotenv
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
# 从仓库根目录运行 python backend/main.py 时,将 backend 加入模块搜索路径
|
|
|
|
|
|
_backend_dir = Path(__file__).resolve().parent
|
|
|
|
|
|
if str(_backend_dir) not in sys.path:
|
|
|
|
|
|
sys.path.insert(0, str(_backend_dir))
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
# 配置日志
|
|
|
|
|
|
logging.basicConfig(
|
|
|
|
|
|
level=logging.INFO,
|
|
|
|
|
|
format='%(asctime)s [%(levelname)s] %(name)s: %(message)s',
|
|
|
|
|
|
datefmt='%Y-%m-%d %H:%M:%S'
|
|
|
|
|
|
)
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
def _repo_root() -> Path:
|
|
|
|
|
|
"""仓库根目录(含 data/、.env、api_server.py 的目录)。"""
|
|
|
|
|
|
return Path(__file__).resolve().parent.parent
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
def _load_project_env():
|
2026-04-14 10:28:22 +08:00
|
|
|
|
"""加载项目根目录 .env,供后续 os.getenv 使用。"""
|
|
|
|
|
|
load_dotenv(_repo_root() / ".env")
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _embedding_model_path() -> str:
|
2026-04-14 18:02:12 +08:00
|
|
|
|
"""仅 ``USE_LOCAL_EMBEDDING=true`` 时需要;远程 Embedding 可为空。"""
|
|
|
|
|
|
return os.getenv("EMBEDDING_MODEL_PATH", "").strip()
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def resolve_sql_dialect(name: str) -> str:
|
|
|
|
|
|
"""CLI / 配置中的方言别名统一为 sqlglot 方言名(SQL Server -> tsql)。"""
|
|
|
|
|
|
n = (name or "sqlserver").lower().strip()
|
|
|
|
|
|
if n in ("sqlserver", "mssql"):
|
|
|
|
|
|
return "tsql"
|
|
|
|
|
|
return n
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def setup_environment():
|
|
|
|
|
|
"""环境检查"""
|
|
|
|
|
|
_load_project_env()
|
2026-04-14 18:02:12 +08:00
|
|
|
|
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "false").strip().lower() in (
|
2026-04-10 16:52:07 +08:00
|
|
|
|
"1", "true", "yes", "on",
|
|
|
|
|
|
)
|
|
|
|
|
|
if use_local_emb:
|
2026-04-14 18:02:12 +08:00
|
|
|
|
mp_str = _embedding_model_path()
|
|
|
|
|
|
if not mp_str:
|
|
|
|
|
|
logger.warning(
|
|
|
|
|
|
"USE_LOCAL_EMBEDDING=true 但未设置 EMBEDDING_MODEL_PATH(本地模型目录)"
|
|
|
|
|
|
)
|
|
|
|
|
|
return False
|
|
|
|
|
|
model_path = Path(mp_str)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
if not model_path.exists():
|
2026-04-14 18:02:12 +08:00
|
|
|
|
logger.warning("本地 Embedding 目录不存在: %s", model_path)
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
"请设置正确的 EMBEDDING_MODEL_PATH,或改用 USE_LOCAL_EMBEDDING=false,"
|
|
|
|
|
|
"并配置 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding"
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
return False
|
|
|
|
|
|
else:
|
|
|
|
|
|
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
|
|
|
|
|
|
oa_key = os.getenv("OPENAI_API_KEY", "").strip()
|
|
|
|
|
|
oa_key_ok = oa_key and not (
|
|
|
|
|
|
oa_key.startswith("http://") or oa_key.startswith("https://")
|
|
|
|
|
|
)
|
|
|
|
|
|
if ms_key:
|
|
|
|
|
|
pass # ModelScope:BASE_URL / MODEL 有默认值,仅需 KEY
|
|
|
|
|
|
elif oa_key_ok:
|
|
|
|
|
|
if not os.getenv("OPENAI_EMBEDDING_MODEL", "").strip():
|
|
|
|
|
|
logger.warning("USE_LOCAL_EMBEDDING=false 但未设置 OPENAI_EMBEDDING_MODEL")
|
|
|
|
|
|
return False
|
|
|
|
|
|
else:
|
|
|
|
|
|
if not os.getenv("DASHSCOPE_API_KEY", "").strip():
|
|
|
|
|
|
logger.warning(
|
|
|
|
|
|
"USE_LOCAL_EMBEDDING=false 但未设置 MODELSCOPE_API_KEY、"
|
|
|
|
|
|
"OPENAI_API_KEY+OPENAI_EMBEDDING_MODEL 或 DASHSCOPE_API_KEY"
|
|
|
|
|
|
)
|
|
|
|
|
|
return False
|
|
|
|
|
|
base = (
|
|
|
|
|
|
os.getenv("DASHSCOPE_BASE_URL") or os.getenv("DASHSCOPE_base_url", "")
|
|
|
|
|
|
).strip()
|
|
|
|
|
|
if not base:
|
|
|
|
|
|
logger.warning("未设置 DASHSCOPE_BASE_URL(或 DASHSCOPE_base_url)")
|
|
|
|
|
|
return False
|
|
|
|
|
|
if not os.getenv("DASHSCOPE_MODEL", "").strip():
|
|
|
|
|
|
logger.warning("未设置 DASHSCOPE_MODEL")
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
2026-04-14 18:21:50 +08:00
|
|
|
|
# 检查Schema文件(支持相对路径和绝对路径)
|
|
|
|
|
|
schema_path_str = os.getenv("SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json")
|
|
|
|
|
|
schema_path = Path(schema_path_str)
|
|
|
|
|
|
|
|
|
|
|
|
# 如果是相对路径,尝试从多个位置查找
|
|
|
|
|
|
if not schema_path.is_absolute():
|
|
|
|
|
|
# 尝试1: PyInstaller 临时目录(单文件模式)
|
|
|
|
|
|
import sys
|
|
|
|
|
|
if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'):
|
|
|
|
|
|
meipass_schema = Path(sys._MEIPASS) / schema_path_str
|
|
|
|
|
|
if meipass_schema.exists():
|
|
|
|
|
|
schema_path = meipass_schema
|
|
|
|
|
|
|
|
|
|
|
|
# 尝试2: 当前工作目录
|
|
|
|
|
|
if not schema_path.exists():
|
|
|
|
|
|
schema_path = Path.cwd() / schema_path_str
|
|
|
|
|
|
|
|
|
|
|
|
# 尝试3: 脚本所在目录
|
|
|
|
|
|
if not schema_path.exists():
|
|
|
|
|
|
script_dir = Path(__file__).resolve().parent.parent
|
|
|
|
|
|
schema_path = script_dir / schema_path_str
|
|
|
|
|
|
|
|
|
|
|
|
# 尝试4: 可执行文件所在目录
|
|
|
|
|
|
if not schema_path.exists() and getattr(sys, 'frozen', False):
|
|
|
|
|
|
exe_dir = Path(sys.executable).parent
|
|
|
|
|
|
schema_path = exe_dir / schema_path_str
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
if not schema_path.exists():
|
|
|
|
|
|
logger.warning(f"Schema文件不存在: {schema_path}")
|
|
|
|
|
|
logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录")
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
|
|
|
|
|
# 检查API Key
|
2026-04-14 18:21:50 +08:00
|
|
|
|
api_key = os.getenv("DEEPSEEK_API_KEY", "").strip()
|
|
|
|
|
|
if not api_key:
|
2026-04-10 16:52:07 +08:00
|
|
|
|
logger.warning("环境变量 DEEPSEEK_API_KEY 未设置")
|
|
|
|
|
|
logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key")
|
|
|
|
|
|
return False
|
|
|
|
|
|
|
2026-04-14 18:21:50 +08:00
|
|
|
|
logger.info(f"[OK] 环境检查通过")
|
|
|
|
|
|
logger.info(f" - Schema: {schema_path}")
|
|
|
|
|
|
logger.info(f" - API Key: {'已配置' if api_key else '未配置'}")
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
return True
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _default_g3sb_meta_path(structure_path: str) -> Optional[str]:
|
|
|
|
|
|
"""若存在与 table_structure 同名的 table_meta 文件则返回其路径。"""
|
|
|
|
|
|
p = Path(structure_path)
|
|
|
|
|
|
if "table_structure" not in p.name:
|
|
|
|
|
|
return None
|
|
|
|
|
|
cand = p.parent / p.name.replace("table_structure", "table_meta")
|
|
|
|
|
|
return str(cand) if cand.is_file() else None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def load_schema(schema_path: str, schema_meta_path: Optional[str] = None):
|
|
|
|
|
|
"""加载 Schema;G3SB structure JSON 会自动尝试配对 table_meta(可用 --schema-meta 指定)。"""
|
|
|
|
|
|
from schema.manager import SchemaManager
|
|
|
|
|
|
|
|
|
|
|
|
meta = (
|
|
|
|
|
|
schema_meta_path
|
|
|
|
|
|
if schema_meta_path is not None
|
|
|
|
|
|
else _default_g3sb_meta_path(schema_path)
|
|
|
|
|
|
)
|
|
|
|
|
|
logger.info(f"加载Schema: {schema_path}")
|
|
|
|
|
|
if meta:
|
|
|
|
|
|
logger.info(f" 表注释(meta): {meta}")
|
|
|
|
|
|
|
|
|
|
|
|
schema_mgr = SchemaManager.load_from_json(
|
|
|
|
|
|
schema_path, g3sb_meta_path=meta
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
stats = schema_mgr.get_statistics()
|
|
|
|
|
|
logger.info(
|
|
|
|
|
|
f"[OK] Schema加载完成: {stats['database']}, "
|
|
|
|
|
|
f"共{stats['total_tables']}张表, {stats['total_columns']}个字段"
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
return schema_mgr
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def create_orchestrator(schema_mgr, args):
|
|
|
|
|
|
"""创建编排器"""
|
|
|
|
|
|
from agents.orchestrator import Text2SQLOrchestrator
|
|
|
|
|
|
from llm.deepseek_client import DeepSeekConfig
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
translate_en = os.getenv("TRANSLATE_EN_TO_ZH", "true").strip().lower() not in (
|
|
|
|
|
|
"0",
|
|
|
|
|
|
"false",
|
|
|
|
|
|
"no",
|
|
|
|
|
|
"off",
|
|
|
|
|
|
)
|
|
|
|
|
|
if getattr(args, "no_translate_en", False):
|
|
|
|
|
|
translate_en = False
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
api_key = (args.api_key or os.getenv("DEEPSEEK_API_KEY") or "").strip()
|
|
|
|
|
|
if not api_key:
|
|
|
|
|
|
raise ValueError(
|
|
|
|
|
|
"未配置 DeepSeek API Key:请在 .env 中设置 DEEPSEEK_API_KEY,"
|
|
|
|
|
|
"或使用命令行参数 --api-key"
|
|
|
|
|
|
)
|
|
|
|
|
|
base_url = (os.getenv("DEEPSEEK_BASE_URL") or "https://api.deepseek.com").strip()
|
|
|
|
|
|
|
|
|
|
|
|
config = DeepSeekConfig(
|
|
|
|
|
|
api_key=api_key,
|
|
|
|
|
|
base_url=base_url,
|
|
|
|
|
|
model_name=args.model or "deepseek-chat",
|
|
|
|
|
|
temperature=args.temperature,
|
|
|
|
|
|
max_tokens=args.max_tokens,
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
orchestrator = Text2SQLOrchestrator(
|
|
|
|
|
|
schema_manager=schema_mgr,
|
|
|
|
|
|
deepseek_config=config,
|
|
|
|
|
|
embedding_model_path=args.embedding_model,
|
|
|
|
|
|
vector_db_path=args.vector_db,
|
|
|
|
|
|
max_retry=args.max_retry,
|
|
|
|
|
|
use_vector_search=not args.no_vector_search,
|
|
|
|
|
|
# Few-shot配置
|
|
|
|
|
|
fewshot_enabled=not args.no_fewshot,
|
|
|
|
|
|
fewshot_top_k=args.fewshot_top_k,
|
|
|
|
|
|
fewshot_min_rating=args.fewshot_min_rating,
|
2026-04-14 10:28:22 +08:00
|
|
|
|
translate_english_to_zh=translate_en,
|
2026-04-10 16:52:07 +08:00
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
return orchestrator
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def single_query(orchestrator, question: str, dialect: str = "tsql"):
|
|
|
|
|
|
"""单次查询"""
|
|
|
|
|
|
import time
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
from agents.orchestrator import GenerationResult
|
|
|
|
|
|
from utils.dialog_classifier import DialogIntent, classify_dialog
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
logger.info(f"[Q] 问题: {question}")
|
|
|
|
|
|
|
2026-04-14 18:02:12 +08:00
|
|
|
|
classified = classify_dialog(question, llm_client=orchestrator.deepseek)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
if classified.intent == DialogIntent.CONVERSATION:
|
|
|
|
|
|
reply = classified.reply_suggestion or ""
|
|
|
|
|
|
logger.info("[Q] 意图: conversation(跳过 SQL 生成)")
|
|
|
|
|
|
print("\n" + "=" * 60)
|
|
|
|
|
|
print("对话 / 非查询输入(未触发 SQL 生成)")
|
|
|
|
|
|
print("=" * 60)
|
|
|
|
|
|
print(reply)
|
|
|
|
|
|
print(f"\n使用表: []")
|
|
|
|
|
|
return GenerationResult(
|
|
|
|
|
|
sql="",
|
|
|
|
|
|
valid=False,
|
|
|
|
|
|
errors=[],
|
|
|
|
|
|
warnings=[],
|
|
|
|
|
|
tables_used=[],
|
|
|
|
|
|
attempts=0,
|
|
|
|
|
|
metadata={"dialog_intent": DialogIntent.CONVERSATION.value},
|
|
|
|
|
|
)
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
start = time.time()
|
|
|
|
|
|
result = orchestrator.generate(
|
|
|
|
|
|
question=question,
|
|
|
|
|
|
dialect=dialect,
|
|
|
|
|
|
top_k_candidates=20
|
|
|
|
|
|
)
|
|
|
|
|
|
elapsed = time.time() - start
|
|
|
|
|
|
|
|
|
|
|
|
print("\n" + "=" * 60)
|
|
|
|
|
|
print("生成结果:")
|
|
|
|
|
|
print("=" * 60)
|
|
|
|
|
|
|
|
|
|
|
|
if result.valid:
|
|
|
|
|
|
print(f"[OK] SQL (耗时 {elapsed:.2f}s, 尝试 {result.attempts} 次):\n")
|
|
|
|
|
|
print(result.sql)
|
|
|
|
|
|
else:
|
|
|
|
|
|
print(f"[FAIL] 生成失败 (尝试 {result.attempts} 次)")
|
|
|
|
|
|
for err in result.errors:
|
|
|
|
|
|
print(f" - {err}")
|
|
|
|
|
|
|
|
|
|
|
|
if result.warnings:
|
|
|
|
|
|
print("\n[WARN] 警告:")
|
|
|
|
|
|
for w in result.warnings:
|
|
|
|
|
|
print(f" - {w}")
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
dbe = result.metadata.get("db_empty_feedback")
|
|
|
|
|
|
if result.valid and dbe:
|
|
|
|
|
|
print("\n[DB 探针 0 — 无数据行] 说明:")
|
|
|
|
|
|
print(dbe)
|
|
|
|
|
|
|
2026-04-10 16:52:07 +08:00
|
|
|
|
print(f"\n使用表: {result.tables_used}")
|
|
|
|
|
|
|
|
|
|
|
|
return result
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def interactive_mode(orchestrator, dialect: str = "tsql"):
|
|
|
|
|
|
"""交互式模式"""
|
|
|
|
|
|
print("\n" + "=" * 60)
|
|
|
|
|
|
print("Text2SQL 交互模式(输入 'quit' 或 'exit' 退出)")
|
2026-04-14 10:28:22 +08:00
|
|
|
|
print("提示:请描述业务数据查询需求;寒暄或「你是谁」等会由对话分类处理,不生成 SQL。")
|
2026-04-10 16:52:07 +08:00
|
|
|
|
print("=" * 60 + "\n")
|
|
|
|
|
|
|
|
|
|
|
|
while True:
|
|
|
|
|
|
try:
|
2026-04-14 10:28:22 +08:00
|
|
|
|
question = input("❓ 请输入业务查询问题: ").strip()
|
2026-04-10 16:52:07 +08:00
|
|
|
|
if question.lower() in ('quit', 'exit', 'q'):
|
|
|
|
|
|
print("再见!")
|
|
|
|
|
|
break
|
|
|
|
|
|
|
|
|
|
|
|
if not question:
|
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
|
|
result = single_query(orchestrator, question, dialect)
|
|
|
|
|
|
print()
|
|
|
|
|
|
|
|
|
|
|
|
except KeyboardInterrupt:
|
|
|
|
|
|
print("\n再见!")
|
|
|
|
|
|
break
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"查询失败: {e}")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def main():
|
|
|
|
|
|
parser = argparse.ArgumentParser(
|
2026-04-14 10:28:22 +08:00
|
|
|
|
description="Text2SQL:运行后进入交互式自然语言生成 SQL(python backend/main.py)"
|
2026-04-10 16:52:07 +08:00
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--schema", "-s",
|
|
|
|
|
|
default="./data/schemas/G3SB_MCDataDictionary_table_structure.json",
|
|
|
|
|
|
help="Schema文件路径(默认: ./data/schemas/G3SB_MCDataDictionary_table_structure.json)"
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--schema-meta",
|
|
|
|
|
|
default=None,
|
|
|
|
|
|
help="G3SB table_meta.json;默认自动使用同目录下文件名含 table_meta 的配对文件",
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--dialect", "-d",
|
|
|
|
|
|
default=os.getenv("TEXT2SQL_DIALECT", "sqlserver").strip(),
|
|
|
|
|
|
choices=["mysql", "postgresql", "sqlite", "tsql", "sqlserver", "mssql"],
|
|
|
|
|
|
help="SQL方言(默认: sqlserver / T-SQL;可用环境变量 TEXT2SQL_DIALECT 覆盖)",
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--api-key",
|
|
|
|
|
|
help="DeepSeek API Key(默认从DEEPSEEK_API_KEY环境变量读取)"
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--model", "-m",
|
|
|
|
|
|
default="deepseek-chat",
|
|
|
|
|
|
help="DeepSeek模型名称(默认: deepseek-chat)"
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--temperature", "-t",
|
|
|
|
|
|
type=float,
|
|
|
|
|
|
default=0.3,
|
|
|
|
|
|
help="生成温度(默认: 0.3)"
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--max-tokens",
|
|
|
|
|
|
type=int,
|
|
|
|
|
|
default=4096,
|
|
|
|
|
|
help="最大token数(默认: 4096)"
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--max-retry",
|
|
|
|
|
|
type=int,
|
|
|
|
|
|
default=2,
|
|
|
|
|
|
help="最大重试次数(默认: 2)"
|
|
|
|
|
|
)
|
|
|
|
|
|
_load_project_env()
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--embedding-model",
|
|
|
|
|
|
default=_embedding_model_path(),
|
|
|
|
|
|
help="Embedding模型路径(默认来自环境变量 EMBEDDING_MODEL_PATH)"
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--vector-db",
|
|
|
|
|
|
default=os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma").strip(),
|
|
|
|
|
|
help="向量数据库路径(默认来自环境变量 VECTOR_DB_PATH)"
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--no-vector-search",
|
|
|
|
|
|
action="store_true",
|
|
|
|
|
|
help="禁用向量检索(使用所有表)"
|
|
|
|
|
|
)
|
|
|
|
|
|
# Few-shot 配置
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--no-fewshot",
|
|
|
|
|
|
action="store_true",
|
|
|
|
|
|
help="禁用few-shot示例增强"
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--fewshot-top-k",
|
|
|
|
|
|
type=int,
|
|
|
|
|
|
default=int(os.getenv("FEWSHOT_TOP_K", "3")),
|
|
|
|
|
|
help="每次使用的few-shot示例数量(默认: 3)"
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--fewshot-min-rating",
|
|
|
|
|
|
type=int,
|
|
|
|
|
|
default=int(os.getenv("FEWSHOT_MIN_RATING", "7")),
|
|
|
|
|
|
help="few-shot示例最低评分(默认: 7)"
|
|
|
|
|
|
)
|
|
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--verbose", "-v",
|
|
|
|
|
|
action="store_true",
|
|
|
|
|
|
help="详细日志"
|
|
|
|
|
|
)
|
2026-04-14 10:28:22 +08:00
|
|
|
|
parser.add_argument(
|
|
|
|
|
|
"--no-translate-en",
|
|
|
|
|
|
action="store_true",
|
|
|
|
|
|
help="关闭英文问句自动译为中文(默认开启;也可用 TRANSLATE_EN_TO_ZH=false)",
|
|
|
|
|
|
)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
args = parser.parse_args()
|
|
|
|
|
|
args.dialect = resolve_sql_dialect(args.dialect)
|
|
|
|
|
|
|
|
|
|
|
|
# 日志级别
|
|
|
|
|
|
if args.verbose:
|
|
|
|
|
|
logging.getLogger().setLevel(logging.DEBUG)
|
|
|
|
|
|
|
|
|
|
|
|
# 环境检查
|
|
|
|
|
|
if not setup_environment():
|
|
|
|
|
|
sys.exit(1)
|
|
|
|
|
|
|
|
|
|
|
|
# 加载Schema
|
|
|
|
|
|
try:
|
|
|
|
|
|
schema_mgr = load_schema(args.schema, args.schema_meta)
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"Schema加载失败: {e}")
|
|
|
|
|
|
sys.exit(1)
|
|
|
|
|
|
|
|
|
|
|
|
# 创建Orchestrator
|
|
|
|
|
|
try:
|
|
|
|
|
|
orchestrator = create_orchestrator(schema_mgr, args)
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
|
logger.error(f"Orchestrator创建失败: {e}")
|
|
|
|
|
|
sys.exit(1)
|
|
|
|
|
|
|
2026-04-14 10:28:22 +08:00
|
|
|
|
interactive_mode(orchestrator, args.dialect)
|
2026-04-10 16:52:07 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__":
|
|
|
|
|
|
main()
|