Files
nba-fan-hub/entity_linker.py
T

222 lines
9.0 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 -*-
"""
实体链接程序(Entity Linker)—— 需求5专用程序
==============================================
作用:大模型回答生成后,用数据库实体库对回答文本做扫描匹配,
找出其中提到的 球队 / 球员 / 人物 / 比赛 实体,供前端:
a) 特殊标记(高亮 + 点击查看详情)
b) 快速查看入口(回答下方卡片)
设计:
- 纯规则 + 数据库匹配,不额外调用大模型(快、稳、零成本)
- 最长匹配优先("洛杉矶湖人" 优先于 "湖人"
- 重叠区间去重,只保留最长的命中
- 比赛实体优先取对话中工具实际查到的比赛(game_refs),避免误匹配
- 别名(字母哥/SGA 等)复用 tools.ALIASES
"""
import os
import re
from db import query, query_one
from tools import ALIASES
_cache = {"names": None, "games_seen": set()}
def _load_names():
"""加载实体名库:[(名字, 实体信息), ...](按名字长度降序)
自动生成简称:球队去掉城市前缀(波士顿凯尔特人→凯尔特人);
球员中文名取最后一段(谢伊·吉尔杰斯-亚历山大→亚历山大),歧义名跳过。"""
if _cache["names"] is not None:
return _cache["names"]
names = []
for t in query("SELECT id, name, name_en, code, city FROM teams"):
info = {"type": "team", "id": t["id"], "name": t["name"]}
for n in {t["name"], t["name_en"], t["code"]}:
if n and len(str(n)) >= 2:
names.append((str(n), info))
# 简称:去掉与城市名的公共前缀('波士顿凯尔特人'→'凯尔特人''俄克拉荷马雷霆'→'雷霆'
short = t["name"]
if t["city"]:
common = os.path.commonprefix([t["name"], t["city"]])
if len(common) >= 2:
short = t["name"][len(common):]
if len(short) >= 2 and short != t["name"]:
names.append((short, info))
for p in query("SELECT id, name, name_en FROM players"):
info = {"type": "player", "id": p["id"], "name": p["name"]}
for n in {p["name"], p["name_en"]}:
if n and len(str(n)) >= 2:
names.append((str(n), info))
# 简称:中文名按 · 分割取最后一段,再按 - 分割取最后一段
# ('谢伊·吉尔杰斯-亚历山大' → '亚历山大''斯蒂芬·库里' → '库里'
for short in _short_names(p["name"]):
if short:
names.append((short, info))
for p in query("SELECT id, name, name_en FROM persons"):
info = {"type": "person", "id": p["id"], "name": p["name"]}
for n in {p["name"], p["name_en"]}:
if n and len(str(n)) >= 2:
names.append((str(n), info))
for short in _short_names(p["name"]):
if short:
names.append((short, info))
# 别名 → 指向正式实体(球队/球员名里查找)
for nick, real in ALIASES.items():
target = None
for n, info in names:
if info["name"] == real and info["type"] in ("player", "team"):
target = info
break
if target:
names.append((nick, target))
# 去歧义:同一简称指向多个不同实体时,全部移除(避免错误标记)
by_name = {}
for n, info in names:
by_name.setdefault(n, set()).add((info["type"], info["id"]))
names = [(n, info) for n, info in names if len(by_name[n]) == 1]
# 按名字长度降序,保证最长匹配优先
names.sort(key=lambda x: -len(x[0]))
_cache["names"] = names
return names
_ASCII_RE = re.compile(r"[A-Za-z0-9 ._'\-]+")
def _short_names(full):
"""生成中文名简称候选:按 · 和 - 逐级取最后一段
'谢伊·吉尔杰斯-亚历山大' → ['吉尔杰斯-亚历山大', '亚历山大']
'格雷格·波波维奇' → ['波波维奇']
'斯蒂芬·库里' → ['库里']"""
if not full or "·" not in full:
return []
seg = full.split("·")[-1].strip()
out = []
if len(seg) >= 2:
out.append(seg)
if "-" in seg:
tail = seg.split("-")[-1].strip()
if len(tail) >= 2 and tail != seg:
out.append(tail)
return out
def _find_spans(text):
"""在 text 中找出所有实体命中区间。
返回 [{start, end, type, id, name}],已去重叠(保留最长),按 start 升序。"""
if not text:
return []
spans = []
for n, info in _load_names():
# 中文/混合名:直接子串查找(两边不接中文字符,避免部分词误匹配)
if re.search(r"[\u4e00-\u9fff]", n):
for m in re.finditer(re.escape(n), text):
spans.append((m.start(), m.end(), info))
else:
# 英文/缩写:需要词边界
for m in re.finditer(r"(?<![A-Za-z0-9])" + re.escape(n) + r"(?![A-Za-z0-9])", text, re.I):
spans.append((m.start(), m.end(), info))
if not spans:
return []
# 按 (start, -len) 排序 → 贪心取最长且不重叠
spans.sort(key=lambda s: (s[0], -(s[1] - s[0])))
out, last_end = [], -1
for st, en, info in spans:
if st >= last_end:
out.append({"start": st, "end": en, "type": info["type"],
"id": info["id"], "name": info["name"]})
last_end = en
return out
def _match_games(text, game_refs):
"""从工具实际查到的比赛里,找出文本中提到的比赛(A队 vs B队 相邻出现)。
返回卡片数据列表(按球队组合去重,避免系列赛多场重复)。"""
cards, seen, seen_pair = [], set(), set()
for g in game_refs or []:
gid = g.get("id")
if not gid or gid in seen:
continue
away, home = g.get("away_team") or "", g.get("home_team") or ""
if not away or not home:
continue
pair = frozenset((away, home))
if pair in seen_pair: # 同一对球队只留一个卡片(优先已结束/最新)
continue
# 两种顺序都试:'雷霆 4-2 凯尔特人' / '凯尔特人不敌雷霆'
hit = None
for a, b in ((away, home), (home, away)):
pat = re.compile(re.escape(a) + r"[^。\n,;]{0,12}?" + re.escape(b))
m = pat.search(text)
if m:
hit = m
break
if hit:
seen.add(gid)
seen_pair.add(pair)
cards.append({
"type": "game", "id": gid,
"home_team": g.get("home_team"), "away_team": g.get("away_team"),
"home_score": g.get("home_score"), "away_score": g.get("away_score"),
"game_time": g.get("game_time"), "status": g.get("status"),
"round_name": g.get("round_name"),
})
return cards
def link_entities(text, game_refs=None):
"""主入口:扫描回答文本。
返回 (spans, cards)
spans: 实体命中区间(前端高亮标记用),按 start 升序、互不重叠
cards: 快速查看卡片数据(每类限量,避免刷屏)
"""
spans = _find_spans(text or "")
cards = []
seen_cards = set()
for sp in spans:
key = (sp["type"], sp["id"])
if key in seen_cards:
continue
seen_cards.add(key)
info = _card_info(sp["type"], sp["id"])
if info:
cards.append(info)
cards += _match_games(text, game_refs)
# 每类最多 4 个卡片
by_type = {}
for c in cards:
by_type.setdefault(c["type"], []).append(c)
cards = []
for t in ("team", "player", "person", "game"):
cards.extend(by_type.get(t, [])[:4])
return spans, cards
def _card_info(etype, eid):
"""取实体摘要数据(前端卡片直接渲染)"""
try:
if etype == "team":
t = query_one("SELECT * FROM teams WHERE id=?", (eid,))
if not t:
return None
return {"type": "team", "id": t["id"], "name": t["name"], "name_en": t["name_en"],
"city": t["city"], "arena": t["arena"], "champion_count": t["champion_count"],
"head_coach": t["head_coach"]}
if etype == "player":
p = query_one("""SELECT p.*, t.name AS team_name FROM players p
LEFT JOIN teams t ON p.team_id=t.id WHERE p.id=?""", (eid,))
if not p:
return None
return {"type": "player", "id": p["id"], "name": p["name"], "name_en": p["name_en"],
"team": p["team_name"], "position": p["position"], "number": p["number"],
"season": {"pts": p["season_pts"], "reb": p["season_reb"], "ast": p["season_ast"]}}
if etype == "person":
p = query_one("SELECT * FROM persons WHERE id=?", (eid,))
if not p:
return None
return {"type": "person", "id": p["id"], "name": p["name"], "role_cn": p["role_cn"],
"title": p["title"]}
except Exception:
return None
return None