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

242 lines
6.9 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
# 从仓库根目录运行 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()