(sap) run black formater

This commit is contained in:
Vasilisa Parshikova 2026-03-13 15:18:50 +04:00
parent 1748ef82f2
commit 39846e491e
4 changed files with 169 additions and 94 deletions

View file

@ -251,11 +251,8 @@ class AsyncSAPStreamIterator:
# -------------------------------
class GenAIHubOrchestration(BaseLLMHTTPHandler):
def _add_stream_param_to_request_body(
self,
data: dict,
provider_config: BaseConfig,
fake_stream: bool
):
self, data: dict, provider_config: BaseConfig, fake_stream: bool
):
if data.get("config", {}).get("stream", None) is not None:
data["config"]["stream"]["enabled"] = True
else:

View file

@ -7,21 +7,22 @@ def validate_different_content(v: Union[str, dict, list]) -> str:
if v in ((), {}, []):
return ""
elif isinstance(v, dict) and "text" in v:
return v['text']
return v["text"]
elif isinstance(v, list):
new_v = []
for item in v:
if isinstance(item, dict) and "text" in item:
if item['text']:
new_v.append(item['text'])
if item["text"]:
new_v.append(item["text"])
elif isinstance(item, str):
new_v.append(item)
return '\n'.join(new_v)
return "\n".join(new_v)
elif isinstance(v, str):
return v
raise ValueError("Content must be a string")
return v
class TextContent(BaseModel):
type_: Literal["text"] = Field(default="text", alias="type")
text: str
@ -80,7 +81,9 @@ class SAPMessage(BaseModel):
role: Literal["system", "developer"] = "system"
content: str
_content_validator = field_validator("content", mode="before")(validate_different_content)
_content_validator = field_validator("content", mode="before")(
validate_different_content
)
class SAPUserMessage(BaseModel):
@ -96,8 +99,9 @@ class SAPAssistantMessage(BaseModel):
refusal: str = ""
tool_calls: list[MessageToolCall] = []
_content_validator = field_validator("content", mode="before")(validate_different_content)
_content_validator = field_validator("content", mode="before")(
validate_different_content
)
class SAPToolChatMessage(BaseModel):
@ -105,7 +109,9 @@ class SAPToolChatMessage(BaseModel):
tool_call_id: str
content: str
_content_validator = field_validator("content", mode="before")(validate_different_content)
_content_validator = field_validator("content", mode="before")(
validate_different_content
)
class ResponseFormat(BaseModel):

View file

