242 lines
6.9 KiB
Python
242 lines
6.9 KiB
Python
#!/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()
|