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 3436e13736a..ad6415e80f3 100644 --- a/litellm/llms/sap/chat/models.py +++ b/litellm/llms/sap/chat/models.py @@ -1,6 +1,6 @@ from typing import Union, Literal, Optional from enum import Enum -import warnings +import warnings from pydantic import BaseModel, Field, field_validator, model_validator @@ -9,16 +9,16 @@ 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") @@ -82,7 +82,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): @@ -98,8 +100,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): @@ -107,7 +110,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 + ) ChatMessage = Union[SAPUserMessage, SAPAssistantMessage, SAPToolChatMessage, SAPMessage] @@ -135,14 +140,14 @@ class KeyValueListPair(BaseModel): class DocumentMetadataKeyValueListPairs(KeyValueListPair): - select_mode: Optional[list[Literal['ignoreIfKeyAbsent']]] = None + select_mode: Optional[list[Literal["ignoreIfKeyAbsent"]]] = None class GroundingSearchConfig(BaseModel): max_chunk_count: Optional[int] = Field(default=None, ge=0) max_document_count: Optional[int] = Field(default=None, ge=0) - @model_validator(mode='after') + @model_validator(mode="after") def validate_max_chunk_count_and_max_document_count(self): if self.max_chunk_count is not None and self.max_document_count is not None: raise ValueError("Cannot specify both maxChunkCount and maxDocumentCount.") @@ -153,7 +158,7 @@ class DocumentGroundingFilter(BaseModel): id_: Optional[str] = Field(default=None, alias="id") data_repository_type: Literal["vector", "help.sap.com"] search_config: Optional[GroundingSearchConfig] = None - data_repositories: Optional[list[str]]= None + data_repositories: Optional[list[str]] = None data_repository_metadata: Optional[list[KeyValueListPair]] = None document_metadata: Optional[list[DocumentMetadataKeyValueListPairs]] = None chunk_metadata: Optional[list[KeyValueListPair]] = None @@ -171,7 +176,9 @@ class DocumentGroundingConfig(BaseModel): class GroundingModuleConfig(BaseModel): - type_: Literal["document_grounding_service"] = Field(default="document_grounding_service", alias="type") + type_: Literal["document_grounding_service"] = Field( + default="document_grounding_service", alias="type" + ) config: DocumentGroundingConfig @@ -284,6 +291,7 @@ class DPIMethodConstant(BaseModel): """ Replaces the entity with the specified value followed by an incrementing number """ + method: Literal["constant"] = "constant" value: str @@ -292,6 +300,7 @@ class DPIMethodFabricatedData(BaseModel): """ Replaces the entity with a randomly generated value appropriate to its type. """ + method: Literal["fabricated_data"] = "fabricated_data" @@ -300,6 +309,7 @@ class DPICustomEntity(BaseModel): regex: Regular expression to match the entity replacement_strategy: Replacement strategy to be used for the entity """ + regex: str replacement_strategy: DPIMethodConstant @@ -309,8 +319,11 @@ class DPIStandardEntity(BaseModel): type: Standard entity type to be masked replacement_strategy: Replacement strategy to be used for the entity """ + type_: SAPMaskingProfileEntity = Field(..., alias="type") - replacement_strategy: Optional[Union[DPIMethodConstant, DPIMethodFabricatedData]] = None + replacement_strategy: Optional[ + Union[DPIMethodConstant, DPIMethodFabricatedData] + ] = None class MaskGroundingInput(BaseModel): @@ -318,6 +331,7 @@ class MaskGroundingInput(BaseModel): Controls whether the input to the grounding module will be masked with the configuration supplied in the masking module """ + enabled: bool = False @@ -338,6 +352,7 @@ class MaskingProviderConfig(BaseModel): mask_grounding_input: A flag indicating whether to mask input to the grounding module. """ + type_: str = Field(default="sap_data_privacy_integration", alias="type") method: Literal["anonymization", "pseudonymization"] entities: list[Union[DPIStandardEntity, DPICustomEntity]] @@ -355,12 +370,14 @@ class MaskingModuleConfig(BaseModel): IMPORTANT: use exactly one of the parameters to set the list of masking provider configurations. DEPRECATED: parameter 'masking_providers' will be removed Sept 15, 2026. Use 'providers' instead. """ + providers: Optional[list[MaskingProviderConfig]] = Field(min_length=1, default=None) - masking_providers: Optional[list[MaskingProviderConfig]] = Field(min_length=1, default=None) + masking_providers: Optional[list[MaskingProviderConfig]] = Field( + min_length=1, default=None + ) @model_validator(mode="after") def enforce_exactly_one_provider_list(self): - has_providers = self.providers is not None has_masking_providers = self.masking_providers is not None @@ -443,7 +460,8 @@ class AzureContentSafetyInput(AzureContentFilter): self_harm: Threshold for self-harm content. prompt_shield: A flag to use prompt shield - """ + """ + prompt_shield: Optional[bool] = False @@ -539,29 +557,34 @@ class FilteringStreamOptions(BaseModel): overlap: Number of characters that should be additionally sent to content filtering services from previous chunks as additional context. """ + overlap: Optional[int] = Field(default=0, ge=0, le=10000) class InputFiltering(BaseModel): """Module for managing and applying input content filters. - Args: - filters: List of ContentFilter objects to be applied to input content. + Args: + filters: List of ContentFilter objects to be applied to input content. """ - filters: list[Union[AzureContentSafetyInputFilterConfig, LlamaGuard38bFilterConfig]] = Field(min_length=1) + + filters: list[ + Union[AzureContentSafetyInputFilterConfig, LlamaGuard38bFilterConfig] + ] = Field(min_length=1) class OutputFiltering(BaseModel): """Module for managing and applying output content filters. - Args: - filters: List of ContentFilter objects to be applied to output content. + Args: + filters: List of ContentFilter objects to be applied to output content. - stream_options: Module-specific streaming options. + stream_options: Module-specific streaming options. """ filters: list[ - Union[AzureContentSafetyOutputFilterConfig, LlamaGuard38bFilterConfig]] = Field(min_length=1) + Union[AzureContentSafetyOutputFilterConfig, LlamaGuard38bFilterConfig] + ] = Field(min_length=1) stream_options: Optional[FilteringStreamOptions] = None @@ -573,6 +596,7 @@ class FilteringModuleConfig(BaseModel): output: Module for filtering and validating output content after generation. """ + input: Optional[InputFiltering] = None output: Optional[OutputFiltering] = None @@ -598,6 +622,7 @@ class SAPDocumentTranslationApplyToSelector(BaseModel): targets the value of "user_input" in placeholder_values specified in the request payload; and considers the value to be in German. """ + category: Literal["placeholders", "template_roles"] items: list[str] source_language: str @@ -612,6 +637,7 @@ class InputTranslationConfig(BaseModel): target_language: Language to which the text should be translated. Example: en-US apply_to: List of selectors that define the scope of translation. """ + source_language: Optional[str] = None target_language: str apply_to: Optional[list[SAPDocumentTranslationApplyToSelector]] = None @@ -633,6 +659,7 @@ class SAPDocumentTranslationInput(BaseModel): config: Configuration object for the translation module. """ + type_: str = Field(default="sap_document_translation", alias="type") translate_messages_history: Optional[bool] = None config: InputTranslationConfig @@ -647,6 +674,7 @@ class SAPDocumentTranslationOutput(BaseModel): config: Configuration object for the translation module. """ + type_: str = Field(default="sap_document_translation", alias="type") config: OutputTranslationConfig @@ -660,6 +688,7 @@ class TranslationModuleConfig(BaseModel): output: Configuration for output translation """ + input: Optional[SAPDocumentTranslationInput] = None output: Optional[SAPDocumentTranslationOutput] = None diff --git a/litellm/llms/sap/chat/transformation.py b/litellm/llms/sap/chat/transformation.py index 5020656808a..8a2e75a8ed3 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 @@ -21,7 +31,12 @@ else: from ..credentials import get_token_creator from .models import ResponseFormatJSONSchema, ResponseFormat, OrchestrationRequest -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, exclude_unset=True) @@ -69,16 +84,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, @@ -90,14 +104,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: @@ -161,7 +174,6 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): params.remove("tool_choice") return params - def validate_environment( self, headers: dict, @@ -177,13 +189,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_ @@ -191,7 +203,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, @@ -220,7 +232,9 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): resp_type = response_format.get("type", None) if resp_type: if resp_type == "json_schema": - response_format = validate_dict(response_format, ResponseFormatJSONSchema) + response_format = validate_dict( + response_format, ResponseFormatJSONSchema + ) else: response_format = validate_dict(response_format, ResponseFormat) response_format = {"response_format": response_format} @@ -228,7 +242,9 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): response_format = {} placeholder_defaults = params.pop("placeholder_defaults", {}) - placeholder_defaults = {"defaults": placeholder_defaults} if placeholder_defaults else {} + placeholder_defaults = ( + {"defaults": placeholder_defaults} if placeholder_defaults else {} + ) optional_modules = {} optional_modules_lst = ["grounding", "masking", "filtering", "translation"] @@ -265,7 +281,9 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): stream_config["delimiters"] = stream_options.get("delimiters") placeholder_values = optional_params.pop("placeholder_values", {}) - placeholder_values = {"placeholder_values": placeholder_values} if placeholder_values else {} + placeholder_values = ( + {"placeholder_values": placeholder_values} if placeholder_values else {} + ) fallback_modules = optional_params.pop("fallback_sap_modules", []) @@ -304,7 +322,7 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): request_body = { "config": { "modules": modules_payload, - **({"stream": stream_config} if stream_config else {}) + **({"stream": stream_config} if stream_config else {}), }, **placeholder_values, } @@ -314,18 +332,18 @@ class GenAIHubOrchestrationConfig(OpenAIGPTConfig): return body 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, @@ -359,17 +377,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 e10bcbf7eae..aeae51bf0bb 100644 --- a/litellm/llms/sap/credentials.py +++ b/litellm/llms/sap/credentials.py @@ -180,7 +180,9 @@ def _resolve_value( return cred.default -def fetch_credentials(service_key: Optional[str] = None, profile: Optional[str] = None, **kwargs) -> Dict[str, str]: +def fetch_credentials( + service_key: Optional[str] = None, profile: Optional[str] = None, **kwargs +) -> Dict[str, str]: """ Resolution order per key: kwargs @@ -196,8 +198,11 @@ def fetch_credentials(service_key: Optional[str] = None, profile: Optional[str] if not config: # Prefer AICORE_SERVICE_KEY if present; otherwise fall back to the VCAP service. - service_like = service_key or sap_service_key or _load_json_env(SERVICE_KEY_ENV_VAR) or _get_vcap_service( - VCAP_AICORE_SERVICE_NAME + service_like = ( + service_key + or sap_service_key + or _load_json_env(SERVICE_KEY_ENV_VAR) + or _get_vcap_service(VCAP_AICORE_SERVICE_NAME) ) out: Dict[str, str] = {} @@ -241,7 +246,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") client_id = credentials.get("client_id") diff --git a/litellm/llms/sap/embed/transformation.py b/litellm/llms/sap/embed/transformation.py index 5d45a69e138..1d1f2ae967d 100644 --- a/litellm/llms/sap/embed/transformation.py +++ b/litellm/llms/sap/embed/transformation.py @@ -52,9 +52,11 @@ class EmbeddingModel(BaseModel): timeout: Optional[int] = Field(default=None, ge=1, le=600) max_retries: Optional[int] = Field(default=None, ge=0, le=5) + class EmbeddingsModelConfig(BaseModel): model: EmbeddingModel + class EmbeddingsModules(BaseModel): embeddings: EmbeddingsModelConfig masking: Optional[MaskingModuleConfig] = None @@ -64,9 +66,11 @@ class EmbeddingInput(BaseModel): text: Union[str, List[str]] type: Optional[Literal["text", "document", "query"]] = None + class EmbeddingConfig(BaseModel): modules: EmbeddingsModules + class EmbeddingRequest(BaseModel): config: EmbeddingConfig input: EmbeddingInput @@ -170,12 +174,7 @@ class GenAIHubEmbeddingConfig(BaseEmbeddingConfig): masking = optional_params.get("masking") masking = {"masking": masking} if masking is not None else {} body = { - "config": { - "modules": { - "embeddings": {"model": model_dict}, - **masking - } - }, + "config": {"modules": {"embeddings": {"model": model_dict}, **masking}}, "input": input_dict, } body = validate_dict(body, EmbeddingRequest)