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:
@@ -0,0 +1,130 @@
|
||||
"""
|
||||
从 Lite NL 会话消息构造 Text2SQL 可用的上文摘要,并判断上一轮是否为数据查询(用于续问意图分类)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_MAX_ASSISTANT_SQL_CHARS = 1400
|
||||
_MAX_ASSISTANT_TEXT_CHARS = 900
|
||||
_MAX_BLOCK_CHARS = 7500
|
||||
|
||||
|
||||
def _truncate(s: str, max_len: int) -> str:
|
||||
s = (s or "").strip()
|
||||
if len(s) <= max_len:
|
||||
return s
|
||||
return s[: max_len - 1].rstrip() + "…"
|
||||
|
||||
|
||||
def summarize_assistant_nl_payload(content: str) -> str:
|
||||
"""将落库的 assistant JSON 转为一小段可读摘要。"""
|
||||
raw = (content or "").strip()
|
||||
if not raw:
|
||||
return ""
|
||||
try:
|
||||
data: Dict[str, Any] = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
return _truncate(raw, _MAX_ASSISTANT_TEXT_CHARS)
|
||||
|
||||
br = data.get("branch_result") or {}
|
||||
if isinstance(br.get("answer"), str) and br["answer"].strip():
|
||||
return "[助手] " + _truncate(br["answer"].strip(), _MAX_ASSISTANT_TEXT_CHARS)
|
||||
|
||||
sql = (br.get("sql") or "").strip() if isinstance(br.get("sql"), str) else ""
|
||||
if sql:
|
||||
parts = ["[上轮 SQL] " + _truncate(sql, _MAX_ASSISTANT_SQL_CHARS)]
|
||||
if br.get("follow_up_required"):
|
||||
parts.append("[状态] 上轮为库探针0行,需补充条件后重问")
|
||||
for key in ("db_empty_feedback", "sql_explain"):
|
||||
v = br.get(key)
|
||||
if isinstance(v, str) and v.strip():
|
||||
parts.append("[说明] " + _truncate(v.strip(), 500))
|
||||
break
|
||||
return "\n".join(parts)
|
||||
|
||||
intent = (data.get("intent") or {}).get("intent")
|
||||
return _truncate(f"[助手] intent={intent}", _MAX_ASSISTANT_TEXT_CHARS)
|
||||
|
||||
|
||||
def last_assistant_was_data_query(items: List[Dict[str, Any]]) -> bool:
|
||||
"""最后一条 assistant 消息是否为数据查询分支(含 SQL 或 DATA_QUERY intent)。"""
|
||||
for m in reversed(items or []):
|
||||
if m.get("role") != "assistant":
|
||||
continue
|
||||
raw = (m.get("content") or "").strip()
|
||||
if not raw:
|
||||
return False
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
return False
|
||||
intent = (data.get("intent") or {}).get("intent")
|
||||
if intent == "DATA_QUERY":
|
||||
return True
|
||||
br = data.get("branch_result") or {}
|
||||
sql = (br.get("sql") or "").strip() if isinstance(br.get("sql"), str) else ""
|
||||
return bool(sql)
|
||||
return False
|
||||
|
||||
|
||||
def messages_to_text2sql_context(
|
||||
items: List[Dict[str, Any]],
|
||||
*,
|
||||
max_chars: int = _MAX_BLOCK_CHARS,
|
||||
max_pairs: int = 8,
|
||||
) -> Tuple[str, int]:
|
||||
"""
|
||||
将会话消息列表转为供模型阅读的「上文」文本(不含本轮用户输入)。
|
||||
|
||||
Returns:
|
||||
(context_text, num_user_turns_included)
|
||||
"""
|
||||
if not items:
|
||||
return "", 0
|
||||
|
||||
pairs: List[Tuple[str, str]] = []
|
||||
i = 0
|
||||
n = len(items)
|
||||
while i < n:
|
||||
u = items[i]
|
||||
if u.get("role") != "user":
|
||||
i += 1
|
||||
continue
|
||||
user_text = (u.get("content") or "").strip()
|
||||
asst_text = ""
|
||||
if i + 1 < n and items[i + 1].get("role") == "assistant":
|
||||
asst_text = summarize_assistant_nl_payload(str(items[i + 1].get("content") or ""))
|
||||
i += 2
|
||||
else:
|
||||
i += 1
|
||||
if user_text or asst_text:
|
||||
pairs.append((user_text, asst_text))
|
||||
|
||||
if not pairs:
|
||||
return "", 0
|
||||
|
||||
pairs = pairs[-max_pairs:]
|
||||
lines: List[str] = []
|
||||
total = 0
|
||||
user_count = 0
|
||||
for idx, (uq, aq) in enumerate(pairs, start=1):
|
||||
chunk_parts = [f"--- 第{idx}轮 ---"]
|
||||
if uq:
|
||||
chunk_parts.append(f"用户:{uq}")
|
||||
if aq:
|
||||
chunk_parts.append(f"{aq}")
|
||||
chunk = "\n".join(chunk_parts)
|
||||
if total + len(chunk) + 2 > max_chars:
|
||||
break
|
||||
lines.append(chunk)
|
||||
total += len(chunk) + 2
|
||||
if uq:
|
||||
user_count += 1
|
||||
|
||||
return "\n\n".join(lines).strip(), user_count
|
||||
Reference in New Issue
Block a user