Refactor embedding configuration to remove local model support, transitioning to a unified remote API approach. Update environment variables and documentation accordingly. Enhance error handling in the orchestrator and related modules to reflect these changes. This update simplifies the embedding process and improves overall system reliability.

This commit is contained in:
陈辅元
2026-04-15 09:49:18 +08:00
parent 4ea3056e95
commit 2d0b1ab36f
21 changed files with 103 additions and 405 deletions
+24 -54
View File
@@ -37,11 +37,6 @@ def _load_project_env():
load_dotenv(_repo_root() / ".env")
def _embedding_model_path() -> str:
"""仅 ``USE_LOCAL_EMBEDDING=true`` 时需要;远程 Embedding 可为空。"""
return os.getenv("EMBEDDING_MODEL_PATH", "").strip()
def resolve_sql_dialect(name: str) -> str:
"""CLI / 配置中的方言别名统一为 sqlglot 方言名(SQL Server -> tsql)。"""
n = (name or "sqlserver").lower().strip()
@@ -53,52 +48,33 @@ def resolve_sql_dialect(name: str) -> str:
def setup_environment():
"""环境检查"""
_load_project_env()
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "false").strip().lower() in (
"1", "true", "yes", "on",
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 use_local_emb:
mp_str = _embedding_model_path()
if not mp_str:
logger.warning(
"USE_LOCAL_EMBEDDING=true 但未设置 EMBEDDING_MODEL_PATH(本地模型目录)"
)
return False
model_path = Path(mp_str)
if not model_path.exists():
logger.warning("本地 Embedding 目录不存在: %s", model_path)
logger.info(
"请设置正确的 EMBEDDING_MODEL_PATH,或改用 USE_LOCAL_EMBEDDING=false,"
"并配置 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding"
)
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:
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
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")
@@ -214,7 +190,6 @@ def create_orchestrator(schema_mgr, args):
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,
@@ -366,11 +341,6 @@ def main():
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(),