Coverage for src/local_deep_research/llm/providers/auto_discovery.py: 98%

105 statements  

« prev     ^ index     » next       coverage.py v7.15.1, created at 2026-07-20 01:24 +0000

1"""Auto-discovery system for OpenAI-compatible providers.""" 

2 

3import importlib 

4import inspect 

5from pathlib import Path 

6from typing import Dict, List, Optional 

7 

8from ...security.secure_logging import logger 

9from .base import BaseLLMProvider, normalize_provider 

10from .openai_base import OpenAICompatibleProvider 

11from ..llm_registry import register_llm 

12from ...security.log_sanitizer import scrub_error 

13 

14 

15class ProviderInfo: 

16 """Information about a discovered provider.""" 

17 

18 def __init__(self, provider_class): 

19 self.provider_class = provider_class 

20 self.provider_key = getattr( 

21 provider_class, 

22 "provider_key", 

23 provider_class.__name__.replace("Provider", "").upper(), 

24 ) 

25 self.provider_name = provider_class.provider_name 

26 self.company_name = getattr( 

27 provider_class, "company_name", provider_class.provider_name 

28 ) 

29 self.is_cloud = getattr(provider_class, "is_cloud", True) 

30 self.requires_auth_for_models = ( 

31 provider_class.requires_auth_for_models() 

32 ) 

33 

34 # Generate display name from attributes 

35 self.display_name = self._generate_display_name() 

36 

37 def _generate_display_name(self): 

38 """Generate a descriptive display name from provider attributes.""" 

39 # Start with the provider name 

40 name_parts = [self.provider_name] 

41 

42 # Add cloud/local indicator 

43 if self.is_cloud is True: 

44 name_parts.append("☁️ Cloud") 

45 elif self.is_cloud is False: 

46 name_parts.append("💻 Local") 

47 

48 return " ".join(name_parts) 

49 

50 def to_dict(self): 

51 """Convert to dictionary for API responses.""" 

52 return { 

53 "value": self.provider_key, 

54 "label": self.display_name, 

55 "is_cloud": self.is_cloud, 

56 } 

57 

58 

59class ProviderDiscovery: 

60 """Discovers and manages OpenAI-compatible providers.""" 

61 

62 _instance = None 

63 _providers: Dict[str, ProviderInfo] = {} 

64 _discovered: bool = False 

65 

66 def __new__(cls): 

67 if cls._instance is None: 

68 cls._instance = super().__new__(cls) 

69 cls._instance._discovered = False 

70 return cls._instance 

71 

72 def discover_providers( 

73 self, force_refresh: bool = False 

74 ) -> Dict[str, ProviderInfo]: 

75 """Discover all providers in the providers directory. 

76 

77 Args: 

78 force_refresh: Force re-discovery even if already done 

79 

80 Returns: 

81 Dictionary mapping provider keys to ProviderInfo objects 

82 """ 

83 if self._discovered and not force_refresh: 

84 return self._providers 

85 

86 self._providers.clear() 

87 # Scan the implementations subdirectory for providers 

88 implementations_dir = Path(__file__).parent / "implementations" 

89 

90 if not implementations_dir.exists(): 

91 logger.warning( 

92 f"Implementations directory not found: {implementations_dir}" 

93 ) 

94 return self._providers 

95 

96 # Scan all Python files in the implementations directory 

97 logger.info(f"Scanning directory: {implementations_dir}") 

98 for file_path in implementations_dir.glob("*.py"): 

99 # Skip special files (like __init__.py) 

100 if file_path.name.startswith("_"): 

101 continue 

102 

103 module_name = file_path.stem 

104 logger.debug(f"Processing module: {module_name} from {file_path}") 

105 try: 

106 # Import the module from implementations subdirectory 

107 module = importlib.import_module( 

108 f".implementations.{module_name}", 

109 package="local_deep_research.llm.providers", 

110 ) 

111 

112 # Find all Provider classes (both OpenAICompatibleProvider and standalone) 

113 logger.debug( 

114 f"Inspecting module {module_name} for Provider classes" 

115 ) 

116 for name, obj in inspect.getmembers(module, inspect.isclass): 

117 if inspect.isclass(obj): 117 ↛ 125line 117 didn't jump to line 125 because the condition on line 117 was always true

118 logger.debug( 

119 f" Found class: {name}, bases: {obj.__bases__}" 

120 ) 

121 # Check if it's a Provider class with a real provider_name. 

122 # BaseLLMProvider now sets a default "unknown" so all 

123 # subclasses inherit the attribute; skip ones that 

124 # haven't customized it (incomplete fixture classes etc.). 

125 if ( 

126 name.endswith("Provider") 

127 and getattr(obj, "provider_name", "unknown") 

128 not in (None, "", "unknown") 

129 and issubclass(obj, BaseLLMProvider) 

130 and obj is not OpenAICompatibleProvider 

131 and obj is not BaseLLMProvider 

132 # Only register classes DEFINED in this module, not ones 

133 # merely imported into its namespace (e.g. 

134 # custom_anthropic_endpoint imports AnthropicProvider as 

135 # its base). Without this, the imported base would be 

136 # re-registered from every module that imports it. 

137 and obj.__module__ == module.__name__ 

138 ): 

139 # Found a provider class 

140 provider_info = ProviderInfo(obj) 

141 self._providers[provider_info.provider_key] = ( 

142 provider_info 

143 ) 

