Refactor api_server.py to import environment setup and schema loading from bootstrap.py, enhancing modularity. Introduce a new function in Text2SQLOrchestrator to prioritize VCUserAccessibleFunction in table selection, improving SQL generation accuracy. Update validation logic to enforce restrictions on CJK characters in SQL string literals, ensuring compliance with business rules. Enhance prompts to clarify SQL generation constraints regarding date conditions and CJK usage.
This commit is contained in:
@@ -0,0 +1,215 @@
|
||||
"""
|
||||
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
|
||||
Reference in New Issue
Block a user