""" 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