Create Relay Bot MVP
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
from app.database.models import Base
|
||||
|
||||
__all__ = ["Base"]
|
||||
@@ -0,0 +1,143 @@
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import (
|
||||
JSON,
|
||||
BigInteger,
|
||||
Boolean,
|
||||
DateTime,
|
||||
Float,
|
||||
Integer,
|
||||
String,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
func,
|
||||
)
|
||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
pass
|
||||
|
||||
|
||||
class TelegramUser(Base):
|
||||
__tablename__ = "telegram_users"
|
||||
telegram_id: Mapped[int] = mapped_column(BigInteger, primary_key=True)
|
||||
display_name: Mapped[str] = mapped_column(String(255), default="")
|
||||
username: Mapped[str | None] = mapped_column(String(255))
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), server_default=func.now())
|
||||
|
||||
|
||||
class AllowedChat(Base):
|
||||
__tablename__ = "allowed_chats"
|
||||
chat_id: Mapped[int] = mapped_column(BigInteger, primary_key=True)
|
||||
kind: Mapped[str] = mapped_column(String(32), default="group")
|
||||
enabled: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
|
||||
|
||||
class Message(Base):
|
||||
__tablename__ = "messages"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("chat_id", "telegram_message_id", name="uq_message_chat_id"),
|
||||
)
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
chat_id: Mapped[int] = mapped_column(BigInteger, index=True)
|
||||
telegram_message_id: Mapped[int] = mapped_column(BigInteger)
|
||||
thread_id: Mapped[int | None] = mapped_column(BigInteger, index=True)
|
||||
user_id: Mapped[int | None] = mapped_column(BigInteger, index=True)
|
||||
user_name: Mapped[str] = mapped_column(String(255), default="")
|
||||
text: Mapped[str] = mapped_column(Text, default="")
|
||||
sent_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True)
|
||||
reply_to_message_id: Mapped[int | None] = mapped_column(BigInteger)
|
||||
edited: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
edited_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True))
|
||||
message_type: Mapped[str] = mapped_column(String(32), default="text")
|
||||
attachment_data: Mapped[dict[str, Any] | None] = mapped_column(JSON)
|
||||
|
||||
|
||||
class Summary(Base):
|
||||
__tablename__ = "summaries"
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
chat_id: Mapped[int] = mapped_column(BigInteger, index=True)
|
||||
thread_id: Mapped[int | None] = mapped_column(BigInteger, index=True)
|
||||
version: Mapped[int] = mapped_column(Integer)
|
||||
text: Mapped[str] = mapped_column(Text)
|
||||
from_message_id: Mapped[int] = mapped_column(BigInteger)
|
||||
to_message_id: Mapped[int] = mapped_column(BigInteger)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), server_default=func.now())
|
||||
|
||||
|
||||
class PromptVersion(Base):
|
||||
__tablename__ = "prompt_versions"
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
version: Mapped[int] = mapped_column(Integer, unique=True)
|
||||
text: Mapped[str] = mapped_column(Text)
|
||||
author_id: Mapped[int] = mapped_column(BigInteger)
|
||||
active: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), server_default=func.now())
|
||||
|
||||
|
||||
class Setting(Base):
|
||||
__tablename__ = "settings"
|
||||
key: Mapped[str] = mapped_column(String(100), primary_key=True)
|
||||
value: Mapped[str] = mapped_column(Text)
|
||||
|
||||
|
||||
class RepositoryState(Base):
|
||||
__tablename__ = "repository_state"
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
branch: Mapped[str] = mapped_column(String(255))
|
||||
commit: Mapped[str] = mapped_column(String(64))
|
||||
synced_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), server_default=func.now())
|
||||
|
||||
|
||||
class RepositoryBinding(Base):
|
||||
"""The active read-only repository for a chat or a forum topic."""
|
||||
|
||||
__tablename__ = "repository_bindings"
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
chat_id: Mapped[int] = mapped_column(BigInteger, index=True)
|
||||
thread_id: Mapped[int | None] = mapped_column(BigInteger, index=True)
|
||||
url: Mapped[str] = mapped_column(String(2_000))
|
||||
branch: Mapped[str] = mapped_column(String(255), default="HEAD")
|
||||
cache_path: Mapped[str] = mapped_column(String(2_000), unique=True)
|
||||
commit: Mapped[str] = mapped_column(String(64))
|
||||
attached_by: Mapped[int] = mapped_column(BigInteger)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), server_default=func.now())
|
||||
synced_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), server_default=func.now())
|
||||
|
||||
|
||||
class ProactiveReply(Base):
|
||||
__tablename__ = "proactive_replies"
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
chat_id: Mapped[int] = mapped_column(BigInteger, index=True)
|
||||
source_message_id: Mapped[int] = mapped_column(BigInteger, index=True)
|
||||
answer_message_id: Mapped[int | None] = mapped_column(BigInteger)
|
||||
fingerprint: Mapped[str] = mapped_column(String(64), index=True)
|
||||
confidence: Mapped[float] = mapped_column(Float)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), server_default=func.now())
|
||||
|
||||
|
||||
class Attachment(Base):
|
||||
__tablename__ = "attachments"
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
chat_id: Mapped[int] = mapped_column(BigInteger, index=True)
|
||||
message_id: Mapped[int] = mapped_column(BigInteger)
|
||||
uploaded_by: Mapped[int] = mapped_column(BigInteger)
|
||||
filename: Mapped[str] = mapped_column(String(512))
|
||||
mime_type: Mapped[str | None] = mapped_column(String(255))
|
||||
size: Mapped[int] = mapped_column(Integer)
|
||||
mode: Mapped[str] = mapped_column(String(32), default="pending")
|
||||
extracted_text: Mapped[str | None] = mapped_column(Text)
|
||||
accepted: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
|
||||
|
||||
class Correction(Base):
|
||||
__tablename__ = "corrections"
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
chat_id: Mapped[int] = mapped_column(BigInteger, index=True)
|
||||
source_message_id: Mapped[int] = mapped_column(BigInteger)
|
||||
text: Mapped[str] = mapped_column(Text)
|
||||
author_id: Mapped[int] = mapped_column(BigInteger)
|
||||
confirmed: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), server_default=func.now())
|
||||
@@ -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
|
||||
@@ -0,0 +1,6 @@
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
|
||||
|
||||
def make_session_factory(database_url: str) -> async_sessionmaker[AsyncSession]:
|
||||
engine = create_async_engine(database_url, pool_pre_ping=True)
|
||||
return async_sessionmaker(engine, expire_on_commit=False)
|
||||
Reference in New Issue
Block a user