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

1#!/usr/bin/env python3 

2"""Asketmc RAG Bot — async entry point. 

3 

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""" 

11 

12from __future__ import annotations 

13 

14import asyncio 

15import logging 

16import os 

17import signal 

18import sys 

19from pathlib import Path 

20from typing import Any 

21 

22from dotenv import load_dotenv 

23from llama_index.core import VectorStoreIndex 

24 

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 

33 

34 

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 ) 

44 

45 

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 

55 

56 candidates = [ 

57 project_root / ".env", 

58 project_root / ".env.local", 

59 project_root / ".env-example", 

60 ] 

61 

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) 

70 

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)}") 

75 

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 } 

93 

94 

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 ) 

110 

111 

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)) 

123 

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) 

128 

129 if not nodes: 

130 return "⚠️ Not enough data." 

131 

132 char_limit = settings["ctx_len_remote"] if use_remote else settings["ctx_len_local"] 

133 ctx_txt = build_context(nodes, qlem, char_limit) 

134 

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) 

140 

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 

147 

148 

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") 

155 

156 log.info("Building document index...") 

157 index: VectorStoreIndex = await build_index() 

158 retriever = index.as_retriever(similarity_top_k=settings["top_k"]) 

159 

160 await init_reranker() 

161 llm = LLMClient(make_llm_config(settings), logger=logging.getLogger("asketmc.llm")) 

162 

163 stop_event = asyncio.Event() 

164 loop = asyncio.get_running_loop() 

165 llm.attach_loop(loop) 

166 

167 shutdown_started = False 

168 

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 

175 

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() 

186 

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 

191 

192 def is_openrouter_blocked_fn() -> bool: 

193 """Synchronous breaker check for bot throttling decisions.""" 

194 return llm.is_remote_blocked_sync() 

195 

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 ) 

211 

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()) 

219 

220 bot_task.add_done_callback(_bot_done) 

221 

222 def _signal_handler() -> None: 

223 asyncio.create_task(shutdown()) 

224 

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 

231 

232 await stop_event.wait() 

233 

234 if not bot_task.done(): 

235 bot_task.cancel() 

236 try: 

237 await bot_task 

238 except asyncio.CancelledError: 

239 pass 

240 

241 

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 

249 

250 try: 

251 asyncio.run(main()) 

252 except KeyboardInterrupt: 

253 logging.getLogger("asketmc.main").warning("Interrupted by user.") 

254 

255 

256if __name__ == "__main__": 

257 cli()