(sap) run black formater

This commit is contained in:
Vasilisa Parshikova 2026-03-13 15:48:48 +04:00 committed by Sameer Kankute
parent f2dd120128
commit 341b559652
No known key found for this signature in database
3 changed files with 61 additions and 27 deletions

View file

@ -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
@ -141,14 +141,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.")
@ -159,7 +159,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
@ -177,7 +177,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
@ -290,6 +292,7 @@ class DPIMethodConstant(BaseModel):
"""
Replaces the entity with the specified value followed by an incrementing number
"""
method: Literal["constant"] = "constant"
value: str
@ -298,6 +301,7 @@ class DPIMethodFabricatedData(BaseModel):
"""
Replaces the entity with a randomly generated value appropriate to its type.
"""
method: Literal["fabricated_data"] = "fabricated_data"
@ -306,6 +310,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
@ -315,8 +320,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):
@ -324,6 +332,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
@ -344,6 +353,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]]
@ -361,12 +371,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
@ -449,7 +461,8 @@ class AzureContentSafetyInput(AzureContentFilter):
self_harm: Threshold for self-harm content.
prompt_shield: A flag to use prompt shield
"""
"""
prompt_shield: Optional[bool] = False
@ -545,29 +558,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
@ -579,6 +597,7 @@ class FilteringModuleConfig(BaseModel):
output: Module for filtering and validating output content after generation.
"""
input: Optional[InputFiltering] = None
output: Optional[OutputFiltering] = None
@ -604,6 +623,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
@ -618,6 +638,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
@ -639,6 +660,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
@ -653,6 +675,7 @@ class SAPDocumentTranslationOutput(BaseModel):
config: Configuration object for the translation module.
"""
type_: str = Field(default="sap_document_translation", alias="type")
config: OutputTranslationConfig
@ -666,6 +689,7 @@ class TranslationModuleConfig(BaseModel):
output: Configuration for output translation
"""
input: Optional[SAPDocumentTranslationInput] = None
output: Optional[SAPDocumentTranslationOutput] = None

View file

@ -31,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)
@ -227,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}
@ -235,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"]
@ -272,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", [])
@ -311,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,
}

View file

@ -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)