main
py 293 lines 10 KB
Raw
1 from __future__ import annotations
2
3 from typing import Any
4
5 import httpx
6 from helpers.api import ApiHandler, Request, Response
7 from helpers.providers import get_provider_config
8 import models
9
10 # Model name substrings to exclude from chat dropdowns and LiteLLM fallback results.
11 _NON_CHAT_EXCLUDE = frozenset({
12 "dall-e",
13 "gpt-image",
14 "image",
15 "tts",
16 "text-to-speech",
17 "whisper",
18 "audio",
19 "transcribe",
20 "transcription",
21 "speech",
22 "realtime",
23 "embedding",
24 "embed",
25 "moderation",
26 "omni-moderation",
27 "vision-preview",
28 })
29 _LOCAL_PLACEHOLDER_KEYS = {
30 "lm_studio": {"lm-studio"},
31 "llama_cpp": {"llama-cpp"},
32 "omlx": {"omlx"},
33 "vllm": {"vllm"},
34 }
35
36
37 class ModelSearch(ApiHandler):
38 async def process(self, input: dict, request: Request) -> dict | Response:
39 provider = str(input.get("provider", "") or "").strip().lower()
40 model_type = str(input.get("model_type", "chat") or "chat").strip().lower()
41 query = str(input.get("query", "") or "").strip().lower()
42 user_api_base = str(input.get("api_base", "") or "").strip()
43
44 if not provider:
45 return {"models": [], "provider": "", "source": "none", "error": ""}
46
47 cfg = self._get_provider_cfg(model_type, provider)
48 ml = self._get_models_list(cfg)
49
50 models_list, source, error = await self._fetch_models(provider, cfg, ml, user_api_base)
51
52 if not models_list:
53 fallback = self._litellm_fallback(provider, cfg)
54 if fallback:
55 models_list = fallback
56 source = "litellm_registry"
57 elif not source:
58 source = "none"
59
60 models_list = self._filter_models(models_list, model_type)
61 if query:
62 models_list = [name for name in models_list if query in name.lower()]
63
64 return {
65 "models": sorted(set(models_list), key=str.lower),
66 "provider": provider,
67 "source": source,
68 "error": error,
69 }
70
71 @staticmethod
72 def _get_provider_cfg(model_type: str, provider: str) -> dict:
73 """Get provider config, falling back to chat config for models_list."""
74 cfg = get_provider_config(model_type, provider) or {}
75 if model_type != "chat" and not cfg.get("models_list"):
76 chat_cfg = get_provider_config("chat", provider) or {}
77 if chat_cfg.get("models_list"):
78 merged = dict(cfg)
79 merged["models_list"] = chat_cfg["models_list"]
80 return merged
81 return cfg
82
83 @staticmethod
84 def _get_models_list(cfg: dict) -> dict:
85 """Extract models_list sub-config."""
86 return cfg.get("models_list") or {}
87
88 async def _fetch_models(
89 self,
90 provider: str,
91 cfg: dict,
92 ml: dict,
93 user_api_base: str = "",
94 ) -> tuple[list[str], str, str]:
95 api_key = models.get_api_key(provider)
96 kwargs = (cfg or {}).get("kwargs", {}) or {}
97 api_base = user_api_base or kwargs.get("api_base", "") or ml.get("default_base", "")
98 effective_ml = dict(ml or {})
99
100 # Ollama's native endpoint is /api/tags, but user-supplied /v1 bases usually
101 # mean the OpenAI-compatible /v1/models endpoint.
102 if provider == "ollama" and user_api_base.rstrip("/").endswith("/v1"):
103 effective_ml["endpoint_url"] = "/models"
104 effective_ml["format"] = "openai"
105
106 url, fmt = self._resolve_url(effective_ml, api_base)
107 if not url:
108 return [], "none", ""
109
110 headers = self._build_headers(provider, api_key, cfg)
111 params = dict(effective_ml.get("params", {}) or {})
112
113 # Google uses query-param auth for the public models list endpoint.
114 if provider == "google" and api_key and api_key != "None":
115 params.setdefault("key", api_key)
116
117 urls: list[tuple[str, str]] = [(url, fmt)]
118 if provider == "ollama" and fmt == "ollama":
119 ps_url = self._ollama_ps_url(url)
120 if ps_url and ps_url != url:
121 urls.append((ps_url, "ollama"))
122
123 combined: list[str] = []
124 errors: list[str] = []
125
126 try:
127 async with httpx.AsyncClient(timeout=10.0) as client:
128 for candidate_url, candidate_fmt in urls:
129 resp = await client.get(candidate_url, headers=headers, params=params)
130 if resp.status_code == 200:
131 combined.extend(self._parse(resp.json(), candidate_fmt))
132 else:
133 errors.append(f"{candidate_url}: HTTP {resp.status_code}")
134 except Exception as exc:
135 errors.append(str(exc))
136
137 if combined:
138 return combined, "provider_endpoint", ""
139 return [], "provider_endpoint", "; ".join(errors)
140
141 @staticmethod
142 def _resolve_url(ml: dict, api_base: str) -> tuple[str | None, str]:
143 fmt = ml.get("format", "openai")
144 endpoint = str(ml.get("endpoint_url", "") or "")
145 default_base = str(ml.get("default_base", "") or "")
146
147 if endpoint.startswith("http://") or endpoint.startswith("https://"):
148 return endpoint, fmt
149
150 base = str(api_base or default_base or "").strip()
151 if not base:
152 return None, fmt
153
154 endpoint = endpoint or "/models"
155 base = base.rstrip("/")
156
157 if not endpoint.startswith("/"):
158 endpoint = "/" + endpoint
159
160 # Avoid doubled /v1/v1 when users enter a base ending in /v1 and metadata
161 # also contains a versioned endpoint.
162 if base.endswith("/v1") and endpoint.startswith("/v1/"):
163 endpoint = endpoint[3:]
164
165 return base + endpoint, fmt
166
167 @staticmethod
168 def _ollama_ps_url(resolved_url: str) -> str:
169 """Return the Ollama running-model endpoint for a resolved native URL."""
170 marker = "/api/"
171 if marker not in resolved_url:
172 return ""
173 return resolved_url.split(marker, 1)[0].rstrip("/") + "/api/ps"
174
175 def _build_headers(self, provider: str, api_key: str, cfg: dict | None) -> dict[str, str]:
176 headers: dict[str, str] = {}
177 has_key = bool(api_key and api_key.strip() and api_key != "None")
178
179 if provider == "anthropic":
180 if has_key:
181 headers["x-api-key"] = api_key
182 headers["anthropic-version"] = "2023-06-01"
183 elif provider == "google":
184 pass
185 elif provider == "azure":
186 if has_key:
187 headers["api-key"] = api_key
188 elif provider != "ollama":
189 if has_key and api_key not in _LOCAL_PLACEHOLDER_KEYS.get(provider, set()):
190 headers["Authorization"] = f"Bearer {api_key}"
191
192 extra = (cfg or {}).get("kwargs", {}).get("extra_headers", {})
193 if isinstance(extra, dict):
194 for key, value in extra.items():
195 if isinstance(value, str):
196 headers[key] = value
197
198 return headers
199
200 def _litellm_fallback(self, provider: str, cfg: dict | None) -> list[str]:
201 try:
202 import litellm
203
204 registry = getattr(litellm, "models_by_provider", None)
205 if not registry:
206 return []
207
208 litellm_provider = (cfg or {}).get("litellm_provider", provider)
209 raw_models = registry.get(litellm_provider, set()) or set()
210 if not raw_models:
211 return []
212
213 prefix = litellm_provider + "/"
214 result: list[str] = []
215 for name in raw_models:
216 clean = str(name or "")
217 clean = clean[len(prefix):] if clean.startswith(prefix) else clean
218 if clean and not self._is_non_chat_model(clean):
219 result.append(clean)
220 return result
221 except Exception:
222 return []
223
224 def _parse(self, data: dict | list, fmt: str) -> list[str]:
225 if isinstance(data, list):
226 return self._parse_list(data)
227
228 if not isinstance(data, dict):
229 return []
230
231 if fmt == "ollama":
232 return self._parse_models_array(data.get("models", []), "name")
233
234 if fmt == "google":
235 result = []
236 for item in data.get("models", []) or []:
237 if not isinstance(item, dict):
238 continue
239 name = str(item.get("name", "") or "")
240 if name.startswith("models/"):
241 name = name[7:]
242 if name:
243 result.append(name)
244 return result
245
246 if "data" in data:
247 return self._parse_models_array(data.get("data", []), "id")
248
249 if "models" in data:
250 return self._parse_models_array(data.get("models", []), "id")
251
252 return []
253
254 @staticmethod
255 def _parse_models_array(items: Any, primary_key: str) -> list[str]:
256 if not isinstance(items, list):
257 return []
258 result = []
259 for item in items:
260 if isinstance(item, str):
261 result.append(item)
262 elif isinstance(item, dict):
263 value = item.get(primary_key) or item.get("id") or item.get("name")
264 if value:
265 result.append(str(value))
266 return result
267
268 def _parse_list(self, data: list) -> list[str]:
269 result = []
270 for item in data:
271 if isinstance(item, str):
272 result.append(item)
273 elif isinstance(item, dict):
274 value = item.get("id") or item.get("name")
275 if value:
276 result.append(str(value))
277 return result
278
279 def _filter_models(self, model_names: list[str], model_type: str) -> list[str]:
280 cleaned = []
281 for name in model_names or []:
282 value = str(name or "").strip()
283 if not value:
284 continue
285 if model_type == "chat" and self._is_non_chat_model(value):
286 continue
287 cleaned.append(value)
288 return cleaned
289
290 @staticmethod
291 def _is_non_chat_model(name: str) -> bool:
292 low = name.lower()
293 return any(token in low for token in _NON_CHAT_EXCLUDE)