@ -1,7 +1,17 @@
"""
Translate from OpenAI's `/v1/chat/completions` to SAP Generative AI Hub's Orchestration Service`v2/completion`
"""
from typing import List, Optional, Union, Dict, Tuple, Any, TYPE_CHECKING, Iterator, AsyncIterator
from typing import (
List,
Optional,
Union,
Dict,
Tuple,
Any,
TYPE_CHECKING,
Iterator,
AsyncIterator,
)
from functools import cached_property
import litellm
import httpx
@ -29,7 +39,12 @@ from .models import (
ResponseFormat,
SAPUserMessage,
)
from .handler import GenAIHubOrchestrationError, AsyncSAPStreamIterator, SAPStreamIterator
from .handler import (
GenAIHubOrchestrationError,
AsyncSAPStreamIterator,
SAPStreamIterator,
)
def validate_dict(data: dict, model) -> dict:
return model(**data).model_dump(by_alias=True)
@ -77,16 +92,15 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
def run_env_setup(self, service_key: Optional[str] = None) -> None:
try:
self.token_creator, self._base_url, self._resource_group = get_token_creator(service_key) # type: ignore
self.token_creator, self._base_url, self._resource_group = get_token_creator(service_key) # type: ignore
except ValueError as err:
raise GenAIHubOrchestrationError(status_code=400, message=err.args[0])
@property
def headers(self) -> Dict[str, str]:
if self.token_creator is None:
self.run_env_setup()
access_token = self.token_creator() # type: ignore
access_token = self.token_creator() # type: ignore
return {
"Authorization": access_token,
"AI-Resource-Group": self.resource_group,
@ -98,14 +112,13 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
def base_url(self) -> str:
if self._base_url is None:
self.run_env_setup()
return self._base_url # type: ignore
return self._base_url # type: ignore
@property
def resource_group(self) -> str:
if self._resource_group is None:
self.run_env_setup()
return self._resource_group # type: ignore
return self._resource_group # type: ignore
@cached_property
def deployment_url(self) -> str:
@ -169,7 +182,6 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
params.remove("tool_choice")
return params
def validate_environment(
self,
headers: dict,
@ -185,13 +197,13 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
return self.headers
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
):
api_base_ = f"{self.deployment_url}/v2/completion"
return api_base_
@ -199,7 +211,7 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
def transform_request(
self,
model: str,
messages: List[Dict[str, str]], # type: ignore
messages: List[Dict[str, str]], # type: ignore
optional_params: dict,
litellm_params: dict,
headers: dict,
@ -240,8 +252,10 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
response_format = model_params.pop("response_format", {})
resp_type = response_format.get("type", None)
if resp_type:
if resp_type== "json_schema":
response_format = validate_dict(response_format, ResponseFormatJSONSchema)
if resp_type == "json_schema":
response_format = validate_dict(
response_format, ResponseFormatJSONSchema
)
else:
response_format = validate_dict(response_format, ResponseFormat)
response_format = {"response_format": response_format}
@ -259,11 +273,7 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
"config": {
"modules": {
"prompt_templating": {
"prompt": {
"template": template,
**tools,
**response_format
},
"prompt": {"template": template, **tools, **response_format},
"model": {
"name": model,
"params": model_params,
@ -278,18 +288,18 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
return config
def transform_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ModelResponse:
logging_obj.post_call(
input=messages,
@ -323,17 +333,17 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig):
if choice.message and choice.message.content:
content = choice.message.content.strip()
# Match ```json ... ``` or ``` ... ```
match = re.match(r'^```(?:json)?\s*\n?(.*?)\n?```$', content, re.DOTALL)
match = re.match(r"^```(?:json)?\s*\n?(.*?)\n?```$", content, re.DOTALL)
if match:
choice.message.content = match.group(1).strip()
return response
def get_model_response_iterator(
self,
streaming_response: Union[Iterator[str], AsyncIterator[str], "ModelResponse"],
sync_stream: bool,
json_mode: Optional[bool] = False,
self,
streaming_response: Union[Iterator[str], AsyncIterator[str], "ModelResponse"],
sync_stream: bool,
json_mode: Optional[bool] = False,
):
if sync_stream:
return SAPStreamIterator(response=streaming_response) # type: ignore

View file

@ -61,6 +61,7 @@ def _load_json_env(var_name: str) -> Optional[Dict[str, Any]]:
except json.JSONDecodeError:
return None
def _str_or_none(value) -> Optional[str]:
try:
return str(value) if value is not None else None
@ -79,11 +80,13 @@ def _get_vcap_service(label: str) -> Optional[Dict[str, Any]]:
return svc
return None
@dataclass
class Source:
name: str
get: Callable[[CredentialsValue], Optional[str]]
@dataclass(frozen=True)
class CredentialsValue:
name: str
@ -125,6 +128,7 @@ CREDENTIAL_VALUES: Final[List[CredentialsValue]] = [
),
]
def init_conf(profile: Optional[str] = None) -> Dict[str, Any]:
"""
Loads config JSON from:
@ -167,6 +171,7 @@ def init_conf(profile: Optional[str] = None) -> Dict[str, Any]:
def _env_name(name: str) -> str:
return f"AICORE_{name.upper()}"
def extract_credentials(source: Source) -> Dict[str, str]:
"""Extract all credentials from a source."""
credentials = {}
@ -176,6 +181,7 @@ def extract_credentials(source: Source) -> Dict[str, str]:
credentials[cv.name] = cv.transform_fn(value) if cv.transform_fn else value
return credentials
def resolve_credentials(sources: List[Source]) -> Dict[str, str]:
"""Extract credentials from the first source that has any defined."""
for source in sources:
@ -185,23 +191,30 @@ def resolve_credentials(sources: List[Source]) -> Dict[str, str]:
return credentials
raise ValueError("No credentials found in any source")
def resolve_resource_group(sources: List[Source]) -> Optional[str]:
"""Find resource_group from the first source that defines it."""
rg_cred = CredentialsValue("resource_group", default="default")
for source in sources:
value = source.get(rg_cred)
if value is not None:
verbose_logger.debug(f"Resolved GEN AI Hub resource_group from source {source.name}")
verbose_logger.debug(
f"Resolved GEN AI Hub resource_group from source {source.name}"
)
return value
return rg_cred.default
def _function_to_resolve_cv_from_service_key(
service_key: Optional[Union[str, dict]], cv: CredentialsValue):
service_key: Optional[Union[str, dict]], cv: CredentialsValue
):
if service_key is None:
return None
val = _str_or_none(
_get_nested(service_key, (("credentials",) + cv.vcap_key) if cv.vcap_key else (cv.name,))
_get_nested(
service_key, (("credentials",) + cv.vcap_key) if cv.vcap_key else (cv.name,)
)
)
if val is None:
return _str_or_none(
_get_nested(service_key, cv.vcap_key if cv.vcap_key else (cv.name,))
@ -209,8 +222,11 @@ def _function_to_resolve_cv_from_service_key(
return val
def fetch_credentials(service_key: Optional[Union[str, dict]] = None, profile: Optional[str] = None, **kwargs) -> Dict[str, str]:
def fetch_credentials(
service_key: Optional[Union[str, dict]] = None,
profile: Optional[str] = None,
**kwargs,
) -> Dict[str, str]:
"""
Resolution order (first-source-wins):
@ -228,25 +244,42 @@ def fetch_credentials(service_key: Optional[Union[str, dict]] = None, profile: O
"""
config = init_conf(profile)
service_key = service_key or litellm.sap_service_key or _load_json_env(SERVICE_KEY_ENV_VAR)
service_key = (
service_key or litellm.sap_service_key or _load_json_env(SERVICE_KEY_ENV_VAR)
)
vcap_service = _get_vcap_service(VCAP_AICORE_SERVICE_NAME)
sources = [
Source("kwargs",
lambda cv: _str_or_none(kwargs.get(cv.name))),
Source("service key",
lambda cv: _function_to_resolve_cv_from_service_key(service_key, cv)), # type: ignore[arg-type]
Source("environment variables",
lambda cv: _str_or_none(os.environ.get(f'AICORE_{cv.name.upper()}'))),
Source("config file",
lambda cv: _str_or_none(config.get(f'AICORE_{cv.name.upper()}')
if config.get(f'AICORE_{cv.name.upper()}') is not None
else config.get(cv.name))),
Source("VCAP service",
lambda cv: (_str_or_none(
_get_nested(vcap_service, (("credentials",) + cv.vcap_key) if cv.vcap_key else (cv.name,))
)
if vcap_service else None)), # type: ignore[arg-type]
Source("kwargs", lambda cv: _str_or_none(kwargs.get(cv.name))),
Source(
"service key",
lambda cv: _function_to_resolve_cv_from_service_key(service_key, cv),
), # type: ignore[arg-type]
Source(
"environment variables",
lambda cv: _str_or_none(os.environ.get(f"AICORE_{cv.name.upper()}")),
),
Source(
"config file",
lambda cv: _str_or_none(
config.get(f"AICORE_{cv.name.upper()}")
if config.get(f"AICORE_{cv.name.upper()}") is not None
else config.get(cv.name)
),
),
Source(
"VCAP service",
lambda cv: (
_str_or_none(
_get_nested(
vcap_service,
(("credentials",) + cv.vcap_key) if cv.vcap_key else (cv.name,),
)
)
if vcap_service
else None
),
), # type: ignore[arg-type]
]
credentials = resolve_credentials(sources)
@ -254,21 +287,22 @@ def fetch_credentials(service_key: Optional[Union[str, dict]] = None, profile: O
resource_group = resolve_resource_group(sources)
if resource_group is not None:
credentials['resource_group'] = resource_group
credentials["resource_group"] = resource_group
if 'cert_url' in credentials:
credentials['auth_url'] = credentials.pop('cert_url')
if "cert_url" in credentials:
credentials["auth_url"] = credentials.pop("cert_url")
return credentials
def validate_credentials(
auth_url: Optional[str] = None,
base_url: Optional[str] = None,
client_id: Optional[str] = None,
client_secret: Optional[str] = None,
cert_str: Optional[str] = None,
key_str: Optional[str] = None,
cert_file_path: Optional[str] = None,
key_file_path: Optional[str] = None
auth_url: Optional[str] = None,
base_url: Optional[str] = None,
client_id: Optional[str] = None,
client_secret: Optional[str] = None,
cert_str: Optional[str] = None,
key_str: Optional[str] = None,
cert_file_path: Optional[str] = None,
key_file_path: Optional[str] = None,
):
if not auth_url or not client_id or not base_url:
raise ValueError(
@ -289,7 +323,10 @@ def validate_credentials(
"(cert_str & key_str), or (cert_file_path & key_file_path)."
)
def _request_token(client_id:str, auth_url: str, timeout:float, cert_pair=None, client_secret=None) -> tuple[str, datetime]:
def _request_token(
client_id: str, auth_url: str, timeout: float, cert_pair=None, client_secret=None
) -> tuple[str, datetime]:
data = {"grant_type": "client_credentials", "client_id": client_id}
if client_secret:
data["client_secret"] = client_secret
@ -313,6 +350,7 @@ def _request_token(client_id:str, auth_url: str, timeout:float, cert_pair=None,
msg = resp.text if resp is not None else getattr(e, "text", str(e))
raise RuntimeError(f"Token request failed: {msg}") from e
def get_token_creator(
service_key: Optional[Union[str, dict]] = None,
profile: Optional[str] = None,
@ -341,7 +379,9 @@ def get_token_creator(
"""
# Resolve credentials using your helper
credentials: Dict[str, str] = fetch_credentials(service_key=service_key, profile=profile, **overrides)
credentials: Dict[str, str] = fetch_credentials(
service_key=service_key, profile=profile, **overrides
)
auth_url = credentials.get("auth_url")
base_url = credentials.get("base_url")
@ -353,7 +393,16 @@ def get_token_creator(
key_file_path = credentials.get("key_file_path")
# Sanity check
validate_credentials(auth_url, base_url, client_id, client_secret, cert_str, key_str, cert_file_path, key_file_path)
validate_credentials(
auth_url,
base_url,
client_id,
client_secret,
cert_str,
key_str,
cert_file_path,
key_file_path,
)
lock = Lock()
token: Optional[str] = None
@ -362,7 +411,12 @@ def get_token_creator(
def _fetch_token() -> tuple[str, datetime]:
# Case 1: secret-based auth
if client_secret:
return _request_token(auth_url=auth_url, client_id=client_id, timeout=timeout, client_secret=client_secret)
return _request_token(
auth_url=auth_url,
client_id=client_id,
timeout=timeout,
client_secret=client_secret,
)
# Case 2: cert/key strings
if cert_str and key_str:
cert_str_fixed = cert_str.replace("\\n", "\n")
@ -374,11 +428,19 @@ def get_token_creator(
f.write(cert_str_fixed)
with open(key_path, "w") as f:
f.write(key_str_fixed)
return _request_token(auth_url=auth_url, client_id=client_id, timeout=timeout,
cert_pair=(cert_path, key_path))
return _request_token(
auth_url=auth_url,
client_id=client_id,
timeout=timeout,
cert_pair=(cert_path, key_path),
)
# Case 3: file-based cert/key
return _request_token(auth_url=auth_url, client_id=client_id, timeout=timeout,
cert_pair=(cert_file_path, key_file_path))
return _request_token(
auth_url=auth_url,
client_id=client_id,
timeout=timeout,
cert_pair=(cert_file_path, key_file_path),
)
def get_token() -> str:
nonlocal token, token_expiry