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