144 

145 # Auto-register the provider directly using the class 

146 register_llm( 

147 normalize_provider(provider_info.provider_key), 

148 obj.create_llm, 

149 ) 

150 logger.info( 

151 f"Auto-registered provider: {provider_info.provider_key}" 

152 ) 

153 

154 logger.info( 

155 f"Discovered provider: {provider_info.provider_key} from {module_name}.py" 

156 ) 

157 

158 except Exception as e: 

159 safe_msg = scrub_error(e) 

160 logger.exception( 

161 f"Error loading provider from {module_name} ({type(e).__name__}): {safe_msg}" 

162 ) 

163 

164 self._discovered = True 

165 logger.info(f"Discovered {len(self._providers)} providers") 

166 return self._providers 

167 

168 def get_provider_info(self, provider_key: str) -> Optional[ProviderInfo]: 

169 """Get information about a specific provider. 

170 

171 Args: 

172 provider_key: The provider key (e.g., 'IONOS', 'GOOGLE') 

173 

174 Returns: 

175 ProviderInfo object or None if not found 

176 """ 

177 if not self._discovered: 

178 self.discover_providers() 

179 return self._providers.get(provider_key.upper()) 

180 

181 def get_provider_options(self) -> List[Dict]: 

182 """Get list of provider options for UI dropdowns. 

183 

184 Returns: 

185 List of dictionaries with 'value' and 'label' keys 

186 """ 

187 if not self._discovered: 

188 self.discover_providers() 

189 

190 options = [] 

191 for provider_info in self._providers.values(): 

192 options.append(provider_info.to_dict()) 

193 

194 # Sort by label 

195 options.sort(key=lambda x: x["label"]) 

196 return options 

197 

198 def get_available_provider_options( 

199 self, settings_snapshot=None 

200 ) -> List[Dict]: 

201 """Get list of available provider options, filtered by availability. 

202 

203 Filters out providers that are not available (e.g., missing API keys). 

204 Useful for contexts where only usable providers should be shown 

205 (e.g., starting a research). For settings/configuration UIs, prefer 

206 get_provider_options() so users can discover and configure new providers. 

207 

208 Args: 

209 settings_snapshot: Settings snapshot for checking provider availability. 

210 Should be provided to correctly check cloud provider API keys. 

211 

212 Returns: 

213 List of dictionaries with 'value' and 'label' keys 

214 """ 

215 if not self._discovered: 215 ↛ 216line 215 didn't jump to line 216 because the condition on line 215 was never true

216 self.discover_providers() 

217 

218 options = [] 

219 for provider_info in self._providers.values(): 

220 if not provider_info.provider_class.is_available( 

221 settings_snapshot=settings_snapshot 

222 ): 

223 logger.debug( 

224 f"Provider {provider_info.provider_key} filtered out " 

225 f"(not available)" 

226 ) 

227 continue 

228 options.append(provider_info.to_dict()) 

229 

230 if not options: 

231 logger.warning( 

232 "No auto-discovered providers passed availability filter. " 

233 "Check that API keys are configured for cloud providers." 

234 ) 

235 

236 # Sort by label 

237 options.sort(key=lambda x: x["label"]) 

238 return options 

239 

240 def get_provider_class(self, provider_key: str): 

241 """Get the provider class for a given key. 

242 

243 Args: 

244 provider_key: The provider key (e.g., 'IONOS', 'GOOGLE') 

245 

246 Returns: 

247 Provider class or None if not found 

248 """ 

249 provider_info = self.get_provider_info(provider_key) 

250 return provider_info.provider_class if provider_info else None 

251 

252 

253# Global instance 

254provider_discovery = ProviderDiscovery() 

255 

256 

257def discover_providers(force_refresh: bool = False) -> Dict[str, ProviderInfo]: 

258 """Discover all available providers. 

259 

260 Args: 

261 force_refresh: Force re-discovery even if already done 

262 

263 Returns: 

264 Dictionary mapping provider keys to ProviderInfo objects 

265 """ 

266 return provider_discovery.discover_providers(force_refresh) 

267 

268 

269def get_discovered_provider_options() -> List[Dict]: 

270 """Get list of discovered provider options for UI dropdowns. 

271 

272 Returns: 

273 List of dictionaries with 'value' and 'label' keys 

274 """ 

275 return provider_discovery.get_provider_options() 

276 

277 

278def get_available_discovered_provider_options( 

279 settings_snapshot=None, 

280) -> List[Dict]: 

281 """Get list of available provider options, filtered by availability. 

282 

283 Only returns providers that pass is_available() check. Useful for 

284 contexts where only usable providers matter (e.g., starting a research). 

285 For settings/configuration UIs, use get_discovered_provider_options() 

286 instead so users can discover and configure new providers. 

287 

288 Args: 

289 settings_snapshot: Settings snapshot for checking provider availability 

290 

291 Returns: 

292 List of dictionaries with 'value' and 'label' keys 

293 """ 

294 return provider_discovery.get_available_provider_options(settings_snapshot) 

295 

296 

297def get_provider_class(provider_key: str): 

298 """Get the provider class for a given key. 

299 

300 Args: 

301 provider_key: The provider key (e.g., 'IONOS', 'GOOGLE') 

302 

303 Returns: 

304 Provider class or None if not found 

305 """ 

306 return provider_discovery.get_provider_class(provider_key)