104 lines
3.2 KiB
Python
104 lines
3.2 KiB
Python
"""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()) |