Create Relay Bot MVP
This commit is contained in:
@@ -0,0 +1,186 @@
|
||||
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
|
||||
Reference in New Issue
Block a user