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:
+24
-54
@@ -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(),
|
||||
|
||||
Reference in New Issue
Block a user