Files
nba-fan-hub/chat.py
T

325 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- 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]