Coverage for src/asketmc_bot/main.py: 0%
121 statements
« prev ^ index » next coverage.py v7.5.1, created at 2026-04-08 17:14 +0000
« prev ^ index » next coverage.py v7.5.1, created at 2026-04-08 17:14 +0000
1#!/usr/bin/env python3
2"""Asketmc RAG Bot — async entry point.
4Responsibilities:
5- Load .env and configuration.
6- Configure logging.
7- Initialize index, reranker, and LLM client.
8- Start Discord bot asynchronously.
9- Gracefully shut down resources on SIGINT/SIGTERM or KeyboardInterrupt.
10"""
12from __future__ import annotations
14import asyncio
15import logging
16import os
17import signal
18import sys
19from pathlib import Path
20from typing import Any
22from dotenv import load_dotenv
23from llama_index.core import VectorStoreIndex
25# Local imports (src-layout safe)
26from asketmc_bot import config as cfg
27from asketmc_bot import discord_bot as discord_module
28from asketmc_bot.index_builder import build_index
29from asketmc_bot.lemma import LEMMA_POOL, extract_lemmas
30from asketmc_bot.llm_client import LLMClient, LLMConfig
31from asketmc_bot.rag_filter import build_context, get_filtered_nodes
32from asketmc_bot.rerank import init_reranker, rerank, shutdown_reranker
35# ── Logging ───────────────────────────────────────────────────────────────────
36def setup_logging(debug: bool = False) -> None:
37 """Configure structured logging."""
38 level = logging.DEBUG if debug else logging.INFO
39 logging.basicConfig(
40 level=level,
41 format="%(asctime)s %(levelname)-8s %(name)s: %(message)s",
42 datefmt="%Y-%m-%d %H:%M:%S",
43 )
46# ── Config & Settings ─────────────────────────────────────────────────────────
47def load_settings() -> dict:
48 """Load .env and config values with safe defaults."""
49 current = Path(__file__).resolve()
50 if current.parent.name == "asketmc_bot":
51 # .../<repo>/src/asketmc_bot/main.py -> parents[2] == <repo>
52 project_root = current.parents[2]
53 else:
54 project_root = current.parent
56 candidates = [
57 project_root / ".env",
58 project_root / ".env.local",
59 project_root / ".env-example",
60 ]
62 env_loaded = False
63 for env_path in candidates:
64 if env_path.exists():
65 load_dotenv(env_path, override=False)
66 env_loaded = True
67 break
68 if not env_loaded:
69 print(f"[WARN] No .env found in {project_root}", file=sys.stderr)
71 required = ["DISCORD_TOKEN", "OPENROUTER_API_KEY"]
72 missing = [v for v in required if not os.getenv(v)]
73 if missing:
74 sys.exit(f"Missing required env vars: {', '.join(missing)}")
76 return {
77 "discord_token": os.getenv("DISCORD_TOKEN"),
78 "openrouter_api_key": os.getenv("OPENROUTER_API_KEY"),
79 "api_url": getattr(cfg, "API_URL", "https://openrouter.ai/api/v1/chat/completions"),
80 "or_model": getattr(cfg, "OR_MODEL", "openrouter/auto"),
81 "or_max_tokens": int(getattr(cfg, "OR_MAX_TOKENS", 512)),
82 "ollama_url": getattr(cfg, "OLLAMA_URL", "http://localhost:11434/api/generate"),
83 "local_model": getattr(cfg, "LOCAL_MODEL", "qwen2.5:7b-instruct-q4_K_M"),
84 "top_k": int(getattr(cfg, "TOP_K", 16)),
85 "ctx_len_remote": int(getattr(cfg, "CTX_LEN_REMOTE", 20_000)),
86 "ctx_len_local": int(getattr(cfg, "CTX_LEN_LOCAL", 12_000)),
87 "http_conn_limit": int(getattr(cfg, "HTTP_CONN_LIMIT", 5)),
88 "or_retries": int(getattr(cfg, "OR_RETRIES", 3)),
89 "http_timeout_total": int(getattr(cfg, "HTTP_TIMEOUT_TOTAL", 240)),
90 "breaker_base_block_sec": int(getattr(cfg, "OPENROUTER_BLOCK_SEC", 120)),
91 "breaker_max_block_sec": int(getattr(cfg, "OPENROUTER_BLOCK_MAX_SEC", 900)),
92 }
95def make_llm_config(s: dict) -> LLMConfig:
96 """Convert settings dict into LLMConfig."""
97 return LLMConfig(
98 api_url=s["api_url"],
99 or_model=s["or_model"],
100 or_max_tokens=s["or_max_tokens"],
101 openrouter_api_key=s["openrouter_api_key"],
102 ollama_url=s["ollama_url"],
103 local_model=s["local_model"],
104 http_conn_limit=s["http_conn_limit"],
105 or_retries=s["or_retries"],
106 http_timeout_total=s["http_timeout_total"],
107 breaker_base_block_sec=s["breaker_base_block_sec"],
108 breaker_max_block_sec=s["breaker_max_block_sec"],
109 )
112# ── Core RAG logic ────────────────────────────────────────────────────────────
113async def generate_rag_answer(
114 retriever: Any,
115 query: str,
116 sys_prompt: str,
117 llm_client: LLMClient,
118 settings: dict,
119 **kwargs: Any,
120) -> str:
121 """Build a RAG prompt and generate an answer via the LLM."""
122 use_remote = bool(kwargs.get("use_remote", True))
124 qlem = extract_lemmas(query)
125 raw_nodes = await retriever.aretrieve(query)
126 reranked_nodes = await rerank(query, raw_nodes)
127 nodes = await get_filtered_nodes(reranked_nodes or raw_nodes, qlem)
129 if not nodes:
130 return "⚠️ Not enough data."
132 char_limit = settings["ctx_len_remote"] if use_remote else settings["ctx_len_local"]
133 ctx_txt = build_context(nodes, qlem, char_limit)
135 if not use_remote:
136 prompt_text = (
137 f"{sys_prompt.strip()}\n\nCONTEXT:\n{ctx_txt.strip()}\n\nQUESTION: {query.strip()}\nANSWER:"
138 )
139 return await llm_client.call_local_llm(prompt_text)
141 text, _used_fallback = await llm_client.query_model(
142 sys_prompt=sys_prompt,
143 ctx_txt=ctx_txt,
144 q=query,
145 )
146 return text
149# ── Application lifecycle ─────────────────────────────────────────────────────
150async def main() -> None:
151 """Main async entry point."""
152 settings = load_settings()
153 setup_logging(getattr(cfg, "DEBUG", False))
154 log = logging.getLogger("asketmc.main")
156 log.info("Building document index...")
157 index: VectorStoreIndex = await build_index()
158 retriever = index.as_retriever(similarity_top_k=settings["top_k"])
160 await init_reranker()
161 llm = LLMClient(make_llm_config(settings), logger=logging.getLogger("asketmc.llm"))
163 stop_event = asyncio.Event()
164 loop = asyncio.get_running_loop()
165 llm.attach_loop(loop)
167 shutdown_started = False
169 async def shutdown() -> None:
170 """Graceful async shutdown for core services."""
171 nonlocal shutdown_started
172 if shutdown_started:
173 return
174 shutdown_started = True
176 log.info("Shutting down core systems...")
177 try:
178 await shutdown_reranker()
179 finally:
180 try:
181 await llm.close()
182 finally:
183 LEMMA_POOL.shutdown(wait=True)
184 log.info("Shutdown complete.")
185 stop_event.set()
187 async def query_model_text(prompt: str) -> tuple[str, bool]:
188 """Simple text query wrapper for the bot."""
189 text, used_fallback = await llm.query_model(messages=[{"role": "user", "content": prompt}])
190 return text, used_fallback
192 def is_openrouter_blocked_fn() -> bool:
193 """Synchronous breaker check for bot throttling decisions."""
194 return llm.is_remote_blocked_sync()
196 bot_task = asyncio.create_task(
197 discord_module.start_bot_async(
198 token=settings["discord_token"],
199 index=index,
200 retriever=retriever,
201 generate_rag_answer=lambda q, p, **kw: generate_rag_answer(
202 retriever, q, p, llm, settings, **kw
203 ),
204 query_model=query_model_text,
205 call_local_llm=llm.call_local_llm,
206 build_index=build_index,
207 is_openrouter_blocked=is_openrouter_blocked_fn,
208 on_core_shutdown=shutdown,
209 )
210 )
212 def _bot_done(t: asyncio.Task) -> None:
213 try:
214 _ = t.result()
215 except asyncio.CancelledError:
216 return
217 except Exception:
218 asyncio.create_task(shutdown())
220 bot_task.add_done_callback(_bot_done)
222 def _signal_handler() -> None:
223 asyncio.create_task(shutdown())
225 for sig_name in ("SIGINT", "SIGTERM"):
226 if hasattr(signal, sig_name):
227 try:
228 loop.add_signal_handler(getattr(signal, sig_name), _signal_handler)
229 except NotImplementedError:
230 pass
232 await stop_event.wait()
234 if not bot_task.done():
235 bot_task.cancel()
236 try:
237 await bot_task
238 except asyncio.CancelledError:
239 pass
242def cli() -> None:
243 """Console-script entrypoint (sync wrapper for async main())."""
244 if sys.platform.startswith("win"):
245 try:
246 asyncio.set_event_loop_policy(asyncio.WindowsSelectorEventLoopPolicy())
247 except (AttributeError, NotImplementedError):
248 pass
250 try:
251 asyncio.run(main())
252 except KeyboardInterrupt:
253 logging.getLogger("asketmc.main").warning("Interrupted by user.")
256if __name__ == "__main__":
257 cli()