Enhance dialog context handling for Text2SQL queries by integrating session history. Introduce new methods for summarizing previous assistant messages and determining if the last interaction was a data query. Update environment configuration for embedding options and improve error handling in SQL generation. Add user-facing delivery messages for successful SQL execution. This update supports more coherent follow-up questions and improves user experience in conversational interactions.
This commit is contained in:
+16
-15
@@ -27,9 +27,6 @@ logging.basicConfig(
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DEFAULT_EMBEDDING_PATH = "./data/models/Qwen3-Embedding-0.6B"
|
||||
|
||||
|
||||
def _repo_root() -> Path:
|
||||
"""仓库根目录(含 data/、.env、api_server.py 的目录)。"""
|
||||
return Path(__file__).resolve().parent.parent
|
||||
@@ -41,7 +38,8 @@ def _load_project_env():
|
||||
|
||||
|
||||
def _embedding_model_path() -> str:
|
||||
return os.getenv("EMBEDDING_MODEL_PATH", _DEFAULT_EMBEDDING_PATH).strip()
|
||||
"""仅 ``USE_LOCAL_EMBEDDING=true`` 时需要;远程 Embedding 可为空。"""
|
||||
return os.getenv("EMBEDDING_MODEL_PATH", "").strip()
|
||||
|
||||
|
||||
def resolve_sql_dialect(name: str) -> str:
|
||||
@@ -55,20 +53,23 @@ def resolve_sql_dialect(name: str) -> str:
|
||||
def setup_environment():
|
||||
"""环境检查"""
|
||||
_load_project_env()
|
||||
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "true").strip().lower() in (
|
||||
use_local_emb = os.getenv("USE_LOCAL_EMBEDDING", "false").strip().lower() in (
|
||||
"1", "true", "yes", "on",
|
||||
)
|
||||
if use_local_emb:
|
||||
# 与 .env 中 EMBEDDING_MODEL_PATH 及 utils.embedding 一致
|
||||
model_path = Path(_embedding_model_path())
|
||||
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(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_*(百炼)")
|
||||
logger.warning("本地 Embedding 目录不存在: %s", model_path)
|
||||
logger.info(
|
||||
"请设置正确的 EMBEDDING_MODEL_PATH,或改用 USE_LOCAL_EMBEDDING=false,"
|
||||
"并配置 OPENAI_* / MODELSCOPE_* / DASHSCOPE_* 远程 Embedding"
|
||||
)
|
||||
return False
|
||||
else:
|
||||
ms_key = os.getenv("MODELSCOPE_API_KEY", "").strip()
|
||||
@@ -206,7 +207,7 @@ def single_query(orchestrator, question: str, dialect: str = "tsql"):
|
||||
|
||||
logger.info(f"[Q] 问题: {question}")
|
||||
|
||||
classified = classify_dialog(question)
|
||||
classified = classify_dialog(question, llm_client=orchestrator.deepseek)
|
||||
if classified.intent == DialogIntent.CONVERSATION:
|
||||
reply = classified.reply_suggestion or ""
|
||||
logger.info("[Q] 意图: conversation(跳过 SQL 生成)")
|
||||
|
||||
Reference in New Issue
Block a user