Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-llm/src/lexigram/ai/llm/model_manager/manager.py: 22%

99 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-25 07:19 +0800

1"""Unified LLM model manager.""" 

2 

3from __future__ import annotations 

4 

5import asyncio 

6from typing import Any 

7 

8from lexigram.ai.llm.model_manager.base import AbstractModelManager 

9from lexigram.ai.llm.model_manager.providers import ( 

10 LMStudioModelManager, 

11 OllamaModelManager, 

12 VLLMModelManager, 

13) 

14from lexigram.ai.llm.model_manager.types import ModelLoadResult 

15from lexigram.logging import ( 

16 get_logger, 

17) 

18 

19logger = get_logger(__name__) 

20 

21 

22class LLMModelManager: 

23 """Unified model manager that routes to specific provider managers.""" 

24 

25 def __init__(self, active_provider: str | None = None): 

26 """Initialize unified model manager.""" 

27 if getattr(self, "_initialized", False): 

28 return 

29 

30 self.managers: dict[str, AbstractModelManager] = {} 

31 self.active_provider = active_provider 

32 

33 # Register default providers (using classes defined in this module) 

34 self.register_provider("ollama", OllamaModelManager()) 

35 self.register_provider("lm-studio", LMStudioModelManager()) 

36 self.register_provider("vllm", VLLMModelManager()) 

37 

38 self._initialized = True 

39 

40 def register_provider(self, provider: str, manager: AbstractModelManager) -> None: 

41 """Register a model manager for a provider.""" 

42 self.managers[provider] = manager 

43 logger.debug("Registered model manager for provider: %s", provider) 

44 

45 def unregister_provider(self, provider: str) -> None: 

46 """Unregister a model manager.""" 

47 if provider in self.managers: 

48 from lexigram.concurrency import TaskManager 

49 

50 task_mgr = TaskManager() 

51 task_mgr.create_background_task( 

52 self.managers[provider].close(), 

53 name=f"provider_shutdown_{provider}", 

54 ) 

55 del self.managers[provider] 

56 logger.info("Unregistered model manager for provider: %s", provider) 

57 

58 async def switch_provider(self, provider: str) -> bool: 

59 """Switch to a different provider and unload models from other providers.""" 

60 if provider not in self.managers: 

61 logger.error("Provider %s not registered", provider) 

62 return False 

63 

64 # Unload models from other providers 

65 await self._unload_other_providers(provider) 

66 

67 self.active_provider = provider 

68 logger.info("Switched to provider: %s", provider) 

69 return True 

70 

71 async def _unload_other_providers(self, keep_provider: str) -> None: 

72 """Unload all models from providers other than the specified one.""" 

73 for provider_name in self.managers: 

74 if provider_name != keep_provider: 

75 await self._unload_all_models(provider_name) 

76 

77 async def _unload_all_models(self, provider: str) -> None: 

78 """Unload all models from a provider.""" 

79 if provider not in self.managers: 

80 return 

81 

82 # Try to get loaded models and unload them specifically 

83 loaded_models = await self.managers[provider].get_loaded_models() 

84 logger.info( 

85 "Found %d loaded models in provider %s: %s", 

86 len(loaded_models), 

87 provider, 

88 loaded_models, 

89 ) 

90 

91 for model in loaded_models: 

92 await self.managers[provider].unload_model(model) 

93 

94 # Check what models remain after unloading 

95 remaining_models = await self.managers[provider].get_loaded_models() 

96 logger.info( 

97 "After unloading, %d models remain in provider %s: %s", 

98 len(remaining_models), 

99 provider, 

100 remaining_models, 

101 ) 

102 

103 # For providers that don't track loaded models well, we still attempt cleanup 

104 # This is especially important for local LLMs with limited GPU memory 

105 logger.info("Unloaded all models from provider: %s", provider) 

106 

107 async def list_models(self, provider: str | None = None) -> list[dict[str, Any]]: 

108 """List models for a provider.""" 

109 target_provider = provider or self.active_provider 

110 if not target_provider or target_provider not in self.managers: 

111 return [] 

112 

113 return await self.managers[target_provider].list_models() 

114 

115 async def load_model( 

116 self, 

117 model_name: str, 

118 provider: str | None = None, 

119 **kwargs: Any, 

120 ) -> ModelLoadResult: 

121 """Load a model.""" 

122 target_provider = provider or self.active_provider 

123 if not target_provider or target_provider not in self.managers: 

124 logger.error("No provider available for loading model %s", model_name) 

125 return ModelLoadResult( 

126 success=False, 

127 model_name=model_name, 

128 error=f"No provider available for loading model {model_name}", 

129 retryable=False, 

130 ) 

131 

132 # Explicitly unload models from other providers before loading new model 

133 logger.info( 

134 "Unloading models from other providers before loading %s in %s", 

135 model_name, 

136 target_provider, 

137 ) 

138 await self._unload_other_providers(target_provider) 

139 

140 # Add a small pause to ensure unloading takes effect 

141 await asyncio.sleep(0.5) 

142 

143 # If switching providers, update active provider 

144 if provider and provider != self.active_provider: 

145 self.active_provider = provider 

146 

147 logger.info("Loading model %s in provider %s", model_name, target_provider) 

148 result = await self.managers[target_provider].load_model(model_name, **kwargs) 

149 

150 if result.success: 

151 logger.info("Successfully loaded model %s", model_name) 

152 else: 

153 logger.error("Failed to load model %s: %s", model_name, result.error) 

154 

155 return result 

156 

157 async def unload_model(self, model_name: str, provider: str | None = None) -> bool: 

158 """Unload a model.""" 

159 target_provider = provider or self.active_provider 

160 if not target_provider or target_provider not in self.managers: 

161 return False 

162 

163 return await self.managers[target_provider].unload_model(model_name) 

164 

165 async def switch_model( 

166 self, 

167 model_name: str, 

168 provider: str | None = None, 

169 **kwargs: Any, 

170 ) -> ModelLoadResult: 

171 """Switch to a different model.""" 

172 target_provider = provider or self.active_provider 

173 if not target_provider or target_provider not in self.managers: 

174 return ModelLoadResult( 

175 success=False, 

176 model_name=model_name, 

177 error=f"Provider {target_provider!r} not registered", 

178 ) 

179 

180 # Explicitly unload models from other providers before switching 

181 logger.info( 

182 "Unloading models from other providers before switching to %s in %s", 

183 model_name, 

184 target_provider, 

185 ) 

186 await self._unload_other_providers(target_provider) 

187 

188 # Add a small pause to ensure unloading takes effect 

189 await asyncio.sleep(0.5) 

190 

191 # If switching providers, update active provider 

192 if provider and provider != self.active_provider: 

193 self.active_provider = provider 

194 

195 logger.info("Switching to model %s in provider %s", model_name, target_provider) 

196 return await self.managers[target_provider].switch_model(model_name, **kwargs) 

197 

198 async def get_loaded_models(self, provider: str | None = None) -> list[str]: 

199 """Get currently loaded models.""" 

200 target_provider = provider or self.active_provider 

201 if not target_provider or target_provider not in self.managers: 

202 return [] 

203 

204 return await self.managers[target_provider].get_loaded_models() 

205 

206 def get_current_provider(self) -> str | None: 

207 """Get the currently active provider.""" 

208 return self.active_provider 

209 

210 async def close(self) -> None: 

211 """Close all managers.""" 

212 for manager in self.managers.values(): 

213 await manager.close() 

214 self.managers.clear() 

215 logger.info("Closed all model managers")