"""SQLAlchemy LLMWiki metadata repository.""" from __future__ import annotations from datetime import UTC, datetime from typing import Any from uuid import uuid4 from sqlalchemy import delete, or_, select from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.llmwiki.base import LlmWikiScope, LlmWikiStore from deerflow.persistence.llmwiki.model import LlmWikiConversationDepositRow, LlmWikiKnowledgeBaseRow def _now() -> datetime: return datetime.now(UTC) def _to_dict(row: LlmWikiKnowledgeBaseRow | LlmWikiConversationDepositRow) -> dict[str, Any]: data = row.to_dict() for key in ("published_at", "external_search_updated_at", "created_at", "updated_at"): value = data.get(key) if isinstance(value, datetime): data[key] = value.isoformat() return data class LlmWikiRepository(LlmWikiStore): def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory async def create_mapping(self, *, weknora_id: str, owner_user_id: str, name: str, description: str, kb_type: str) -> dict[str, Any]: row = LlmWikiKnowledgeBaseRow( id=str(uuid4()), weknora_id=weknora_id, owner_user_id=owner_user_id, name=name, description=description, kb_type=kb_type, publication_status="private", created_at=_now(), updated_at=_now(), ) async with self._sf() as session: session.add(row) await session.commit() return _to_dict(row) async def list_visible(self, user_id: str, *, scope: LlmWikiScope = "all", is_admin: bool = False) -> list[dict[str, Any]]: stmt = select(LlmWikiKnowledgeBaseRow) if scope == "personal" and not is_admin: stmt = stmt.where(LlmWikiKnowledgeBaseRow.owner_user_id == user_id) elif scope == "public": stmt = stmt.where(LlmWikiKnowledgeBaseRow.publication_status == "published") elif scope == "all" and not is_admin: stmt = stmt.where( or_( LlmWikiKnowledgeBaseRow.owner_user_id == user_id, LlmWikiKnowledgeBaseRow.publication_status == "published", ) ) stmt = stmt.order_by(LlmWikiKnowledgeBaseRow.created_at.desc()) async with self._sf() as session: result = await session.execute(stmt) return [_to_dict(row) for row in result.scalars()] async def get_authorized(self, mapping_id: str, user_id: str, *, write: bool, is_admin: bool) -> dict[str, Any] | None: stmt = select(LlmWikiKnowledgeBaseRow).where(LlmWikiKnowledgeBaseRow.id == mapping_id) if not is_admin: condition = LlmWikiKnowledgeBaseRow.owner_user_id == user_id if not write: condition = or_(condition, LlmWikiKnowledgeBaseRow.publication_status == "published") stmt = stmt.where(condition) async with self._sf() as session: row = (await session.execute(stmt)).scalar_one_or_none() return _to_dict(row) if row else None async def get_by_weknora_id(self, weknora_id: str) -> dict[str, Any] | None: stmt = select(LlmWikiKnowledgeBaseRow).where(LlmWikiKnowledgeBaseRow.weknora_id == weknora_id) async with self._sf() as session: row = (await session.execute(stmt)).scalar_one_or_none() return _to_dict(row) if row else None async def update_mapping( self, mapping_id: str, *, name: str | None = None, description: str | None = None, wiki_index_enabled: bool | None = None, external_search_enabled: bool | None = None, ) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(LlmWikiKnowledgeBaseRow, mapping_id) if row is None: return None if name is not None: row.name = name if description is not None: row.description = description if wiki_index_enabled is not None: row.wiki_index_enabled = wiki_index_enabled if external_search_enabled is not None: row.external_search_enabled = external_search_enabled row.external_search_updated_at = _now() row.updated_at = _now() await session.commit() return _to_dict(row) async def adopt_conversation_deposit_mapping( self, mapping_id: str, *, name: str, description: str, ) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(LlmWikiKnowledgeBaseRow, mapping_id) if row is None: return None now = _now() row.owner_user_id = "system" row.name = name row.description = description row.publication_status = "published" row.published_by = "system" row.reviewed_by = None row.published_at = row.published_at or now row.updated_at = now await session.commit() return _to_dict(row) async def delete_mapping(self, mapping_id: str) -> bool: async with self._sf() as session: row = await session.get(LlmWikiKnowledgeBaseRow, mapping_id) if row is None: return False await session.delete(row) await session.commit() return True async def request_publish(self, mapping_id: str, actor_user_id: str, *, is_admin: bool) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(LlmWikiKnowledgeBaseRow, mapping_id) if row is None or (not is_admin and row.owner_user_id != actor_user_id): return None row.publication_status = "published" row.published_by = actor_user_id row.reviewed_by = None row.published_at = _now() row.updated_at = _now() await session.commit() return _to_dict(row) async def review_publish(self, mapping_id: str, reviewer_user_id: str, *, approved: bool) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(LlmWikiKnowledgeBaseRow, mapping_id) if row is None or row.publication_status != "pending": return None row.publication_status = "published" if approved else "rejected" row.reviewed_by = reviewer_user_id row.published_at = _now() if approved else None row.updated_at = _now() await session.commit() return _to_dict(row) async def unpublish(self, mapping_id: str, actor_user_id: str, *, is_admin: bool) -> dict[str, Any] | None: async with self._sf() as session: row = await session.get(LlmWikiKnowledgeBaseRow, mapping_id) if row is None or (not is_admin and row.owner_user_id != actor_user_id): return None row.publication_status = "private" if is_admin: row.reviewed_by = actor_user_id row.published_at = None row.updated_at = _now() await session.commit() return _to_dict(row) async def list_publications(self, *, status: str | None = None) -> list[dict[str, Any]]: stmt = select(LlmWikiKnowledgeBaseRow).where(LlmWikiKnowledgeBaseRow.publication_status != "private") if status: stmt = stmt.where(LlmWikiKnowledgeBaseRow.publication_status == status) stmt = stmt.order_by(LlmWikiKnowledgeBaseRow.updated_at.desc()) async with self._sf() as session: result = await session.execute(stmt) return [_to_dict(row) for row in result.scalars()] async def get_conversation_deposit_mapping(self) -> dict[str, Any] | None: stmt = ( select(LlmWikiKnowledgeBaseRow) .where( LlmWikiKnowledgeBaseRow.owner_user_id == "system", LlmWikiKnowledgeBaseRow.name == "对话沉淀", LlmWikiKnowledgeBaseRow.publication_status == "published", ) .order_by(LlmWikiKnowledgeBaseRow.created_at.asc()) .limit(1) ) async with self._sf() as session: row = (await session.execute(stmt)).scalar_one_or_none() return _to_dict(row) if row else None async def get_conversation_deposit( self, *, thread_id: str, assistant_message_id: str, ) -> dict[str, Any] | None: stmt = select(LlmWikiConversationDepositRow).where( LlmWikiConversationDepositRow.thread_id == thread_id, LlmWikiConversationDepositRow.assistant_message_id == assistant_message_id, ) async with self._sf() as session: row = (await session.execute(stmt)).scalar_one_or_none() return _to_dict(row) if row else None async def upsert_conversation_deposit( self, *, thread_id: str, human_message_id: str | None, assistant_message_id: str, knowledge_base_mapping_id: str, weknora_knowledge_id: str | None, title: str, created_by_user_id: str | None, ) -> dict[str, Any]: now = _now() async with self._sf() as session: stmt = select(LlmWikiConversationDepositRow).where( LlmWikiConversationDepositRow.thread_id == thread_id, LlmWikiConversationDepositRow.assistant_message_id == assistant_message_id, ) row = (await session.execute(stmt)).scalar_one_or_none() if row is None: row = LlmWikiConversationDepositRow( id=str(uuid4()), thread_id=thread_id, human_message_id=human_message_id, assistant_message_id=assistant_message_id, knowledge_base_mapping_id=knowledge_base_mapping_id, weknora_knowledge_id=weknora_knowledge_id, title=title, created_by_user_id=created_by_user_id, created_at=now, updated_at=now, ) session.add(row) else: row.human_message_id = human_message_id row.knowledge_base_mapping_id = knowledge_base_mapping_id row.weknora_knowledge_id = weknora_knowledge_id row.title = title row.updated_at = now await session.commit() return _to_dict(row) async def delete_conversation_deposits( self, *, thread_id: str, message_ids: list[str] | None = None, ) -> list[dict[str, Any]]: conditions = [LlmWikiConversationDepositRow.thread_id == thread_id] if message_ids: conditions.append( or_( LlmWikiConversationDepositRow.human_message_id.in_(message_ids), LlmWikiConversationDepositRow.assistant_message_id.in_(message_ids), ) ) stmt = select(LlmWikiConversationDepositRow).where(*conditions) async with self._sf() as session: rows = (await session.execute(stmt)).scalars().all() deleted = [_to_dict(row) for row in rows] if rows: await session.execute(delete(LlmWikiConversationDepositRow).where(LlmWikiConversationDepositRow.id.in_([row.id for row in rows]))) await session.commit() return deleted