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