first commit

This commit is contained in:
陈辅元
2026-04-10 16:52:07 +08:00
commit 84fe545640
87 changed files with 20842 additions and 0 deletions
+406
View File
@@ -0,0 +1,406 @@
#!/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()