Files
ai-g3sb-backman2.0/main.py
T
2026-04-10 16:52:07 +08:00

407 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 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()