187 lines
6.3 KiB
Python
187 lines
6.3 KiB
Python
from sqlalchemy import delete, desc, func, select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.database.models import (
|
|
Message,
|
|
PromptVersion,
|
|
RepositoryBinding,
|
|
Setting,
|
|
Summary,
|
|
TelegramUser,
|
|
)
|
|
|
|
|
|
class MemoryRepository:
|
|
def __init__(self, session: AsyncSession):
|
|
self.session = session
|
|
|
|
async def save_message(self, message: Message) -> Message:
|
|
existing = await self.session.scalar(
|
|
select(Message).where(
|
|
Message.chat_id == message.chat_id,
|
|
Message.telegram_message_id == message.telegram_message_id,
|
|
)
|
|
)
|
|
if existing:
|
|
existing.text, existing.edited, existing.edited_at = (
|
|
message.text,
|
|
True,
|
|
message.edited_at,
|
|
)
|
|
result = existing
|
|
else:
|
|
self.session.add(message)
|
|
result = message
|
|
await self.session.commit()
|
|
return result
|
|
|
|
async def context_messages(
|
|
self, chat_id: int, thread_id: int | None, limit: int = 80
|
|
) -> list[Message]:
|
|
stmt = select(Message).where(Message.chat_id == chat_id)
|
|
stmt = (
|
|
stmt.where(Message.thread_id.is_(thread_id))
|
|
if thread_id is None
|
|
else stmt.where(Message.thread_id == thread_id)
|
|
)
|
|
rows = await self.session.scalars(stmt.order_by(desc(Message.sent_at)).limit(limit))
|
|
return list(reversed(rows.all()))
|
|
|
|
async def latest_summary(self, chat_id: int, thread_id: int | None) -> Summary | None:
|
|
stmt = select(Summary).where(Summary.chat_id == chat_id)
|
|
stmt = (
|
|
stmt.where(Summary.thread_id.is_(thread_id))
|
|
if thread_id is None
|
|
else stmt.where(Summary.thread_id == thread_id)
|
|
)
|
|
return await self.session.scalar(stmt.order_by(desc(Summary.version)).limit(1))
|
|
|
|
async def save_summary(self, summary: Summary) -> None:
|
|
self.session.add(summary)
|
|
await self.session.commit()
|
|
|
|
async def delete_personal_memory(self, user_id: int) -> int:
|
|
result = await self.session.execute(delete(Message).where(Message.user_id == user_id))
|
|
await self.session.execute(delete(TelegramUser).where(TelegramUser.telegram_id == user_id))
|
|
await self.session.commit()
|
|
return result.rowcount or 0
|
|
|
|
|
|
class PromptRepository:
|
|
def __init__(self, session: AsyncSession):
|
|
self.session = session
|
|
|
|
async def set(self, text: str, author_id: int) -> PromptVersion:
|
|
await self.session.execute(PromptVersion.__table__.update().values(active=False))
|
|
version = (await self.session.scalar(select(func.max(PromptVersion.version)))) or 0
|
|
prompt = PromptVersion(version=version + 1, text=text, author_id=author_id, active=True)
|
|
self.session.add(prompt)
|
|
await self.session.commit()
|
|
return prompt
|
|
|
|
async def active(self) -> PromptVersion | None:
|
|
return await self.session.scalar(
|
|
select(PromptVersion).where(PromptVersion.active.is_(True))
|
|
)
|
|
|
|
async def history(self, limit: int = 10) -> list[PromptVersion]:
|
|
rows = await self.session.scalars(
|
|
select(PromptVersion).order_by(desc(PromptVersion.version)).limit(limit)
|
|
)
|
|
return list(rows)
|
|
|
|
async def rollback(self) -> PromptVersion | None:
|
|
current = await self.active()
|
|
if not current:
|
|
return None
|
|
previous = await self.session.scalar(
|
|
select(PromptVersion)
|
|
.where(PromptVersion.version < current.version)
|
|
.order_by(desc(PromptVersion.version))
|
|
)
|
|
if not previous:
|
|
return None
|
|
current.active, previous.active = False, True
|
|
await self.session.commit()
|
|
return previous
|
|
|
|
|
|
class SettingsRepository:
|
|
def __init__(self, session: AsyncSession):
|
|
self.session = session
|
|
|
|
async def get(self, key: str, default: str | None = None) -> str | None:
|
|
row = await self.session.get(Setting, key)
|
|
return row.value if row else default
|
|
|
|
async def set(self, key: str, value: str) -> None:
|
|
row = await self.session.get(Setting, key)
|
|
if row:
|
|
row.value = value
|
|
else:
|
|
self.session.add(Setting(key=key, value=value))
|
|
await self.session.commit()
|
|
|
|
|
|
class RepositoryBindingRepository:
|
|
def __init__(self, session: AsyncSession):
|
|
self.session = session
|
|
|
|
async def active(self, chat_id: int, thread_id: int | None) -> RepositoryBinding | None:
|
|
if thread_id is not None:
|
|
binding = await self.session.scalar(
|
|
select(RepositoryBinding).where(
|
|
RepositoryBinding.chat_id == chat_id,
|
|
RepositoryBinding.thread_id == thread_id,
|
|
)
|
|
)
|
|
if binding:
|
|
return binding
|
|
return await self.session.scalar(
|
|
select(RepositoryBinding).where(
|
|
RepositoryBinding.chat_id == chat_id,
|
|
RepositoryBinding.thread_id.is_(None),
|
|
)
|
|
)
|
|
|
|
async def upsert(
|
|
self,
|
|
*,
|
|
chat_id: int,
|
|
thread_id: int | None,
|
|
url: str,
|
|
branch: str,
|
|
cache_path: str,
|
|
commit: str,
|
|
attached_by: int,
|
|
) -> RepositoryBinding:
|
|
stmt = select(RepositoryBinding).where(RepositoryBinding.chat_id == chat_id)
|
|
stmt = (
|
|
stmt.where(RepositoryBinding.thread_id.is_(None))
|
|
if thread_id is None
|
|
else stmt.where(RepositoryBinding.thread_id == thread_id)
|
|
)
|
|
existing = await self.session.scalar(stmt)
|
|
if existing:
|
|
existing.url, existing.branch, existing.cache_path, existing.commit = (
|
|
url,
|
|
branch,
|
|
cache_path,
|
|
commit,
|
|
)
|
|
existing.attached_by = attached_by
|
|
result = existing
|
|
else:
|
|
result = RepositoryBinding(
|
|
chat_id=chat_id,
|
|
thread_id=thread_id,
|
|
url=url,
|
|
branch=branch,
|
|
cache_path=cache_path,
|
|
commit=commit,
|
|
attached_by=attached_by,
|
|
)
|
|
self.session.add(result)
|
|
await self.session.commit()
|
|
return result
|