- 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>
578 lines
20 KiB
Python
578 lines
20 KiB
Python
"""群聊人设卡 —— 数据层:仓储
|
||
|
||
只做存储读写,不关心数据来源(群事件、历史导入走同一个入口)。
|
||
业务层与数据来源层都只通过这里的函数访问数据库。
|
||
"""
|
||
|
||
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())
|