#!/usr/bin/env python3 """ Text2SQL 多智能体系统 - CLI 入口 用法:在项目根目录执行 python backend/main.py """ import os import sys import argparse import logging from pathlib import Path from typing import Optional from dotenv import load_dotenv # 从仓库根目录运行 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)) # 配置日志 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__) def _repo_root() -> Path: """仓库根目录(含 data/、.env、api_server.py 的目录)。""" return Path(__file__).resolve().parent.parent def _load_project_env(): """加载项目根目录 .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(): """环境检查""" _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 临时目录(单文件模式) 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 if not schema_path.exists(): logger.warning(f"Schema文件不存在: {schema_path}") logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录") return False # 检查API Key api_key = os.getenv("DEEPSEEK_API_KEY", "").strip() if not api_key: logger.warning("环境变量 DEEPSEEK_API_KEY 未设置") logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key") return False logger.info(f"[OK] 环境检查通过") logger.info(f" - Schema: {schema_path}") logger.info(f" - API Key: {'已配置' if api_key else '未配置'}") 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 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 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, 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 def single_query(orchestrator, question: str, dialect: str = "tsql"): """单次查询""" import time from agents.orchestrator import GenerationResult from utils.dialog_classifier import DialogIntent, classify_dialog logger.info(f"[Q] 问题: {question}") classified = classify_dialog(question, llm_client=orchestrator.deepseek) 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}, ) 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}") dbe = result.metadata.get("db_empty_feedback") if result.valid and dbe: print("\n[DB 探针 0 — 无数据行] 说明:") print(dbe) print(f"\n使用表: {result.tables_used}") return result def interactive_mode(orchestrator, dialect: str = "tsql"): """交互式模式""" print("\n" + "=" * 60) print("Text2SQL 交互模式(输入 'quit' 或 'exit' 退出)") print("提示:请描述业务数据查询需求;寒暄或「你是谁」等会由对话分类处理,不生成 SQL。") print("=" * 60 + "\n") while True: try: question = input("❓ 请输入业务查询问题: ").strip() 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( description="Text2SQL:运行后进入交互式自然语言生成 SQL(python backend/main.py)" ) 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( "--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="详细日志" ) parser.add_argument( "--no-translate-en", action="store_true", help="关闭英文问句自动译为中文(默认开启;也可用 TRANSLATE_EN_TO_ZH=false)", ) 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) interactive_mode(orchestrator, args.dialect) if __name__ == "__main__": main()