123 lines
4.5 KiB
Python
123 lines
4.5 KiB
Python
import asyncio
|
|
import logging
|
|
|
|
from aiogram import Bot, Dispatcher
|
|
from aiogram.client.default import DefaultBotProperties
|
|
from aiogram.enums import ParseMode
|
|
|
|
from app.attachments.service import AttachmentService
|
|
from app.bot.handlers import make_router
|
|
from app.config import Settings
|
|
from app.database.repositories import (
|
|
MemoryRepository,
|
|
PromptRepository,
|
|
RepositoryBindingRepository,
|
|
)
|
|
from app.database.session import make_session_factory
|
|
from app.llm.client import DeepSeekClient
|
|
from app.memory.context_builder import ContextBuilder
|
|
from app.memory.summarizer import Summarizer
|
|
from app.proactive.classifier import MessageClassifier
|
|
from app.repository.service import RepositoryManager, RepositoryService
|
|
from app.services import AnswerService
|
|
|
|
|
|
async def main() -> None:
|
|
settings = Settings()
|
|
logging.basicConfig(
|
|
level=settings.log_level, format="%(asctime)s %(levelname)s %(name)s %(message)s"
|
|
)
|
|
sessions = make_session_factory(settings.database_url)
|
|
llm = DeepSeekClient(
|
|
settings.deepseek_api_key.get_secret_value(),
|
|
settings.deepseek_base_url,
|
|
settings.deepseek_model,
|
|
)
|
|
repository = RepositoryService(
|
|
settings.repository_url,
|
|
settings.repository_branch,
|
|
settings.repository_path,
|
|
settings.repository_access_token.get_secret_value()
|
|
if settings.repository_access_token
|
|
else None,
|
|
settings.max_repository_file_bytes,
|
|
)
|
|
if settings.repository_url:
|
|
try:
|
|
await asyncio.to_thread(repository.sync)
|
|
except Exception:
|
|
logging.getLogger(__name__).exception("initial_repository_sync_failed")
|
|
repository_manager = RepositoryManager(
|
|
settings.repository_path.parent / "repositories",
|
|
settings.repository_access_token.get_secret_value()
|
|
if settings.repository_access_token
|
|
else None,
|
|
settings.max_repository_file_bytes,
|
|
)
|
|
builder = ContextBuilder(settings.context_token_limit)
|
|
dispatcher = Dispatcher()
|
|
|
|
# Use a session-owning facade because each Telegram update needs a separate transaction.
|
|
class Facade:
|
|
def __init__(self):
|
|
self.repository = repository
|
|
|
|
async def repository_for(self, chat_id, thread_id):
|
|
async with sessions() as session:
|
|
binding = await RepositoryBindingRepository(session).active(chat_id, thread_id)
|
|
if binding:
|
|
return repository_manager.service_for(
|
|
binding.url, binding.branch, binding.cache_path
|
|
)
|
|
return repository if settings.repository_url else None
|
|
|
|
async def attach(self, chat_id, thread_id, url, attached_by):
|
|
service = repository_manager.service_for(url)
|
|
commit = await asyncio.to_thread(service.sync)
|
|
async with sessions() as session:
|
|
binding = await RepositoryBindingRepository(session).upsert(
|
|
chat_id=chat_id,
|
|
thread_id=thread_id,
|
|
url=service.url,
|
|
branch=service.branch,
|
|
cache_path=str(service.path),
|
|
commit=commit,
|
|
attached_by=attached_by,
|
|
)
|
|
return binding
|
|
|
|
async def answer(self, chat_id, thread_id, question):
|
|
async with sessions() as session:
|
|
selected_repository = await self.repository_for(chat_id, thread_id)
|
|
return await AnswerService(
|
|
llm,
|
|
MemoryRepository(session),
|
|
PromptRepository(session),
|
|
builder,
|
|
selected_repository,
|
|
).answer(chat_id, thread_id, question)
|
|
|
|
facade = Facade()
|
|
# Handler replaces this repository with a per-update transaction before use.
|
|
summarizer = Summarizer(None, builder, llm, settings.summary_trigger_tokens)
|
|
dispatcher.include_router(
|
|
make_router(
|
|
settings=settings,
|
|
session_factory=sessions,
|
|
answer_service=facade,
|
|
summarizer=summarizer,
|
|
classifier=MessageClassifier(llm),
|
|
attachment_service=AttachmentService(settings.max_attachment_bytes),
|
|
)
|
|
)
|
|
bot = Bot(
|
|
settings.telegram_bot_token.get_secret_value(),
|
|
default=DefaultBotProperties(parse_mode=ParseMode.MARKDOWN),
|
|
)
|
|
logging.getLogger(__name__).info("bot_starting")
|
|
await dispatcher.start_polling(bot, allowed_updates=dispatcher.resolve_used_update_types())
|
|
|
|
|
|
if __name__ == "__main__":
|
|
asyncio.run(main())
|