""" 从 Lite NL 会话消息构造 Text2SQL 可用的上文摘要,并判断上一轮是否为数据查询(用于续问意图分类)。 """ from __future__ import annotations import json import logging import re 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 # 续问/指代常见触发词:用于判定是否需要注入会话上文,避免多话题串扰。 # 规则应偏保守:宁可把新话题当作“无上文”也不要带入无关上下文污染 SQL。 _FOLLOW_UP_PATTERN = re.compile( r"(再|继续|同样|沿用|还是|照旧|刚才|上面|上一条|上一个|上次|上述|这个SQL|这条SQL|按上面|按刚才|在此基础上|同一口径|同口径|按之前|改成|改为|加上|加一下|筛一下|过滤一下|补充|补一下|追加|优化一下|调整一下|换成|换为)", re.IGNORECASE, ) def is_likely_follow_up(user_text: str) -> bool: """ 仅基于用户本轮文本做“是否续问”的轻量判断。 True -> 允许注入少量上文(最近 1~2 轮),用于指代消解/沿用口径 False -> 新话题,禁用上文,避免上下文污染 """ t = (user_text or "").strip() if not t: return False # 很短的“再来一个/同上/继续”等通常是续问 if len(t) <= 10 and _FOLLOW_UP_PATTERN.search(t): return True return bool(_FOLLOW_UP_PATTERN.search(t)) 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