Enhance environment configuration and error handling in API server and backend. Implement dynamic loading of .env files for various execution contexts, improve schema file path resolution, and refine vector search error handling in the orchestrator. Update ChromaDB integration to support memory mode and add persistence options for few-shot learning. Include additional logging for better traceability.

This commit is contained in:
lasean.zhou
2026-04-14 18:21:50 +08:00
parent ca8bc5e7de
commit 4ea3056e95
8 changed files with 161 additions and 41 deletions
+46 -3
View File
@@ -44,7 +44,45 @@ logging.basicConfig(
) )
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
load_dotenv(Path(__file__).resolve().parent / ".env") # 加载 .env 文件(支持 PyInstaller 打包后的目录结构)
def _find_env_file() -> Path:
"""查找 .env 文件,支持多种运行环境"""
import sys
# 尝试1: 当前工作目录
cwd_env = Path.cwd() / ".env"
if cwd_env.exists():
return cwd_env
# 尝试2: PyInstaller 临时目录(单文件模式)
if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'):
meipass_env = Path(sys._MEIPASS) / ".env"
if meipass_env.exists():
return meipass_env
# 尝试3: 脚本/可执行文件所在目录
if getattr(sys, 'frozen', False):
# PyInstaller 打包后
exe_dir = Path(sys.executable).parent
env_path = exe_dir / ".env"
if env_path.exists():
return env_path
else:
# 开发环境
script_dir = Path(__file__).resolve().parent
env_path = script_dir / ".env"
if env_path.exists():
return env_path
# 默认返回当前目录
return cwd_env
env_file = _find_env_file()
if env_file.exists():
load_dotenv(env_file)
logger.info(f"[OK] 已加载配置文件: {env_file}")
else:
logger.warning(f"[WARN] 未找到 .env 文件: {env_file}")
orchestrator = None orchestrator = None
schema_manager = None schema_manager = None
@@ -78,7 +116,7 @@ def get_orchestrator():
max_retry = int(os.getenv("MAX_RETRY", "2")) max_retry = int(os.getenv("MAX_RETRY", "2"))
embedding_model = None embedding_model = None
vector_db = os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma") vector_db = os.getenv("VECTOR_DB_PATH", "./data/embeddings/chroma")
no_vector_search = False no_vector_search = False # 启用向量搜索(ChromaDB 已修复)
no_fewshot = not os.getenv("FEWSHOT_ENABLED", "true").lower() in ("true", "1", "yes") no_fewshot = not os.getenv("FEWSHOT_ENABLED", "true").lower() in ("true", "1", "yes")
fewshot_top_k = int(os.getenv("FEWSHOT_TOP_K", "3")) fewshot_top_k = int(os.getenv("FEWSHOT_TOP_K", "3"))
fewshot_min_rating = int(os.getenv("FEWSHOT_MIN_RATING", "7")) fewshot_min_rating = int(os.getenv("FEWSHOT_MIN_RATING", "7"))
@@ -644,6 +682,7 @@ async def admin_visibility():
if __name__ == "__main__": if __name__ == "__main__":
import uvicorn import uvicorn
import multiprocessing
port = int(os.getenv("API_PORT", "8041")) port = int(os.getenv("API_PORT", "8041"))
host = os.getenv("API_HOST", "0.0.0.0") host = os.getenv("API_HOST", "0.0.0.0")
@@ -651,8 +690,12 @@ if __name__ == "__main__":
logger.info(f"启动服务: http://{host}:{port}") logger.info(f"启动服务: http://{host}:{port}")
logger.info(f"API文档: http://{host}:{port}/docs") logger.info(f"API文档: http://{host}:{port}/docs")
# PyInstaller + multiprocessing(spawn)兼容
multiprocessing.freeze_support()
# PyInstaller 打包后必须使用 app 对象,不能使用字符串模块名
uvicorn.run( uvicorn.run(
"api_server:app", app,
host=host, host=host,
port=port, port=port,
reload=False, reload=False,
+25 -14
View File
@@ -175,25 +175,36 @@ class Text2SQLOrchestrator:
""" """
if not self.use_vector_search: if not self.use_vector_search:
# 不使用向量检索时,返回所有表 # 不使用向量检索时,返回所有表
logger.info("[Orchestrator] 向量搜索已禁用,使用所有表")
return self.schema_manager.list_tables() return self.schema_manager.list_tables()
indexer = self._get_vector_index() try:
indexer = self._get_vector_index()
# 确保索引已构建 # 确保索引已构建
if indexer.count() == 0: if indexer.count() == 0:
logger.info("向量索引为空,正在构建...") logger.info("向量索引为空,正在构建...")
indexer.build_index(self.schema_manager, force_rebuild=True) indexer.build_index(self.schema_manager, force_rebuild=True)
logger.info(f"[OK] 索引构建完成,共 {indexer.count()} 张表")
# 检索 # 检索
results = indexer.search( logger.info(f"[Orchestrator] 开始向量检索: query='{question[:50]}...'")
query=question, results = indexer.search(
top_k=top_k, query=question,
score_threshold=0.1 # 降低阈值以提高召回率(原0.2) top_k=top_k,
) score_threshold=0.1 # 降低阈值以提高召回率
)
candidate_tables = [r["table_name"] for r in results] candidate_tables = [r["table_name"] for r in results]
logger.debug(f"粗筛候选表:{candidate_tables[:10]}...(共{len(candidate_tables)}个)") logger.debug(f"粗筛候选表:{candidate_tables[:10]}...(共{len(candidate_tables)}个)")
return candidate_tables return candidate_tables
except Exception as e:
logger.error(f"[Orchestrator] 向量检索失败: {e},降级为使用所有表", exc_info=True)
import traceback
logger.error(traceback.format_exc())
# 降级:返回所有表
return self.schema_manager.list_tables()
def _llm_select_tables( def _llm_select_tables(
self, self,
+33 -3
View File
@@ -100,19 +100,49 @@ def setup_environment():
logger.warning("未设置 DASHSCOPE_MODEL") logger.warning("未设置 DASHSCOPE_MODEL")
return False return False
# 检查Schema文件 # 检查Schema文件(支持相对路径和绝对路径)
schema_path = Path("./data/schemas/G3SB_MCDataDictionary_table_structure.json") schema_path_str = os.getenv("SCHEMA_PATH", "./data/schemas/G3SB_MCDataDictionary_table_structure.json")
schema_path = Path(schema_path_str)
# 如果是相对路径,尝试从多个位置查找
if not schema_path.is_absolute():
# 尝试1: PyInstaller 临时目录(单文件模式)
import sys
if getattr(sys, 'frozen', False) and hasattr(sys, '_MEIPASS'):
meipass_schema = Path(sys._MEIPASS) / schema_path_str
if meipass_schema.exists():
schema_path = meipass_schema
# 尝试2: 当前工作目录
if not schema_path.exists():
schema_path = Path.cwd() / schema_path_str
# 尝试3: 脚本所在目录
if not schema_path.exists():
script_dir = Path(__file__).resolve().parent.parent
schema_path = script_dir / schema_path_str
# 尝试4: 可执行文件所在目录
if not schema_path.exists() and getattr(sys, 'frozen', False):
exe_dir = Path(sys.executable).parent
schema_path = exe_dir / schema_path_str
if not schema_path.exists(): if not schema_path.exists():
logger.warning(f"Schema文件不存在: {schema_path}") logger.warning(f"Schema文件不存在: {schema_path}")
logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录") logger.info("请将G3SB Schema文件放置在 ./data/schemas/ 目录")
return False return False
# 检查API Key # 检查API Key
if not os.getenv("DEEPSEEK_API_KEY"): api_key = os.getenv("DEEPSEEK_API_KEY", "").strip()
if not api_key:
logger.warning("环境变量 DEEPSEEK_API_KEY 未设置") logger.warning("环境变量 DEEPSEEK_API_KEY 未设置")
logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key") logger.info("请在 .env 文件中配置,或 export DEEPSEEK_API_KEY=your_key")
return False return False
logger.info(f"[OK] 环境检查通过")
logger.info(f" - Schema: {schema_path}")
logger.info(f" - API Key: {'已配置' if api_key else '未配置'}")
return True return True
+9 -7
View File
@@ -20,7 +20,7 @@ class SchemaIndexer:
功能: 功能:
1. 为表结构构建向量索引 1. 为表结构构建向量索引
2. 基于问题的表检索 2. 基于问题的表检索
3. 持久化存储和加载 3. Chroma 使用内存模式(EphemeralClient),进程退出后不保留;``persist_dir`` 仅作配置/统计引用
""" """
def __init__( def __init__(
@@ -34,17 +34,15 @@ class SchemaIndexer:
Args: Args:
embedder: Embedding模型实例 embedder: Embedding模型实例
persist_dir: 向量数据库持久化目录 persist_dir: 历史配置中的向量库路径(仅展示与统计,不落盘)
collection_name: 集合名称 collection_name: 集合名称
""" """
self.embedder = embedder self.embedder = embedder
self.persist_dir = Path(persist_dir) self.persist_dir = Path(persist_dir)
self.persist_dir.mkdir(parents=True, exist_ok=True)
self.collection_name = collection_name self.collection_name = collection_name
# 初始化ChromaDB客户端 # 内存 Chroma:不落盘,每次进程需重新 build_index
self.client = chromadb.PersistentClient( self.client = chromadb.EphemeralClient(
path=str(self.persist_dir),
settings=ChromaSettings(anonymized_telemetry=False), settings=ChromaSettings(anonymized_telemetry=False),
) )
@@ -54,7 +52,10 @@ class SchemaIndexer:
metadata={"hnsw:space": "cosine"}, # 使用余弦相似度 metadata={"hnsw:space": "cosine"}, # 使用余弦相似度
) )
logger.info(f"[OK] 初始化SchemaIndexer: collection={self.collection_name}") logger.info(
f"[OK] 初始化SchemaIndexer(Chroma内存): collection={self.collection_name}, "
f"配置路径引用={self.persist_dir}"
)
def build_index( def build_index(
self, self,
@@ -249,4 +250,5 @@ class SchemaIndexer:
"total_columns": total_columns, "total_columns": total_columns,
"avg_columns": total_columns / count if count > 0 else 0, "avg_columns": total_columns / count if count > 0 else 0,
"persist_dir": str(self.persist_dir), "persist_dir": str(self.persist_dir),
"chroma_mode": "memory",
} }
+26 -11
View File
@@ -1,8 +1,9 @@
""" """
Few-shot 经验样本的 Chroma 向量库存储与检索。 Few-shot 经验样本的 Chroma 向量库存储与检索。
与 SchemaIndexer 一致:使用项目统一 ``get_embedder()``,持久化目录默认 与 SchemaIndexer 一致:使用项目统一 ``get_embedder()``;默认 Chroma 为内存
``./data/embeddings/chroma_fewshot``,集合名 ``fewshot_samples``。 (``EphemeralClient``)。离线灌库脚本可设 ``persist_to_disk=True`` 写入磁盘。
集合名默认 ``fewshot_samples``。
""" """
from __future__ import annotations from __future__ import annotations
@@ -34,7 +35,16 @@ class FewShotChromaStore:
embedder: Any, embedder: Any,
persist_dir: Optional[str] = None, persist_dir: Optional[str] = None,
collection_name: Optional[str] = None, collection_name: Optional[str] = None,
*,
persist_to_disk: bool = False,
): ):
"""
Args:
embedder: 编码器实例。
persist_dir: 磁盘模式下的持久化目录;内存模式下仍解析为配置引用路径。
collection_name: 集合名;可空,空则读环境变量或默认。
persist_to_disk: 为 True 时使用 ``PersistentClient`` 落盘(如灌库脚本)。
"""
self.embedder = embedder self.embedder = embedder
# Chroma 的 collection 名;显式参数优先,否则读 FEWSHOT_CHROMA_COLLECTION,再回退默认 # Chroma 的 collection 名;显式参数优先,否则读 FEWSHOT_CHROMA_COLLECTION,再回退默认
explicit = (collection_name or "").strip() explicit = (collection_name or "").strip()
@@ -43,20 +53,25 @@ class FewShotChromaStore:
self.persist_dir = Path( self.persist_dir = Path(
(persist_dir or os.getenv("FEWSHOT_CHROMA_PATH") or DEFAULT_FEWSHOT_CHROMA_DIR).strip() (persist_dir or os.getenv("FEWSHOT_CHROMA_PATH") or DEFAULT_FEWSHOT_CHROMA_DIR).strip()
) )
self.persist_dir.mkdir(parents=True, exist_ok=True) self.persist_to_disk = persist_to_disk
self.client = chromadb.PersistentClient( if persist_to_disk:
path=str(self.persist_dir), self.persist_dir.mkdir(parents=True, exist_ok=True)
settings=ChromaSettings(anonymized_telemetry=False), self.client = chromadb.PersistentClient(
) path=str(self.persist_dir),
settings=ChromaSettings(anonymized_telemetry=False),
)
else:
self.client = chromadb.EphemeralClient(
settings=ChromaSettings(anonymized_telemetry=False),
)
self.collection = self.client.get_or_create_collection( self.collection = self.client.get_or_create_collection(
name=self.collection_name, name=self.collection_name,
metadata={"hnsw:space": "cosine"}, metadata={"hnsw:space": "cosine"},
) )
backend = f"磁盘 {self.persist_dir}" if persist_to_disk else f"内存(配置路径={self.persist_dir})"
logger.info( logger.info(
"[OK] FewShotChromaStore: path=%s collection=%s count=%s", f"[OK] FewShotChromaStore: {backend} collection={self.collection_name} "
self.persist_dir, f"count={self.collection.count()}"
self.collection_name,
self.collection.count(),
) )
def count(self) -> int: def count(self) -> int:
Binary file not shown.
+21 -2
View File
@@ -2,7 +2,7 @@
import os import os
from pathlib import Path from pathlib import Path
from PyInstaller.utils.hooks import collect_data_files, collect_submodules from PyInstaller.utils.hooks import collect_data_files, collect_submodules, collect_dynamic_libs
# 获取项目根目录(spec 文件所在目录) # 获取项目根目录(spec 文件所在目录)
import sys import sys
@@ -18,6 +18,24 @@ datas = [
# ('data/experiences', 'data/experiences'), # ('data/experiences', 'data/experiences'),
] ]
# ChromaDB 在 Windows 单文件 exe 下如缺少 native 动态库,可能在持久化访问时直接“硬退出”
# 这里显式收集 chromadb 与 hnswlib(向量索引)的动态库/二进制依赖
binaries = []
try:
binaries += collect_dynamic_libs("chromadb")
except Exception:
pass
try:
binaries += collect_dynamic_libs("hnswlib")
except Exception:
pass
# 某些 chromadb 版本会依赖包内数据文件(如 migrations/默认配置等),一并收集更稳妥
try:
datas += collect_data_files("chromadb")
except Exception:
pass
# 收集所有需要的隐藏导入 # 收集所有需要的隐藏导入
hiddenimports = [ hiddenimports = [
# Web 框架 # Web 框架
@@ -57,6 +75,7 @@ hiddenimports = [
'backend.schema.models', 'backend.schema.models',
'backend.utils', 'backend.utils',
'backend.utils.dialog_classifier', 'backend.utils.dialog_classifier',
'backend.utils.dialog_context',
'backend.utils.embedding', 'backend.utils.embedding',
'backend.utils.fewshot_selector', 'backend.utils.fewshot_selector',
'backend.utils.question_locale', 'backend.utils.question_locale',
@@ -118,7 +137,7 @@ excludes = [
a = Analysis( a = Analysis(
['api_server.py'], ['api_server.py'],
pathex=[str(_REPO_DIR)], pathex=[str(_REPO_DIR)],
binaries=[], binaries=binaries,
datas=datas, datas=datas,
hiddenimports=hiddenimports, hiddenimports=hiddenimports,
hookspath=[], hookspath=[],
+1 -1
View File
@@ -90,7 +90,7 @@ def main() -> int:
if not persist.is_absolute(): if not persist.is_absolute():
persist = _REPO_ROOT / persist persist = _REPO_ROOT / persist
store = FewShotChromaStore(embedder, persist_dir=str(persist)) store = FewShotChromaStore(embedder, persist_dir=str(persist), persist_to_disk=True)
n = store.build_from_samples(rows, force_rebuild=args.force) n = store.build_from_samples(rows, force_rebuild=args.force)
logger.info("完成: 写入 %s 条 → %s", n, persist) logger.info("完成: 写入 %s 条 → %s", n, persist)
return 0 return 0