Files
2026-07-24 22:36:04 +03:00

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