216 lines
7.5 KiB
Python
216 lines
7.5 KiB
Python
"""
|
||||
|
|
Text2SQL 共享启动逻辑:环境检查、Schema 加载、Orchestrator 构造。
|
|||
|
|
|
|||
|
|
供 `main` CLI 与 `api_server` 复用,避免 API 层依赖 CLI 入口模块。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
from __future__ import annotations
|
|||
|
|
|
|||
|
|
import logging
|
|||
|
|
import os
|
|||
|
|
import sys
|
|||
|
|
from pathlib import Path
|
|||
|
|
from typing import Any, Optional
|
|||
|
|
|
|||
|
|
from dotenv import load_dotenv
|
|||
|
|
|
|||
|
|
logger = logging.getLogger(__name__)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _repo_root() -> Path:
|
|||
|
|
"""仓库根目录(含 data/、.env、api_server.py 的目录)。"""
|
|||
|
|
return Path(__file__).resolve().parent.parent
|
|||
|
|
|
|||
|
|
|
|||
|
|
def load_project_env() -> None:
|
|||
|
|
"""加载项目根目录 .env,供后续 os.getenv 使用。"""
|
|||
|
|
load_dotenv(_repo_root() / ".env")
|
|||
|
|
|
|||
|
|
|
|||
|
|
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() -> bool:
|
|||
|
|
"""环境检查"""
|
|||
|
|
load_project_env()
|
|||
|
|
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("未设置 OPENAI_EMBEDDING_MODEL(远程 Embedding 模型名)")
|
|||
|
|
return False
|
|||
|
|
else:
|
|||
|
|
if not os.getenv("DASHSCOPE_API_KEY", "").strip():
|
|||
|
|
logger.warning(
|
|||
|
|
"未设置 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
|
|||
|
|
|
|||
|
|
# 检查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 临时目录(单文件模式)
|
|||
|
|
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: 仓库根目录(本文件位于 backend/)
|
|||
|
|
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
|
|||
|
|
|
|||
|
|
if not schema_path.exists():
|
|||
|
|
logger.warning(f"Schema文件不存在: {schema_path}")
|
|||
|
|
logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录")
|
|||
|
|
return False
|
|||
|
|
|
|||
|
|
# 检查 LLM Key(DeepSeek / OpenAI 可切换)
|
|||
|
|
llm_sc = (os.getenv("LLM_SERVICE_CODE") or "").strip().lower()
|
|||
|
|
if llm_sc and llm_sc not in ("deepseek", "openai"):
|
|||
|
|
logger.warning("未知 LLM_SERVICE_CODE=%r(仅支持 deepseek/openai)", llm_sc)
|
|||
|
|
return False
|
|||
|
|
if llm_sc == "openai":
|
|||
|
|
if not (os.getenv("OPENAI_API_KEY") or "").strip():
|
|||
|
|
logger.warning("LLM_SERVICE_CODE=openai 但 OPENAI_API_KEY 未设置")
|
|||
|
|
return False
|
|||
|
|
else:
|
|||
|
|
if not (os.getenv("DEEPSEEK_API_KEY") or "").strip():
|
|||
|
|
logger.warning("环境变量 DEEPSEEK_API_KEY 未设置(默认 LLM_SERVICE_CODE=deepseek)")
|
|||
|
|
logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key")
|
|||
|
|
return False
|
|||
|
|
|
|||
|
|
logger.info(f"[OK] 环境检查通过")
|
|||
|
|
logger.info(f" - Schema: {schema_path}")
|
|||
|
|
logger.info(
|
|||
|
|
" - LLM: %s",
|
|||
|
|
(llm_sc or "deepseek(auto)"),
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
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: Any, args: Any):
|
|||
|
|
"""创建编排器"""
|
|||
|
|
from agents.orchestrator import Text2SQLOrchestrator
|
|||
|
|
from llm.router import create_llm_client, resolve_llm_service_code
|
|||
|
|
from llm.deepseek_client import DeepSeekConfig
|
|||
|
|
|
|||
|
|
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
|
|||
|
|
|
|||
|
|
sc = resolve_llm_service_code()
|
|||
|
|
if sc == "deepseek":
|
|||
|
|
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()
|
|||
|
|
cfg = DeepSeekConfig(
|
|||
|
|
api_key=api_key,
|
|||
|
|
base_url=base_url,
|
|||
|
|
model_name=args.model or "deepseek-chat",
|
|||
|
|
temperature=args.temperature,
|
|||
|
|
max_tokens=args.max_tokens,
|
|||
|
|
)
|
|||
|
|
llm_client = create_llm_client("deepseek", **cfg.__dict__)
|
|||
|
|
else:
|
|||
|
|
# openai:完全由 OPENAI_* 决定;同时沿用 temperature/max_tokens 作为默认值覆盖
|
|||
|
|
llm_client = create_llm_client(
|
|||
|
|
"openai",
|
|||
|
|
temperature=args.temperature,
|
|||
|
|
max_tokens=args.max_tokens,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
orchestrator = Text2SQLOrchestrator(
|
|||
|
|
schema_manager=schema_mgr,
|
|||
|
|
llm_client=llm_client,
|
|||
|
|
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,
|
|||
|
|
translate_english_to_zh=translate_en,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
return orchestrator
|