Files

578 lines
20 KiB
Python
Raw Permalink Normal View History

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