Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/pass_through_endpoints/passthrough_endpoint_router.py: 36%
134 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1import json
2from collections.abc import Callable
3from typing import TYPE_CHECKING, Final
5import litellm
6from litellm._logging import verbose_router_logger
7from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
8 LiteLLM_ManagedVectorStore,
9)
10from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
11from litellm.secret_managers.main import get_secret_str
12from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES
13from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
14from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict
16if TYPE_CHECKING: 16 ↛ 17line 16 didn't jump to line 17 because the condition on line 16 was never true
17 from litellm.router import Router
20def _get_proxy_llm_router() -> "Router | None":
21 from litellm.proxy.proxy_server import llm_router
23 return llm_router
26def _get_str_value(values: dict[str, object] | None, key: str) -> str | None:
27 value: Final = values.get(key) if values is not None else None
28 return value if isinstance(value, str) else None
31def _credential_identity(credentials: VERTEX_CREDENTIALS_TYPES | None) -> str | None:
32 """
33 A hashable stand-in for a credential, so two deployments can be compared for holding the same one
34 """
35 if isinstance(credentials, dict):
36 return json.dumps(credentials, sort_keys=True)
37 return credentials
40class PassthroughEndpointRouter:
41 """
42 Use this class to Get credentials for pass-through endpoints
43 """
45 def __init__(
46 self,
47 llm_router_getter: "Callable[[], Router | None]" = _get_proxy_llm_router,
48 ):
49 self.llm_router_getter: Final = llm_router_getter
50 self.deployment_key_to_vertex_credentials: dict[str, VertexPassThroughCredentials] = {}
51 self.default_vertex_config: VertexPassThroughCredentials | None = None
53 def get_credentials(
54 self,
55 custom_llm_provider: str,
56 region_name: str | None,
57 ) -> str | None:
58 deployment_api_key: Final = self._get_deployment_api_key(
59 custom_llm_provider=custom_llm_provider,
60 region_name=region_name,
61 )
62 if deployment_api_key is not None: 62 ↛ 63line 62 didn't jump to line 63 because the condition on line 62 was never true
63 return deployment_api_key
64 verbose_router_logger.debug(
65 "No pass-through deployment credentials found for %s, looking for env variable", custom_llm_provider
66 )
67 _env_variable_name: Final = self._get_default_env_variable_name_passthrough_endpoint(
68 custom_llm_provider=custom_llm_provider,
69 )
70 return get_secret_str(_env_variable_name)
72 def _get_deployment_api_key(
73 self,
74 custom_llm_provider: str,
75 region_name: str | None,
76 ) -> str | None:
77 llm_router: Final = self.llm_router_getter()
78 if llm_router is None: 78 ↛ 79line 78 didn't jump to line 79 because the condition on line 78 was never true
79 return None
80 deployments: Final = llm_router.get_model_list() or ()
81 return next(
82 (
83 api_key
84 for deployment in deployments
85 if (
86 api_key := self._resolve_matching_deployment_api_key(
87 litellm_params=deployment["litellm_params"],
88 custom_llm_provider=custom_llm_provider,
89 region_name=region_name,
90 )
91 )
92 is not None
93 ),
94 None,
95 )
97 def _resolve_matching_deployment_api_key(
98 self,
99 litellm_params: LiteLLMParamsTypedDict,
100 custom_llm_provider: str,
101 region_name: str | None,
102 ) -> str | None:
103 if litellm_params.get("use_in_pass_through") is not True: 103 ↛ 105line 103 didn't jump to line 105 because the condition on line 103 was always true
104 return None
105 if self._get_deployment_provider(litellm_params) != custom_llm_provider:
106 return None
107 credential_name: Final = litellm_params.get("litellm_credential_name")
108 credential_values: Final = (
109 CredentialAccessor.get_credential_values(credential_name) if credential_name is not None else None
110 )
111 api_base: Final = _get_str_value(credential_values, "api_base") or litellm_params.get("api_base")
112 deployment_region: Final = self._get_region_name_from_api_base(
113 custom_llm_provider=custom_llm_provider,
114 api_base=api_base,
115 )
116 if deployment_region != region_name:
117 return None
118 return _get_str_value(credential_values, "api_key") or litellm_params.get("api_key")
120 def _get_deployment_provider(self, litellm_params: LiteLLMParamsTypedDict) -> str | None:
121 model: Final = litellm_params.get("model")
122 if model is None:
123 return None
124 try:
125 _, provider, _, _ = litellm.get_llm_provider(
126 model=model,
127 custom_llm_provider=litellm_params.get("custom_llm_provider"),
128 )
129 except litellm.exceptions.BadRequestError:
130 return None
131 return provider
133 def get_vertex_credentials_from_router_deployments(self, model: str | None) -> VertexPassThroughCredentials | None:
134 """
135 Resolve vertex pass-through credentials from the live router deployments flagged ``use_in_pass_through``.
137 ``deployment_key_to_vertex_credentials`` is only reachable when the caller names a project and location,
138 which WebSocket clients never do, so DB-stored deployments need this lookup to be usable at all.
140 With no model to go on, only deployments that agree on a project, a location, and a credential answer:
141 guessing between two Vertex projects would mint a token for one and later send the other one's model name
142 """
143 llm_router: Final = self.llm_router_getter()
144 if llm_router is None:
145 return None
146 resolved: Final = tuple(
147 (deployment, credentials)
148 for deployment in (llm_router.get_model_list() or ())
149 if (credentials := self._resolve_vertex_deployment_credentials(deployment["litellm_params"])) is not None
150 )
151 matched: Final = next(
152 (
153 credentials
154 for deployment, credentials in resolved
155 if model is not None and self._deployment_matches_model(deployment, model)
156 ),
157 None,
158 )
159 if matched is not None:
160 return matched
161 targets: Final = frozenset(
162 (
163 credentials.vertex_project,
164 credentials.vertex_location,
165 _credential_identity(credentials.vertex_credentials),
166 )
167 for _, credentials in resolved
168 )
169 if len(targets) != 1:
170 return None
171 return resolved[0][1]
173 def _resolve_vertex_deployment_credentials(
174 self, litellm_params: LiteLLMParamsTypedDict
175 ) -> VertexPassThroughCredentials | None:
176 if litellm_params.get("use_in_pass_through") is not True:
177 return None
178 if self._get_deployment_provider(litellm_params) != "vertex_ai":
179 return None
180 credential_name: Final = litellm_params.get("litellm_credential_name")
181 credential_values: Final = (
182 CredentialAccessor.get_credential_values(credential_name) if credential_name is not None else None
183 )
184 vertex_project: Final = _get_str_value(credential_values, "vertex_project") or litellm_params.get(
185 "vertex_project"
186 )
187 vertex_location: Final = _get_str_value(credential_values, "vertex_location") or litellm_params.get(
188 "vertex_location"
189 )
190 stored_credentials: Final = (
191 credential_values.get("vertex_credentials") if credential_values is not None else None
192 )
193 vertex_credentials: Final = (
194 stored_credentials if isinstance(stored_credentials, (str, dict)) else None
195 ) or litellm_params.get("vertex_credentials")
196 if vertex_project is None or vertex_location is None:
197 return None
198 return VertexPassThroughCredentials(
199 vertex_project=vertex_project,
200 vertex_location=vertex_location,
201 vertex_credentials=vertex_credentials,
202 )
204 @staticmethod
205 def _deployment_matches_model(deployment: DeploymentTypedDict, model: str) -> bool:
206 upstream_model: Final = deployment["litellm_params"].get("model")
207 return model in (
208 deployment.get("model_name"),
209 upstream_model,
210 upstream_model.split("/", 1)[-1] if upstream_model is not None else None,
211 )
213 def _get_vertex_env_vars(self) -> VertexPassThroughCredentials:
214 """
215 Helper to get vertex pass through config from environment variables
217 The following environment variables are used:
218 - DEFAULT_VERTEXAI_PROJECT (project id)
219 - DEFAULT_VERTEXAI_LOCATION (location)
220 - DEFAULT_GOOGLE_APPLICATION_CREDENTIALS (path to credentials file)
221 """
222 return VertexPassThroughCredentials(
223 vertex_project=get_secret_str("DEFAULT_VERTEXAI_PROJECT"),
224 vertex_location=get_secret_str("DEFAULT_VERTEXAI_LOCATION"),
225 vertex_credentials=get_secret_str("DEFAULT_GOOGLE_APPLICATION_CREDENTIALS"),
226 )
228 def set_default_vertex_config(self, config: dict | None = None):
229 """Sets vertex configuration from provided config and/or environment variables
231 Args:
232 config (Optional[dict]): Configuration dictionary
233 Example: {
234 "vertex_project": "my-project-123",
235 "vertex_location": "us-central1",
236 "vertex_credentials": "os.environ/GOOGLE_CREDS"
237 }
238 """
239 # Initialize config dictionary if None
240 if config is None: 240 ↛ 244line 240 didn't jump to line 244 because the condition on line 240 was always true
241 self.default_vertex_config = self._get_vertex_env_vars()
242 return
244 if isinstance(config, dict):
245 for key, value in config.items():
246 if isinstance(value, str) and value.startswith("os.environ/"):
247 config[key] = get_secret_str(value)
249 self.default_vertex_config = VertexPassThroughCredentials(**config)
251 def add_vertex_credentials(
252 self,
253 project_id: str,
254 location: str,
255 vertex_credentials: VERTEX_CREDENTIALS_TYPES | None,
256 ):
257 """
258 Add the vertex credentials for the given project-id, location
259 """
261 deployment_key: Final = self._get_deployment_key(
262 project_id=project_id,
263 location=location,
264 )
265 if deployment_key is None:
266 verbose_router_logger.debug("No deployment key found for project-id, location")
267 return
268 vertex_pass_through_credentials: Final = VertexPassThroughCredentials(
269 vertex_project=project_id,
270 vertex_location=location,
271 vertex_credentials=vertex_credentials,
272 )
273 self.deployment_key_to_vertex_credentials[deployment_key] = vertex_pass_through_credentials
275 def _get_deployment_key(self, project_id: str | None, location: str | None) -> str | None:
276 """
277 Get the deployment key for the given project-id, location
278 """
279 if project_id is None or location is None: 279 ↛ 281line 279 didn't jump to line 281 because the condition on line 279 was always true
280 return None
281 return f"{project_id}-{location}"
283 def get_vector_store_credentials(self, vector_store_id: str) -> LiteLLM_ManagedVectorStore | None:
284 """
285 Get the vector store credentials for the given vector store id
286 """
287 if litellm.vector_store_registry is None:
288 return None
289 vector_store_to_run: Final[LiteLLM_ManagedVectorStore | None] = (
290 litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
291 vector_store_id=vector_store_id
292 )
293 )
294 return vector_store_to_run
296 def get_vertex_credentials(
297 self, project_id: str | None, location: str | None
298 ) -> VertexPassThroughCredentials | None:
299 """
300 Get the vertex credentials for the given project-id, location
301 """
302 deployment_key: Final = self._get_deployment_key(
303 project_id=project_id,
304 location=location,
305 )
307 if deployment_key is None: 307 ↛ 309line 307 didn't jump to line 309 because the condition on line 307 was always true
308 return self.default_vertex_config
309 if deployment_key in self.deployment_key_to_vertex_credentials:
310 return self.deployment_key_to_vertex_credentials[deployment_key]
311 else:
312 return self.default_vertex_config
314 def _get_region_name_from_api_base(
315 self,
316 custom_llm_provider: str,
317 api_base: str | None,
318 ) -> str | None:
319 """
320 Get the region name from the API base.
322 Each provider might have a different way of specifying the region in the API base - this is where you can use conditional logic to handle that.
323 """
324 if custom_llm_provider == "assemblyai":
325 if api_base and "eu" in api_base:
326 return "eu"
327 return None
329 @staticmethod
330 def _get_default_env_variable_name_passthrough_endpoint(
331 custom_llm_provider: str,
332 ) -> str:
333 return f"{custom_llm_provider.upper()}_API_KEY"