1"""OpenAPI-compatible model manager implementation.
2
3Generic provider for OpenAI-compatible APIs like LM Studio, VLLM, Ollama, etc.
4"""
5
6from __future__ import annotations
7
8import asyncio
9from typing import Any
10
11try:
12 import aiohttp
13except ImportError:
14 aiohttp = None # type: ignore[assignment]
15
16from lexigram.ai.llm.model_manager.base import AbstractModelManager
17from lexigram.ai.llm.model_manager.types import ModelLoadResult
18from lexigram.logging import (
19 get_logger,
20)
21
22logger = get_logger(__name__)
23
24
25class OpenAPICompatibleModelManager(AbstractModelManager):
26 """Generic model manager for OpenAI-compatible APIs.
27
28 Supports LM Studio, VLLM, Ollama, and other OpenAI-compatible endpoints.
29 """
30
31 def __init__(
32 self,
33 base_url: str = "http://localhost:8000",
34 provider_name: str = "openapi-compatible",
35 api_version: str = "v1",
36 supports_model_loading: bool = False,
37 supports_model_switching: bool = False,
38 ):
39 """Initialize OpenAPI-compatible model manager.
40
41 Args:
42 base_url: Base URL of the API server
43 provider_name: Name identifier for the provider
44 api_version: API version (v1, v1beta, etc.)
45 supports_model_loading: Whether the API supports explicit model loading
46 supports_model_switching: Whether the API supports model switching
47 """
48 super().__init__(base_url, f"{provider_name}-model-manager")
49 self.api_version = api_version
50 self.supports_model_loading = supports_model_loading
51 self.supports_model_switching = supports_model_switching
52 self._loaded_models: dict[str, Any] = {} # model_name -> handle/metadata
53
54 async def list_models(self) -> list[dict[str, Any]]:
55 """List available models via OpenAI-compatible API."""
56 client = await self._get_client()
57 resp = await client.get(f"/{self.api_version}/models")
58 resp.raise_for_status()
59 data = resp.json()
60 if asyncio.iscoroutine(data):
61 data = await data
62
63 models = data["data"]
64 return [
65 {
66 "name": model["id"],
67 "details": model,
68 }
69 for model in models
70 ]
71
72 async def load_model(self, model_name: str, **kwargs: Any) -> ModelLoadResult:
73 """Load a model via OpenAI-compatible API."""
74 if model_name in self._loaded_models:
75 logger.info("Model %s already loaded", model_name)
76 return ModelLoadResult(success=True, model_name=model_name)
77
78 if not self.supports_model_loading:
79 # For APIs that don't support explicit loading, just test connectivity
80 try:
81 await self._test_model_access(model_name)
82 self._loaded_models[model_name] = None
83 logger.info("Successfully prepared model %s", model_name)
84 return ModelLoadResult(success=True, model_name=model_name)
85 except (ConnectionError, TimeoutError, OSError, ValueError) as e:
86 return ModelLoadResult(
87 success=False,
88 model_name=model_name,
89 error=f"Failed to access model {model_name}: {e}",
90 retryable=True,
91 )
92
93 # For APIs that support explicit loading (future implementation)
94 # This would make specific API calls to load models
95 logger.warning("Explicit model loading not implemented for %s", self.name)
96 return ModelLoadResult(
97 success=False,
98 model_name=model_name,
99 error=f"Explicit model loading not supported by {self.name}",
100 )
101
102 async def unload_model(self, model_name: str) -> bool:
103 """Unload a model."""
104 if model_name not in self._loaded_models:
105 logger.info("Model %s not loaded", model_name)
106 return True
107
108 if not self.supports_model_loading:
109 # For APIs that don't support explicit unloading, just remove from tracking
110 del self._loaded_models[model_name]
111 logger.info("Model %s marked as unloaded", model_name)
112 return True
113
114 # For APIs that support explicit unloading (future implementation)
115 logger.warning("Explicit model unloading not implemented for %s", self.name)
116 return False
117
118 async def switch_model(self, model_name: str, **kwargs: Any) -> ModelLoadResult:
119 """Switch to a different model."""
120 if not self.supports_model_switching:
121 # For APIs that don't support switching, unload all and load new one
122 current_models = list(self._loaded_models.keys())
123 for model in current_models:
124 if model != model_name:
125 await self.unload_model(model)
126
127 return await self.load_model(model_name, **kwargs)
128
129 # For APIs that support explicit switching (future implementation)
130 logger.warning("Explicit model switching not implemented for %s", self.name)
131 return ModelLoadResult(
132 success=False,
133 model_name=model_name,
134 error=f"Model switching not supported by {self.name}",
135 )
136
137 async def get_loaded_models(self) -> list[str]:
138 """Get currently loaded models."""
139 return list(self._loaded_models.keys())
140
141 async def _test_model_access(self, model_name: str) -> None:
142 """Test if a model is accessible by making a minimal API call."""
143 client = await self._get_client()
144
145 # Make a minimal completion request to test model access
146 test_payload = {
147 "model": model_name,
148 "messages": [{"role": "user", "content": "test"}],
149 "max_tokens": 1,
150 "temperature": 0,
151 }
152
153 resp = await client.post(
154 f"/{self.api_version}/chat/completions",
155 json=test_payload,
156 )
157 resp.raise_for_status()
158
159
160class LMStudioModelManager(OpenAPICompatibleModelManager):
161 """Model manager for LM Studio using OpenAPI-compatible base."""
162
163 def __init__(self, base_url: str = "http://localhost:1234"):
164 super().__init__(
165 base_url=base_url,
166 provider_name="lmstudio",
167 api_version="v1",
168 supports_model_loading=False, # LM Studio loads on demand
169 supports_model_switching=False, # No explicit switching API
170 )
171
172 async def get_loaded_models(self) -> list[str]:
173 """Get currently loaded models in LM Studio."""
174 client = await self._get_client()
175 response = await client.list_loaded_models() # type: ignore[attr-defined]
176 return [m.identifier for m in response]
177
178
179class VLLMModelManager(OpenAPICompatibleModelManager):
180 """Model manager for VLLM using OpenAPI-compatible base."""
181
182 def __init__(self, base_url: str = "http://localhost:8000"):
183 super().__init__(
184 base_url=base_url,
185 provider_name="vllm",
186 api_version="v1",
187 supports_model_loading=False, # VLLM manages loading internally
188 supports_model_switching=False, # No explicit switching API
189 )