Files
ai-g3sb-backman2.0/backend/main.py
T

413 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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()