"""KnowledgeBase 知识库 Repository。""" from sqlalchemy import func, select from sqlalchemy.orm import Session from app.models.document import Document from app.models.knowledge_base import KnowledgeBase class KnowledgeBaseRepository: def __init__(self, session: Session) -> None: self._session = session def get_by_id(self, kb_id: str) -> KnowledgeBase | None: return self._session.get(KnowledgeBase, kb_id) def get_by_token_hash(self, token_hash: str) -> KnowledgeBase | None: stmt = select(KnowledgeBase).where(KnowledgeBase.token_hash == token_hash) return self._session.scalars(stmt).first() def list_by_user( self, user_id: str, *, page: int = 1, page_size: int = 20 ) -> tuple[list[KnowledgeBase], int]: """分页获取用户的知识库列表。返回 (items, total)。""" count_stmt = ( select(func.count()) .select_from(KnowledgeBase) .where( KnowledgeBase.user_id == user_id, KnowledgeBase.status != "DELETED", ) ) total = self._session.scalar(count_stmt) or 0 stmt = ( select(KnowledgeBase) .where( KnowledgeBase.user_id == user_id, KnowledgeBase.status != "DELETED", ) .order_by(KnowledgeBase.created_at.desc()) .offset((page - 1) * page_size) .limit(page_size) ) items = list(self._session.scalars(stmt).all()) return items, total def create( self, *, user_id: str, name: str, description: str | None, token_hash: str, token_encrypted: str, token_hint: str, ) -> KnowledgeBase: kb = KnowledgeBase( user_id=user_id, name=name, description=description, enabled=True, status="active", token_hash=token_hash, token_encrypted=token_encrypted, token_hint=token_hint, ) self._session.add(kb) self._session.flush() return kb def update(self, kb: KnowledgeBase, **fields) -> None: for key, value in fields.items(): if value is not None: setattr(kb, key, value) self._session.flush() def delete(self, kb: KnowledgeBase) -> None: """软删除。""" kb.status = "DELETED" self._session.flush() def count_documents(self, kb_id: str) -> int: stmt = ( select(func.count()) .select_from(Document) .where( Document.knowledge_base_id == kb_id, Document.status != "DELETED", ) ) return self._session.scalar(stmt) or 0 def get_document_count_map(self, user_id: str) -> dict[str, int]: """批量获取用户所有知识库的文档数量(避免 N+1)。""" stmt = ( select(Document.knowledge_base_id, func.count()) .where( Document.user_id == user_id, Document.status != "DELETED", ) .group_by(Document.knowledge_base_id) ) return dict(self._session.execute(stmt).all())