mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
(sap) run black formater
This commit is contained in:
parent
a105cb336d
commit
4c816c53ad
5 changed files with 127 additions and 77 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue