#!/usr/bin/env python3 """ Text2SQL 多智能体系统 - CLI 入口 用法:在项目根目录执行 python backend/main.py """ import os import sys import argparse import logging from pathlib import Path # 从仓库根目录运行 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__) from bootstrap import ( create_orchestrator, load_project_env, load_schema, resolve_sql_dialect, setup_environment, ) 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="关闭问句归一中文(默认开启:中英均先归一句中文以利 SQL 一致;也可用 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()