From 39846e491e40a9e15acaeaad24eb2ac439bc8efd Mon Sep 17 00:00:00 2001 From: Vasilisa Parshikova Date: Fri, 13 Mar 2026 15:18:50 +0400 Subject: [PATCH] (sap) run black formater --- litellm/llms/sap/chat/handler.py | 7 +- litellm/llms/sap/chat/models.py | 22 ++-- litellm/llms/sap/chat/transformation.py | 92 ++++++++------- litellm/llms/sap/credentials.py | 142 +++++++++++++++++------- 4 files changed, 169 insertions(+), 94 deletions(-) diff --git a/litellm/llms/sap/chat/handler.py b/litellm/llms/sap/chat/handler.py index 1390b2a4785..713143d895f 100755 --- a/litellm/llms/sap/chat/handler.py +++ b/litellm/llms/sap/chat/handler.py @@ -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: diff --git a/litellm/llms/sap/chat/models.py b/litellm/llms/sap/chat/models.py index 1b09ce9a756..8ca2aa7a690 100644 --- a/litellm/llms/sap/chat/models.py +++ b/litellm/llms/sap/chat/models.py @@ -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): diff --git a/litellm/llms/sap/chat/transformation.py b/litellm/llms/sap/chat/transformation.py index a019ba1767a..7f6bab4a1d5 100755 --- a/litellm/llms/sap/chat/transformation.py +++ b/litellm/llms/sap/chat/transformation.py @@ -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 diff --git a/litellm/llms/sap/credentials.py b/litellm/llms/sap/credentials.py index 5b6d4d40fd9..b6d8be18f5a 100644 --- a/litellm/llms/sap/credentials.py +++ b/litellm/llms/sap/credentials.py @@ -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