Files
nba-fan-hub/chat.py
T

324 lines
15 KiB
Python
Raw Normal View History

# -*- coding: utf-8 -*-
"""
对话管线(RAG 混合检索 + 函数调用):
1. 系统提示词 + 用户问题 → DeepSeek(带 8 个工具)
2. 模型选择工具 → 执行(SQLite 结构化查询 + Chroma 向量检索)
3. 工具结果回填 → DeepSeek 生成最终答案(数据准确,标注来源)
4. 兜底:LLM 或工具异常时,用关键词预检索注入上下文后直接回答
"""
import json
import logging
import re
import llm
import tools as tools_mod
from db import query
log = logging.getLogger("chat")
SYSTEM_PROMPT = """你是「NBA球迷大全」智能助手,为球迷提供比赛、球员、球队、新闻、人物等准确信息。
工作准则:
1. 用户问到时事、数据、赛程类问题,必须先调用工具查询数据库,用工具返回的真实数据回答,严禁编造具体比分、数据、日期。
2. 工具结果就是权威数据源。回答时引用关键数据(比分、时间、数据),并标注来源(如「数据库」「新闻」)。
3. 多步问题可以连续调用多个工具(如先查球队,再查该队比赛)。比较多个球员/球队时,把每个名字拆开单独调用一次工具(例如“约基奇和字母哥谁强”应分别调用 search_players("约基奇") 与 search_players("字母哥"))。
4. 查不到时如实说"数据库暂未收录",可以基于常识补充介绍,但要明确区分"数据库数据"与"常识补充"。
5. 回答使用中文,简洁有条理,可适当使用小标题或列表;不要啰嗦。
6. 严禁在回答正文中输出 tool_calls、XML 或函数调用代码——需要数据时直接调用工具函数,或直接如实回答;调用了工具就用工具返回的数据作答。
7. 当前赛季为 2025-26 赛季,总决赛已于 2026年6月结束,雷霆 4-2 击败凯尔特人夺冠。"""
def _tool_result_to_text(name, result):
"""把工具结果压缩成给模型的文本(控制 token 量)"""
if not result or result.get("_source") == "error":
return f"[工具 {name}] 查询失败:{result.get('error','未知错误') if result else '无结果'}"
results = result.get("results", [])
if not results:
return f"[工具 {name}] 未找到相关记录。"
lines = [f"[工具 {name}] 查询到 {len(results)} 条记录:"]
for r in results[:8]:
lines.append(json.dumps(r, ensure_ascii=False)[:600])
return "\n".join(lines)
def _grounding_context(user_msg):
"""兜底检索:用关键词在数据库里快速找相关记录,注入上下文"""
ctx = []
for name, fn in (("teams", tools_mod.search_teams), ("players", tools_mod.search_players),
("games", tools_mod.search_games), ("news", tools_mod.search_news),
("persons", tools_mod.search_persons)):
try:
r = fn(user_msg, limit=3)
if r.get("results"):
ctx.append(_tool_result_to_text(name, r))
except Exception:
continue
return "\n".join(ctx)
def chat_once(user_msg, history=None):
"""单轮对话。返回 (reply, sources, used_tools, news_refs)
sources : 供前端展示的信息来源卡片
news_refs : 本次回答用到的新闻/百科资讯列表(前端折叠展示为链接)
"""
return _chat_impl(user_msg, history)
def _chat_impl(user_msg, history=None):
history = history or []
messages = [{"role": "system", "content": SYSTEM_PROMPT}]
for h in history[-10:]:
messages.append({"role": "user", "content": h.get("user", "")})
if h.get("assistant"):
messages.append({"role": "assistant", "content": h["assistant"]})
messages.append({"role": "user", "content": user_msg})
used_tools, sources, news_refs = [], [], []
# ---- 第 1 轮:带工具
try:
resp = llm.chat(messages, tools=tools_mod.TOOLS)
except Exception as e:
log.warning("LLM 首轮失败(%s),走兜底管线", e)
ctx = _grounding_context(user_msg)
msgs = messages + ([{"role": "system", "content": f"以下是数据库检索到的可能相关信息(未直接命中时勿强行引用):\n{ctx}"}] if ctx else [])
try:
resp2 = llm.chat(msgs)
return llm.parse_content(resp2), _mk_sources(ctx), [t for t in ("grounding",) if ctx], []
except Exception as e2:
return (f"抱歉,大模型服务暂时不可用({e2})。你可以稍后再试,或直接浏览下方数据页面。", [], [], [])
# ---- 工具执行(单轮)+ 实体覆盖补全 → 最终无工具作答
all_executed = []
calls = llm.extract_tool_calls(resp)
if not calls:
calls = _parse_text_tool_calls(llm.parse_content(resp), user_msg)
if not calls:
return llm.parse_content(resp) or "(模型未返回内容)", [], [], []
executed = []
for c in calls:
if c["name"] in [e[0] for e in executed]:
continue
result = tools_mod.run_tool(c["name"], c["arguments"])
log.info("[tools] %s(%s) -> %d条", c["name"],
json.dumps(c["arguments"], ensure_ascii=False)[:120], len(result.get("results", [])))
executed.append((c["name"], c["id"], c["arguments"], result))
used_tools.append(c["name"])
if result.get("results"):
sources.append({"tool": c["name"], "items": result["results"][:3]})
if c["name"] == "search_news":
for it in result["results"][:5]:
if it.get("id"):
news_refs.append({"id": it["id"], "title": it.get("title", ""),
"source": it.get("source", ""),
"publish_time": it.get("publish_time", "")})
elif c["name"] == "get_game_detail" and isinstance(result, dict) and result.get("id"):
sources.append({"tool": "get_game_detail", "items": [result]})
all_executed += executed
# 覆盖补全:用户问题里提到的其他实体(多球员/多球队对比)自动补查,避免模型漏调
covered = set()
for _n, _cid, _args, result in executed:
for r in result.get("results", []):
covered.add(r.get("name") or r.get("team") or "")
for term in tools_mod._split_terms(user_msg)[:4]:
for fn, key in ((tools_mod.search_players, "name"), (tools_mod.search_teams, "name")):
try:
r = fn(term, limit=3)
except Exception:
continue
for row in r.get("results", []):
nm = row.get(key) or row.get("name") or ""
if nm and nm not in covered:
cid = f"cover_{key}_{len(all_executed)}"
executed.append(("search_players" if key == "name" else "search_teams", cid, {"query": term}, r))
covered.add(nm)
used_tools.append("search_players" if key == "name" else "search_teams")
if r.get("results"):
sources.append({"tool": "search_players" if key == "name" else "search_teams", "items": r["results"][:3]})
break
# 新闻引用去重(按 id
seen_nid, news_refs_u = set(), []
for n in news_refs:
if n["id"] not in seen_nid:
seen_nid.add(n["id"])
news_refs_u.append(n)
news_refs = news_refs_u
messages.append({"role": "assistant", "content": None,
"reasoning_content": llm.extract_reasoning(resp),
"tool_calls": [
{"id": cid, "type": "function",
"function": {"name": name, "arguments": json.dumps(args, ensure_ascii=False)}}
for name, cid, args, _ in executed]})
for name, cid, _args, result in executed:
messages.append({"role": "tool", "tool_call_id": cid, "content": _tool_result_to_text(name, result)})
try:
resp = llm.chat(messages) # 最终轮:不带工具,让模型基于数据作答
except Exception as e:
log.warning("LLM 最终轮失败(%s),返回工具结果摘要", e)
return _finalize(all_executed, user_msg, messages, sources, used_tools), sources, used_tools, news_refs
reply = llm.parse_content(resp) or ""
if not reply or "<tool_calls>" in reply or "search_" in reply:
# 模型又输出工具调用文本 → 注入上下文再答一次
return _finalize(all_executed, user_msg, messages, sources, used_tools), sources, used_tools, news_refs
return reply, sources, used_tools, news_refs
def _finalize(executed, user_msg, messages, sources, used_tools):
"""工具循环结束后:把检索结果注入上下文,让模型再总结一次(无工具),失败则给原始摘要"""
if not executed:
return "抱歉,没有检索到相关信息。请换个问法,或直接浏览下方数据页面。"
ctx = "\n\n".join(_tool_result_to_text(name, result) for name, _cid, _args, result in executed)
msgs = [{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": user_msg},
{"role": "assistant", "content": f"我已经查询了数据库,检索结果如下:\n{ctx}\n\n请基于以上数据回答用户的问题(不要提及'工具',直接给出答案;数据不足时如实说明)。"}]
try:
r = llm.chat(msgs)
return llm.parse_content(r) or _results_summary(executed)
except Exception:
return _results_summary(executed)
def _results_summary(executed):
"""把已执行工具的结果整理成给用户的摘要文本"""
parts = ["以下是数据库查到的相关信息:"]
for name, _cid, _args, result in executed:
parts.append(_tool_result_to_text(name, result))
return "\n".join(parts)
def _parse_text_tool_calls(content, user_msg=""):
"""解析模型以正文形式输出的工具调用(兜底),返回与 extract_tool_calls 同构的列表
支持:真实工具名 / 模型臆造的工具名(get_game_stats、search_player_stats 等)→ 映射到最接近的真实工具"""
if not content:
return []
names = re.findall(r"[a-z_]+_[a-z_]+", content)
real = [n for n in names if n in tools_mod.TOOL_HANDLERS]
game_id = None
m = re.search(r"game_id[\"':=>]+\s*(\d+)", content)
if m:
game_id = int(m.group(1))
out = []
if real:
for i, n in enumerate(dict.fromkeys(real)):
q = re.findall(r"[\"'“”]([^\"'“”]{1,80})[\"'“”]", content)
qv = q[i] if i < len(q) else user_msg
args = {"query": qv} if n != "get_game_detail" else {"game_id": game_id or 0}
out.append({"name": n, "arguments": args, "id": f"text_{i}"})
return out
# 模型臆造的工具名 → 智能映射
game_intent = bool(re.search(r"总决赛|比赛|技术统计|统计|G\d|第.场|比分|对位", user_msg or ""))
if game_id and any("game" in n or "stat" in n for n in names):
out.append({"name": "get_game_detail", "arguments": {"game_id": game_id}, "id": "text_g"})
elif game_intent and any("stat" in n or "score" in n or "game" in n for n in names):
out.append({"name": "search_games", "arguments": {"query": user_msg}, "id": "text_g2"})
elif any("player" in n or "stat" in n for n in names):
out.append({"name": "search_players", "arguments": {"query": user_msg}, "id": "text_p"})
elif any("game" in n or "match" in n for n in names):
out.append({"name": "search_games", "arguments": {"query": user_msg}, "id": "text_g2"})
elif any("news" in n for n in names):
out.append({"name": "search_news", "arguments": {"query": user_msg}, "id": "text_n"})
return out
def _mk_sources(ctx):
return [{"tool": "grounding", "items": [{"note": "关键词预检索上下文"}]}] if ctx else []
def suggest_questions():
"""快捷问题(前端展示用,从站点配置读取,管理后台可编辑)"""
try:
from admin import get_config
raw = get_config().get("suggestions") or ""
arr = json.loads(raw)
if isinstance(arr, list) and arr:
return [str(x).strip() for x in arr if str(x).strip()]
except Exception:
pass
return DEFAULT_SUGGESTIONS
DEFAULT_SUGGESTIONS = [
"最近一场比赛结果",
"湖人本赛季战绩怎么样",
"库里本赛季场均数据",
"2026年总决赛谁赢了",
"SGA拿了什么荣誉",
"NBA工资帽是什么",
"介绍一下波波维奇",
"今天有什么新闻",
"西部排名",
"雷霆和凯尔特人总决赛G6数据",
]
def boot_info():
"""对话界面启动信息:开场白 + 快捷问题(均可后台配置)"""
try:
from admin import get_config
cfg = get_config()
except Exception:
cfg = {}
return {
"site_name": cfg.get("site_name", "NBA球迷大全"),
"site_subtitle": cfg.get("site_subtitle", "比赛 · 球员 · 球队 · 资讯 · 人物 · 百科"),
"welcome_text": cfg.get("welcome_text", "你好,我是**NBA球迷大全**助手!可以问我任何关于比赛、球员、球队、新闻、人物的问题,我会基于数据库给你准确答案~"),
"welcome_hint": cfg.get("welcome_hint", "试试:"),
"suggestions": suggest_questions(),
"footer_text": cfg.get("footer_text", "NBA球迷大全 · 数据为模拟演示数据(2025-26 赛季)"),
}
def _parse_json_array(text):
"""从 LLM 输出中解析 JSON 数组(容错:直接 JSON / 提取中括号段)"""
if not text:
return []
text = text.strip()
try:
arr = json.loads(text)
if isinstance(arr, list):
return arr
except Exception:
pass
m = re.search(r"\[.*\]", text, re.S)
if m:
try:
arr = json.loads(m.group(0))
if isinstance(arr, list):
return arr
except Exception:
pass
return []
def predict_suggestions(history=None, n=3):
"""基于对话历史,让大模型预测用户接下来最可能追问的 n 个问题(底部快捷语句)。
每个问题不超过 30 字;LLM 异常时回退到默认快捷问题。"""
n = max(1, min(int(n or 3), 6))
history = history or []
msgs = [{"role": "system", "content": (
f"你是「NBA球迷大全」智能助手。根据对话历史,站在用户角度预测他接下来最可能追问的{n}个问题。\n"
"要求:\n"
"1. 每个问题不超过30个汉字,简洁口语化\n"
"2. 必须是用户会直接发送的提问,不要编号、不要引号、不要解释\n"
"3. 只输出JSON数组,例如:[\"库里今天拿了几分\",\"湖人下一场什么时候\"],不要输出任何其他内容")}]
for h in history[-6:]:
msgs.append({"role": "user", "content": h.get("user", "")})
if h.get("assistant"):
msgs.append({"role": "assistant", "content": str(h["assistant"])[:600]})
if len(msgs) == 1:
return DEFAULT_SUGGESTIONS[:n]
try:
resp = llm.chat(msgs, temperature=0.9, max_tokens=200)
arr = _parse_json_array(llm.parse_content(resp))
out = [str(x).strip()[:30] for x in arr if str(x).strip()][:n]
if out:
return out
except Exception as e:
log.warning("快捷问题预测失败(%s),回退默认", e)
return DEFAULT_SUGGESTIONS[:n]