#!/usr/bin/env python3 """ Text2SQL 多智能体系统 - CLI演示入口 用法: python main.py "查询2024年1月销售额最高的前5个产品" python main.py --question "查询所有状态为Active的账户数量" python main.py --interactive # 交互式模式 """ import os import sys import argparse import logging from pathlib import Path from typing import Optional from dotenv import load_dotenv # 配置日志 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__) _DEFAULT_EMBEDDING_PATH = "./data/models/Qwen3-Embedding-0.6B" def _load_project_env(): """加载项目根目录 .env(与 main.py 同目录),供后续 os.getenv 使用。""" load_dotenv(Path(__file__).resolve().parent / ".env") def _embedding_model_path() -> str: return os.getenv("EMBEDDING_MODEL_PATH", _DEFAULT_EMBEDDING_PATH).strip() 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() use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "true").strip().lower() in ( "1", "true", "yes", "on", ) if use_local_emb: # 与 .env 中 EMBEDDING_MODEL_PATH 及 utils.embedding 一致 model_path = Path(_embedding_model_path()) if not model_path.exists(): logger.warning(f"Embedding模型不存在: {model_path}") logger.info("请先下载模型:") logger.info(" modelscope download --model 'Qwen/Qwen3-Embedding-0.6B' " f"--local_dir '{model_path}'") logger.info("或使用远程 Embedding API:USE_LOCAL_EMBEDDING=false,并配置 " "MODELSCOPE_API_KEY、或 OPENAI_API_KEY+OPENAI_EMBEDDING_MODEL" "(及可选 OPENAI_BASE_URL)、或 DASHSCOPE_*(百炼)") 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 # 检查Schema文件 schema_path = Path("./data/schemas/G3SB_MCDataDictionary_table_structure.json") if not schema_path.exists(): logger.warning(f"Schema文件不存在: {schema_path}") logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录") return False # 检查API Key if not os.getenv("DEEPSEEK_API_KEY"): logger.warning("环境变量 DEEPSEEK_API_KEY 未设置") logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key") return False 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 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, ) return orchestrator def single_query(orchestrator, question: str, dialect: str = "tsql"): """单次查询""" import time logger.info(f"[Q] 问题: {question}") 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}") print(f"\n使用表: {result.tables_used}") return result def interactive_mode(orchestrator, dialect: str = "tsql"): """交互式模式""" print("\n" + "=" * 60) print("Text2SQL 交互模式(输入 'quit' 或 'exit' 退出)") 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 batch_mode(orchestrator, questions: list, dialect: str = "tsql"): """批量查询模式""" print(f"\n批量模式:共 {len(questions)} 个问题\n") results = [] for i, question in enumerate(questions, 1): print(f"[{i}/{len(questions)}] {question}") result = single_query(orchestrator, question, dialect) results.append(result) print() # 统计 success_count = sum(1 for r in results if r.valid) print("=" * 60) print(f"统计: {success_count}/{len(questions)} 成功 " f"({success_count/len(questions)*100:.1f}%)") return results def main(): parser = argparse.ArgumentParser( description="Text2SQL 多智能体系统 - 自然语言生成SQL" ) parser.add_argument( "question", nargs="?", help="自然语言问题(如不提供则进入交互模式)" ) 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( "--interactive", "-i", action="store_true", help="交互模式" ) parser.add_argument( "--batch", "-b", help="批量文件路径(每行一个问题)" ) parser.add_argument( "--verbose", "-v", action="store_true", help="详细日志" ) 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) # 根据参数选择模式 if args.interactive or (not args.question and not args.batch): interactive_mode(orchestrator, args.dialect) elif args.batch: with open(args.batch, 'r', encoding='utf-8') as f: questions = [line.strip() for line in f if line.strip()] batch_mode(orchestrator, questions, args.dialect) elif args.question: single_query(orchestrator, args.question, args.dialect) else: parser.print_help() if __name__ == "__main__": main()