Files
nba-fan-hub/entity_linker.py
T

231 lines
9.5 KiB
Python
Raw Normal View History

# -*- 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, mark_mode="first"):
"""主入口:扫描回答文本。
返回 (spans, cards)
spans: 实体命中区间(前端高亮标记用),按 start 升序、互不重叠
mark_mode="first" 时同一实体(type,id)只保留首次出现;"all" 时全部标记
cards: 快速查看卡片数据(每类限量,避免刷屏)
"""
spans = _find_spans(text or "")
if mark_mode == "first":
seen, kept = set(), []
for sp in spans: # spans 已按 start 升序 → 保留首次
key = (sp["type"], sp["id"])
if key in seen:
continue
seen.add(key)
kept.append(sp)
spans = kept
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