131 lines
4.0 KiB
Python
131 lines
4.0 KiB
Python
"""
|
|||
|
|
从 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
|