Files
HeXi/hexi/plugins/nonebot_plugin_galgame_card/repository.py
T
sansenhoshiandClaude b61d09f09f Add HeXi bot codebase: custom plugins, web frontends, tests
- hexi core: message handling, rate limiting, cooldown, plugin manager
- Custom plugins: BF stats, daily check-in, quotes, persona cards, etc.
- Community plugins vendored under hexi/plugins with local fixes
- Web admin frontends (learning-chat, persona-admin), unified hexi/web
- Tests for rate_limit/cooldown/memes/persona; poetry.lock

Co-Authored-By: Claude <noreply@anthropic.com>
2026-09-01 13:13:40 +08:00

578 lines
20 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.
"""群聊人设卡 —— 数据层:仓储
只做存储读写,不关心数据来源(群事件、历史导入走同一个入口)。
业务层与数据来源层都只通过这里的函数访问数据库。
"""
from datetime import datetime
from typing import Optional
from nonebot import require
require("nonebot_plugin_orm")
from nonebot_plugin_orm import get_session
from sqlalchemy import delete, func, select
from .models import (
PersonaChatLog,
PersonaGroup,
PersonaImage,
PersonaImpression,
PersonaSummary,
PersonaUser,
)
# 单用户单群语料滚动保留上限(超出淘汰最旧,DESIGN.md §6)
ROLLING_WINDOW = 3000
# 发言段链条:同一说话人两条消息间隔不超过该秒数,视为同一段发言
FOLLOWS_MAX_GAP = 300 # 5 分钟
# ---------- 管理群列表(数据来源总闸门) ----------
async def ensure_schema() -> None:
"""schema 管理(幂等):补列 + 建新表。
orm 自动同步已被禁用(见 __init__.py),且 orm 的启动建表早于本插件
模型导入(新表不会自动建),因此建表/补列都走这里。
"""
from sqlalchemy import text
from sqlalchemy.exc import OperationalError
statements = [
"ALTER TABLE persona_group ADD COLUMN group_name VARCHAR(64) DEFAULT ''",
"ALTER TABLE persona_chat_log ADD COLUMN follows_id INTEGER",
"ALTER TABLE persona_chat_log ADD COLUMN target_inherited BOOLEAN DEFAULT 0",
"ALTER TABLE persona_chat_log ADD COLUMN image_count INTEGER DEFAULT 0",
"ALTER TABLE persona_chat_log ADD COLUMN image_hashes TEXT",
]
async with get_session() as session:
for stmt in statements:
try:
await session.execute(text(stmt))
await session.commit()
except OperationalError:
pass # 列已存在
# 幂等创建缺失的表(本插件全部模型共用同一 metadata)
def _create_tables(sync_session) -> None:
PersonaImage.metadata.create_all(sync_session.connection())
await session.run_sync(_create_tables)
async def list_groups() -> list[PersonaGroup]:
"""全部被管理的群(数据库直读,不依赖 OneBot)"""
async with get_session() as session:
result = await session.execute(
select(PersonaGroup).order_by(PersonaGroup.group_id)
)
return list(result.scalars().all())
async def add_group(group_id: int, group_name: str = "") -> None:
"""加入管理列表(幂等:已存在则更新名称),默认不采集"""
async with get_session() as session:
row = await session.get(PersonaGroup, group_id)
if row is None:
session.add(PersonaGroup(group_id=group_id, group_name=group_name))
elif group_name and row.group_name != group_name:
row.group_name = group_name
await session.commit()
async def remove_group(group_id: int) -> None:
"""移出管理列表(数据保留,重新加入可续上)"""
async with get_session() as session:
await session.execute(
delete(PersonaGroup).where(PersonaGroup.group_id == group_id)
)
await session.commit()
async def backfill_group_names(names: dict[int, str]) -> None:
"""用 OneBot 群列表回填已管理群的名称(仅在拉取群列表时顺带做)"""
async with get_session() as session:
rows = (await session.execute(select(PersonaGroup))).scalars().all()
changed = False
for row in rows:
name = names.get(row.group_id, "")
if name and row.group_name != name:
row.group_name = name
changed = True
if changed:
await session.commit()
async def set_group_enabled(group_id: int, enabled: bool) -> None:
"""开/关某群的采集(幂等:不存在则创建,存在则更新)"""
async with get_session() as session:
row = await session.get(PersonaGroup, group_id)
if row is None:
session.add(PersonaGroup(group_id=group_id, enabled=enabled))
else:
row.enabled = enabled
row.updated_at = datetime.now()
await session.commit()
async def is_group_enabled(group_id: int) -> bool:
async with get_session() as session:
result = await session.execute(
select(PersonaGroup).where(PersonaGroup.group_id == group_id)
)
row = result.scalar()
return bool(row and row.enabled)
async def enabled_groups() -> list[int]:
"""全部开启采集的群(调度/采集层扫描用)"""
async with get_session() as session:
result = await session.execute(
select(PersonaGroup.group_id).where(PersonaGroup.enabled.is_(True))
)
return [row[0] for row in result.all()]
# ---------- 参与者 ----------
async def is_joined(user_id: int, group_id: int) -> bool:
async with get_session() as session:
result = await session.execute(
select(PersonaUser).where(
PersonaUser.user_id == user_id, PersonaUser.group_id == group_id
)
)
return result.scalar() is not None
async def join(user_id: int, group_id: int) -> None:
async with get_session() as session:
session.add(PersonaUser(user_id=user_id, group_id=group_id))
await session.commit()
async def leave(user_id: int, group_id: int) -> None:
async with get_session() as session:
await session.execute(
delete(PersonaUser).where(
PersonaUser.user_id == user_id, PersonaUser.group_id == group_id
)
)
await session.commit()
async def joined_users(group_id: int) -> list[int]:
"""某群全部参与者(调度层扫描用)"""
async with get_session() as session:
result = await session.execute(
select(PersonaUser.user_id).where(PersonaUser.group_id == group_id)
)
return [row[0] for row in result.all()]
# ---------- 语料 ----------
async def add_chat_log(
user_id: int,
group_id: int,
content: str,
nickname: str = "",
target_user_id: Optional[int] = None,
reply_to_content: Optional[str] = None,
image_count: int = 0,
image_hashes: Optional[str] = None,
) -> int:
"""落库一条语料,返回新语料 id
发言段链条:同一说话人的上一条语料记为 follows_id(间隔 ≤ FOLLOWS_MAX_GAP);
本条无显式目标(回复/@)时从链条继承 target_user_id——对上面那句的解释/补充
仍视为发给同一对象,不丢语境。
"""
async with get_session() as session:
prev = (
await session.execute(
select(PersonaChatLog)
.where(
PersonaChatLog.user_id == user_id,
PersonaChatLog.group_id == group_id,
)
.order_by(PersonaChatLog.id.desc())
.limit(1)
)
).scalar()
follows_id = None
inherited = False
if prev is not None and (
datetime.now() - prev.created_at
).total_seconds() <= FOLLOWS_MAX_GAP:
follows_id = prev.id
if target_user_id is None:
target_user_id = prev.target_user_id
inherited = target_user_id is not None
log = PersonaChatLog(
user_id=user_id,
group_id=group_id,
nickname=nickname,
content=content,
target_user_id=target_user_id,
target_inherited=inherited,
follows_id=follows_id,
image_count=image_count,
image_hashes=image_hashes,
reply_to_content=reply_to_content,
)
session.add(log)
await session.commit()
await session.refresh(log)
return log.id
async def count_chat_log(user_id: int, group_id: int) -> int:
async with get_session() as session:
result = await session.execute(
select(func.count(PersonaChatLog.id)).where(
PersonaChatLog.user_id == user_id, PersonaChatLog.group_id == group_id
)
)
return result.scalar_one()
async def latest_chat_logs(
user_id: int, group_id: int, limit: int, before_id: Optional[int] = None
) -> list[PersonaChatLog]:
"""按 id 倒序取最新 limit 条;before_id 用于翻页"""
stmt = (
select(PersonaChatLog)
.where(PersonaChatLog.user_id == user_id, PersonaChatLog.group_id == group_id)
.order_by(PersonaChatLog.id.desc())
.limit(limit)
)
if before_id is not None:
stmt = stmt.where(PersonaChatLog.id <= before_id)
async with get_session() as session:
result = await session.execute(stmt)
return list(result.scalars().all())
async def chat_logs_since(
user_id: int, group_id: int, since_id: int, limit: int
) -> list[PersonaChatLog]:
"""按 id 正序取 id > since_id 的语料(增量印象生成的输入)"""
async with get_session() as session:
result = await session.execute(
select(PersonaChatLog)
.where(
PersonaChatLog.user_id == user_id,
PersonaChatLog.group_id == group_id,
PersonaChatLog.id > since_id,
)
.order_by(PersonaChatLog.id.asc())
.limit(limit)
)
return list(result.scalars().all())
async def prune_chat_log(user_id: int, group_id: int, keep: int = ROLLING_WINDOW) -> int:
"""淘汰滚动窗口外的旧语料,返回删除条数"""
async with get_session() as session:
result = await session.execute(
select(PersonaChatLog.id)
.where(
PersonaChatLog.user_id == user_id, PersonaChatLog.group_id == group_id
)
.order_by(PersonaChatLog.id.desc())
.offset(keep)
)
stale_ids = [row[0] for row in result.all()]
if stale_ids:
await session.execute(
delete(PersonaChatLog).where(PersonaChatLog.id.in_(stale_ids))
)
await session.commit()
return len(stale_ids)
# ---- 统计与分页(Web 管理后台用) ----
async def count_chat_log_group(group_id: int) -> int:
async with get_session() as session:
result = await session.execute(
select(func.count(PersonaChatLog.id)).where(
PersonaChatLog.group_id == group_id
)
)
return result.scalar_one()
async def count_impressions_group(group_id: int) -> int:
async with get_session() as session:
result = await session.execute(
select(func.count(PersonaImpression.id)).where(
PersonaImpression.group_id == group_id
)
)
return result.scalar_one()
async def count_summaries_group(group_id: int) -> int:
async with get_session() as session:
result = await session.execute(
select(func.count(PersonaSummary.id)).where(
PersonaSummary.group_id == group_id
)
)
return result.scalar_one()
async def corpus_counts_by_group() -> dict[int, int]:
"""各群语料条数(一次聚合查询,替代逐群 count)"""
async with get_session() as session:
result = await session.execute(
select(PersonaChatLog.group_id, func.count(PersonaChatLog.id)).group_by(
PersonaChatLog.group_id
)
)
return {int(row[0]): int(row[1]) for row in result.all()}
async def impression_counts_by_group() -> dict[int, int]:
async with get_session() as session:
result = await session.execute(
select(PersonaImpression.group_id, func.count(PersonaImpression.id)).group_by(
PersonaImpression.group_id
)
)
return {int(row[0]): int(row[1]) for row in result.all()}
async def summary_counts_by_group() -> dict[int, int]:
async with get_session() as session:
result = await session.execute(
select(PersonaSummary.group_id, func.count(PersonaSummary.id)).group_by(
PersonaSummary.group_id
)
)
return {int(row[0]): int(row[1]) for row in result.all()}
async def participant_counts_by_group() -> dict[int, int]:
async with get_session() as session:
result = await session.execute(
select(PersonaUser.group_id, func.count(PersonaUser.user_id)).group_by(
PersonaUser.group_id
)
)
return {int(row[0]): int(row[1]) for row in result.all()}
async def chat_log_stats_by_user(group_id: int) -> list[tuple[int, int]]:
"""群内每人语料条数((user_id, count),按条数倒序)"""
async with get_session() as session:
result = await session.execute(
select(PersonaChatLog.user_id, func.count(PersonaChatLog.id))
.where(PersonaChatLog.group_id == group_id)
.group_by(PersonaChatLog.user_id)
.order_by(func.count(PersonaChatLog.id).desc())
)
return [(int(row[0]), int(row[1])) for row in result.all()]
async def list_logs(
group_id: int,
user_id: Optional[int] = None,
offset: int = 0,
limit: int = 50,
) -> list[PersonaChatLog]:
"""语料分页(倒序,新在前)"""
stmt = select(PersonaChatLog).where(PersonaChatLog.group_id == group_id)
if user_id is not None:
stmt = stmt.where(PersonaChatLog.user_id == user_id)
stmt = stmt.order_by(PersonaChatLog.id.desc()).offset(offset).limit(limit)
async with get_session() as session:
result = await session.execute(stmt)
return list(result.scalars().all())
async def count_logs_filtered(group_id: int, user_id: Optional[int] = None) -> int:
stmt = select(func.count(PersonaChatLog.id)).where(
PersonaChatLog.group_id == group_id
)
if user_id is not None:
stmt = stmt.where(PersonaChatLog.user_id == user_id)
async with get_session() as session:
result = await session.execute(stmt)
return result.scalar_one()
async def list_impressions_group(
group_id: int, offset: int = 0, limit: int = 50
) -> list[PersonaImpression]:
async with get_session() as session:
result = await session.execute(
select(PersonaImpression)
.where(PersonaImpression.group_id == group_id)
.order_by(PersonaImpression.id.desc())
.offset(offset)
.limit(limit)
)
return list(result.scalars().all())
async def list_summaries_group(group_id: int) -> list[PersonaSummary]:
async with get_session() as session:
result = await session.execute(
select(PersonaSummary)
.where(PersonaSummary.group_id == group_id)
.order_by(PersonaSummary.version.desc())
)
return list(result.scalars().all())
async def delete_logs(group_id: int, user_id: Optional[int] = None) -> int:
"""删除语料(可只删某人),返回删除条数"""
stmt = delete(PersonaChatLog).where(PersonaChatLog.group_id == group_id)
if user_id is not None:
stmt = stmt.where(PersonaChatLog.user_id == user_id)
async with get_session() as session:
result = await session.execute(stmt)
await session.commit()
return result.rowcount or 0
async def clear_group_data(group_id: int) -> None:
"""清空某群全部数据:语料 + 印象 + 画像 + 参与者(群开关行保留)"""
async with get_session() as session:
await session.execute(delete(PersonaChatLog).where(PersonaChatLog.group_id == group_id))
await session.execute(delete(PersonaImpression).where(PersonaImpression.group_id == group_id))
await session.execute(delete(PersonaSummary).where(PersonaSummary.group_id == group_id))
await session.execute(delete(PersonaUser).where(PersonaUser.group_id == group_id))
await session.commit()
# ---------- 印象 ----------
async def add_impression(
user_id: int,
group_id: int,
content: str,
cover_from_id: int,
cover_to_id: int,
model: str = "",
) -> None:
async with get_session() as session:
session.add(
PersonaImpression(
user_id=user_id,
group_id=group_id,
content=content,
cover_from_id=cover_from_id,
cover_to_id=cover_to_id,
model=model,
)
)
await session.commit()
async def latest_impression(user_id: int, group_id: int) -> Optional[PersonaImpression]:
async with get_session() as session:
result = await session.execute(
select(PersonaImpression)
.where(
PersonaImpression.user_id == user_id,
PersonaImpression.group_id == group_id,
)
.order_by(PersonaImpression.id.desc())
.limit(1)
)
return result.scalar()
async def count_impressions(user_id: int, group_id: int) -> int:
async with get_session() as session:
result = await session.execute(
select(func.count(PersonaImpression.id)).where(
PersonaImpression.user_id == user_id,
PersonaImpression.group_id == group_id,
)
)
return result.scalar_one()
async def list_impressions(
user_id: int, group_id: int, limit: int
) -> list[PersonaImpression]:
"""按 id 正序取最近 limit 条印象(画像生成的输入,最旧在前)"""
async with get_session() as session:
result = await session.execute(
select(PersonaImpression)
.where(
PersonaImpression.user_id == user_id,
PersonaImpression.group_id == group_id,
)
.order_by(PersonaImpression.id.desc())
.limit(limit)
)
return list(reversed(result.scalars().all()))
# ---------- 画像 ----------
async def latest_summary(user_id: int, group_id: int) -> Optional[PersonaSummary]:
async with get_session() as session:
result = await session.execute(
select(PersonaSummary)
.where(
PersonaSummary.user_id == user_id,
PersonaSummary.group_id == group_id,
)
.order_by(PersonaSummary.version.desc())
.limit(1)
)
return result.scalar()
async def add_summary(
user_id: int,
group_id: int,
card_text: str,
corpus_count: int = 0,
impression_count: int = 0,
model: str = "",
) -> int:
"""新增画像版本,version 自动 +1,返回新版本号。
注意:并发调用可能撞 UNIQUE(user_id, group_id, version),由调度层加锁保护。
"""
prev = await latest_summary(user_id, group_id)
version = (prev.version + 1) if prev else 1
async with get_session() as session:
session.add(
PersonaSummary(
user_id=user_id,
group_id=group_id,
version=version,
card_text=card_text,
corpus_count=corpus_count,
impression_count=impression_count,
model=model,
)
)
await session.commit()
return version
async def list_summaries(user_id: int, group_id: int) -> list[PersonaSummary]:
"""全部历史版本(版本对比用),旧版在前"""
async with get_session() as session:
result = await session.execute(
select(PersonaSummary)
.where(
PersonaSummary.user_id == user_id,
PersonaSummary.group_id == group_id,
)
.order_by(PersonaSummary.version.asc())
)
return list(result.scalars().all())