Files
ai-g3sb-backman2.0/backend/utils/dialog_context.py
T

131 lines
4.0 KiB
Python
Raw Normal View History

"""
从 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