mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 5744a5fcc2 into 2e1a98f521
This commit is contained in:
commit
00348f128b
2 changed files with 195 additions and 19 deletions
|
|
@ -2,13 +2,14 @@
|
|||
Translate from OpenAI's `/v1/chat/completions` to SAP Generative AI Hub's Orchestration Service`v2/completion`
|
||||
"""
|
||||
|
||||
import os
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from functools import cached_property
|
||||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
|
@ -39,6 +40,10 @@ from .models import (
|
|||
SAPUserMessage,
|
||||
)
|
||||
|
||||
# Env var that an operator can set to pin a specific orchestration deployment URL.
|
||||
# An empty string is treated as unset and falls through to auto-discovery.
|
||||
_AICORE_ORCHESTRATION_DEPLOYMENT_URL_ENV_VAR = "AICORE_ORCHESTRATION_DEPLOYMENT_URL"
|
||||
|
||||
# Keys routed outside SAP orchestration `model.params` (prompt, stream, fallbacks, etc.)
|
||||
_SAP_MODEL_PARAMS_EXCLUDED_KEYS: Final[frozenset[str]] = frozenset(
|
||||
{
|
||||
|
|
@ -135,6 +140,7 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
|
|||
self.token_creator = None
|
||||
self._base_url = None
|
||||
self._resource_group = None
|
||||
self._cached_deployment_url: str | None = None # None = not yet resolved
|
||||
|
||||
def run_env_setup(self, service_key: str | None = None) -> None:
|
||||
try:
|
||||
|
|
@ -166,23 +172,63 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
|
|||
self.run_env_setup()
|
||||
return self._resource_group
|
||||
|
||||
@cached_property
|
||||
@property
|
||||
def deployment_url(self) -> str:
|
||||
# Keep a short, tight client lifecycle here to avoid fd leaks
|
||||
client: Final = litellm.module_level_client
|
||||
# with httpx.Client(timeout=30) as client:
|
||||
deployments: Final = client.get(f"{self.base_url}/lm/deployments", headers=self.headers).json()
|
||||
valid: Final[list[tuple[str, str]]] = []
|
||||
for dep in deployments.get("resources", []):
|
||||
if dep.get("scenarioId") == "orchestration":
|
||||
cfg = client.get(
|
||||
f"{self.base_url}/lm/configurations/{dep['configurationId']}",
|
||||
headers=self.headers,
|
||||
).json()
|
||||
if cfg.get("executableId") == "orchestration":
|
||||
valid.append((dep["deploymentUrl"], dep["createdAt"]))
|
||||
# newest first
|
||||
return sorted(valid, key=lambda x: x[1], reverse=True)[0][0]
|
||||
"""Resolve the orchestration deployment URL, with caching.
|
||||
|
||||
Resolution order:
|
||||
1. ``_AICORE_ORCHESTRATION_DEPLOYMENT_URL_ENV_VAR`` env var (operator-level pin).
|
||||
An empty string is treated as unset and falls through to discovery.
|
||||
2. Auto-discovery via ``/lm/deployments`` (one network call, cached).
|
||||
|
||||
A per-request override via ``optional_params["deployment_url"]`` is
|
||||
handled one level up in ``get_complete_url`` before this property
|
||||
is ever called.
|
||||
"""
|
||||
cached = self._cached_deployment_url
|
||||
if cached is None:
|
||||
env = os.environ.get(_AICORE_ORCHESTRATION_DEPLOYMENT_URL_ENV_VAR)
|
||||
cached = env or self._resolve_deployment_url()
|
||||
self._cached_deployment_url = cached
|
||||
return cached
|
||||
|
||||
def _resolve_deployment_url(self) -> str:
|
||||
resources: Final = (
|
||||
litellm.module_level_client.get(
|
||||
f"{self.base_url}/lm/deployments",
|
||||
headers=self.headers,
|
||||
params={ # mutable-ok: one-shot query-params dict passed directly to httpx, not stored
|
||||
"scenarioId": "orchestration",
|
||||
"executableIds": ["orchestration"],
|
||||
"status": "RUNNING",
|
||||
},
|
||||
)
|
||||
.json()
|
||||
.get("resources", [])
|
||||
)
|
||||
candidates: Final[list[tuple[str, str, str]]] = sorted(
|
||||
(
|
||||
(dep["deploymentUrl"], dep["createdAt"], dep.get("configurationName", ""))
|
||||
for dep in resources
|
||||
if dep.get("deploymentUrl")
|
||||
),
|
||||
key=lambda c: c[1],
|
||||
reverse=True,
|
||||
)
|
||||
if not candidates:
|
||||
raise GenAIHubOrchestrationError(
|
||||
status_code=404,
|
||||
message="No orchestration deployment found in SAP AI Core.",
|
||||
)
|
||||
if len(candidates) > 1:
|
||||
verbose_logger.warning(
|
||||
"SAP: %d orchestration deployments found; using newest (name=%r, url=%r). Others ignored: %s.",
|
||||
len(candidates),
|
||||
candidates[0][2],
|
||||
candidates[0][0],
|
||||
tuple(v[2] or v[0] for v in candidates[1:]),
|
||||
)
|
||||
return candidates[0][0]
|
||||
|
||||
@classmethod
|
||||
def get_config(cls):
|
||||
|
|
@ -249,8 +295,9 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
|
|||
litellm_params: dict,
|
||||
stream: bool | None = None,
|
||||
):
|
||||
api_base_: Final = f"{self.deployment_url}/v2/completion"
|
||||
return api_base_
|
||||
# Per-request override wins; deployment_url handles env var + discovery.
|
||||
base = optional_params.get("deployment_url") or self.deployment_url
|
||||
return f"{base}/v2/completion"
|
||||
|
||||
def _build_prompt_module(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import warnings
|
||||
import pytest
|
||||
from unittest.mock import MagicMock, patch
|
||||
from pydantic import ValidationError
|
||||
|
||||
|
||||
|
|
@ -639,3 +640,131 @@ class TestSAPTransformationIntegration:
|
|||
config["config"]["modules"][1]["translation"]["input"]["type"]
|
||||
== "sap_document_translation"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers shared by deployment-resolution tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_config():
|
||||
from litellm.llms.sap.chat.transformation import GenAIHubOrchestrationConfig
|
||||
|
||||
cfg = GenAIHubOrchestrationConfig()
|
||||
cfg.token_creator = lambda: "Bearer TEST"
|
||||
cfg._base_url = "https://api.test-sap.com"
|
||||
cfg._resource_group = "test-group"
|
||||
return cfg
|
||||
|
||||
|
||||
def _mock_client(*names: str):
|
||||
"""Fake httpx client returning a /lm/deployments response for the given deployment names.
|
||||
|
||||
The backend filters by scenarioId/executableIds/status, so there is no
|
||||
second call to /lm/configurations; the mock only needs one response shape.
|
||||
"""
|
||||
resources = [
|
||||
{
|
||||
"scenarioId": "orchestration",
|
||||
"configurationId": f"cfg-{n}",
|
||||
"deploymentUrl": f"https://deploy-{n}.sap.com",
|
||||
"createdAt": f"2024-01-{i + 1:02d}T00:00:00Z",
|
||||
"configurationName": n,
|
||||
}
|
||||
for i, n in enumerate(names)
|
||||
]
|
||||
resp = MagicMock()
|
||||
resp.json.return_value = {"resources": resources}
|
||||
client = MagicMock()
|
||||
client.get.return_value = resp
|
||||
return client
|
||||
|
||||
|
||||
class TestDeploymentResolution:
|
||||
def test_single_deployment_returns_url(self):
|
||||
cfg = _make_config()
|
||||
with patch("litellm.module_level_client", _mock_client("orch-a")): # test-quality-ok: patching the HTTP client at the litellm transport boundary
|
||||
url = cfg._resolve_deployment_url()
|
||||
assert url == "https://deploy-orch-a.sap.com"
|
||||
|
||||
def test_multiple_deployments_picks_newest(self):
|
||||
# "older" has createdAt 2024-01-01, "newer" has 2024-01-02
|
||||
cfg = _make_config()
|
||||
with patch("litellm.module_level_client", _mock_client("older", "newer")): # test-quality-ok: patching the HTTP client at the litellm transport boundary
|
||||
url = cfg._resolve_deployment_url()
|
||||
assert url == "https://deploy-newer.sap.com"
|
||||
|
||||
def test_multiple_deployments_emits_warning(self):
|
||||
cfg = _make_config()
|
||||
with patch("litellm.module_level_client", _mock_client("older", "newer")): # test-quality-ok: patching the HTTP client at the litellm transport boundary
|
||||
with patch("litellm.llms.sap.chat.transformation.verbose_logger") as mock_log: # test-quality-ok: patching module-level logger to intercept warning calls
|
||||
url = cfg._resolve_deployment_url()
|
||||
mock_log.warning.assert_called_once()
|
||||
assert 2 == mock_log.warning.call_args[0][1]
|
||||
assert url == "https://deploy-newer.sap.com"
|
||||
|
||||
def test_no_deployments_raises(self):
|
||||
from litellm.llms.sap.chat.handler import GenAIHubOrchestrationError
|
||||
|
||||
cfg = _make_config()
|
||||
with patch("litellm.module_level_client", _mock_client()): # test-quality-ok: patching the HTTP client at the litellm transport boundary
|
||||
with pytest.raises(GenAIHubOrchestrationError) as exc:
|
||||
cfg._resolve_deployment_url()
|
||||
assert "No orchestration deployment found" in str(exc.value)
|
||||
|
||||
def test_query_params_sent_to_backend(self):
|
||||
"""The whole point of this refactor: filtering must happen server-side."""
|
||||
cfg = _make_config()
|
||||
mock = _mock_client("orch-a")
|
||||
with patch("litellm.module_level_client", mock): # test-quality-ok: patching the HTTP client at the litellm transport boundary
|
||||
cfg._resolve_deployment_url()
|
||||
_, kwargs = mock.get.call_args
|
||||
assert kwargs["params"] == {
|
||||
"scenarioId": "orchestration",
|
||||
"executableIds": ["orchestration"],
|
||||
"status": "RUNNING",
|
||||
}
|
||||
|
||||
|
||||
class TestGetCompleteUrl:
|
||||
def test_optional_params_used_first(self):
|
||||
"""Step 1: deployment_url in optional_params skips discovery entirely."""
|
||||
cfg = _make_config()
|
||||
explicit = "https://custom.sap.com/deployments/abc"
|
||||
mock = MagicMock()
|
||||
with patch("litellm.module_level_client", mock): # test-quality-ok: patching the HTTP client at the litellm transport boundary
|
||||
url = cfg.get_complete_url(None, None, "gpt-4o", {"deployment_url": explicit}, {})
|
||||
assert url == f"{explicit}/v2/completion"
|
||||
mock.get.assert_not_called()
|
||||
|
||||
def test_env_var_used_when_no_optional_param(self):
|
||||
"""Step 2: AICORE_ORCHESTRATION_DEPLOYMENT_URL skips discovery."""
|
||||
cfg = _make_config()
|
||||
env_url = "https://env.sap.com/deployments/env"
|
||||
mock = MagicMock()
|
||||
with patch("litellm.module_level_client", mock): # test-quality-ok: patching the HTTP client at the litellm transport boundary
|
||||
with patch.dict("os.environ", {"AICORE_ORCHESTRATION_DEPLOYMENT_URL": env_url}):
|
||||
url = cfg.get_complete_url(None, None, "gpt-4o", {}, {})
|
||||
assert url == f"{env_url}/v2/completion"
|
||||
mock.get.assert_not_called()
|
||||
|
||||
def test_optional_params_beats_env_var(self):
|
||||
"""Step 1 takes precedence over step 2."""
|
||||
cfg = _make_config()
|
||||
opt_url = "https://opt.sap.com/deployments/opt"
|
||||
env_url = "https://env.sap.com/deployments/env"
|
||||
mock = MagicMock()
|
||||
with patch("litellm.module_level_client", mock): # test-quality-ok: patching the HTTP client at the litellm transport boundary
|
||||
with patch.dict("os.environ", {"AICORE_ORCHESTRATION_DEPLOYMENT_URL": env_url}):
|
||||
url = cfg.get_complete_url(None, None, "gpt-4o", {"deployment_url": opt_url}, {})
|
||||
assert url == f"{opt_url}/v2/completion"
|
||||
mock.get.assert_not_called()
|
||||
|
||||
def test_discovery_used_when_no_override(self):
|
||||
"""Step 3: no optional_param, no env var — discovery runs."""
|
||||
cfg = _make_config()
|
||||
env = {"AICORE_ORCHESTRATION_DEPLOYMENT_URL": ""} # empty = falsy
|
||||
with patch("litellm.module_level_client", _mock_client("orch-a")): # test-quality-ok: patching the HTTP client at the litellm transport boundary
|
||||
with patch.dict("os.environ", env):
|
||||
url = cfg.get_complete_url(None, None, "gpt-4o", {}, {})
|
||||
assert url == "https://deploy-orch-a.sap.com/v2/completion"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue