282 lines
12 KiB
Python
282 lines
12 KiB
Python
"""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
|