codekingpro/portable-devtools
114k
1from enum import Enum2from typing import TYPE_CHECKING, Any, Dict, Iterator, List, Mapping, Optional3 4from langchain_core.embeddings import Embeddings5from langchain_core.utils import pre_init6from pydantic import BaseModel, ConfigDict7 8if TYPE_CHECKING:9 import oci10 11CUSTOM_ENDPOINT_PREFIX = "ocid1.generativeaiendpoint"12 13 14class OCIAuthType(Enum):15 """OCI authentication types as enumerator."""16 17 API_KEY = 118 SECURITY_TOKEN = 219 INSTANCE_PRINCIPAL = 320 RESOURCE_PRINCIPAL = 421 22 23class OCIGenAIEmbeddings(BaseModel, Embeddings):24 """OCI embedding models.25 26 To authenticate, the OCI client uses the methods described in27 https://docs.oracle.com/en-us/iaas/Content/API/Concepts/sdk_authentication_methods.htm28 29 The authentifcation method is passed through auth_type and should be one of:30 API_KEY (default), SECURITY_TOKEN, INSTANCE_PRINCIPLE, RESOURCE_PRINCIPLE31 32 Make sure you have the required policies (profile/roles) to33 access the OCI Generative AI service. If a specific config profile is used,34 you must pass the name of the profile (~/.oci/config) through auth_profile.35 If a specific config file location is used, you must pass36 the file location where profile name configs present37 through auth_file_location38 39 To use, you must provide the compartment id40 along with the endpoint url, and model id41 as named parameters to the constructor.42 43 Example:44 .. code-block:: python45 46 from langchain_classic.embeddings import OCIGenAIEmbeddings47 48 embeddings = OCIGenAIEmbeddings(49 model_id="MY_EMBEDDING_MODEL",50 service_endpoint="https://inference.generativeai.us-chicago-1.oci.oraclecloud.com",51 compartment_id="MY_OCID"52 )53 """54 55 client: Any = None #: :meta private:56 57 service_models: Any = None #: :meta private:58 59 auth_type: Optional[str] = "API_KEY"60 """Authentication type, could be 61 62 API_KEY, 63 SECURITY_TOKEN, 64 INSTANCE_PRINCIPLE, 65 RESOURCE_PRINCIPLE66 67 If not specified, API_KEY will be used68 """69 70 auth_profile: Optional[str] = "DEFAULT"71 """The name of the profile in ~/.oci/config72 If not specified , DEFAULT will be used 73 """74 75 auth_file_location: Optional[str] = "~/.oci/config"76 """Path to the config file.77 If not specified, ~/.oci/config will be used78 """79 80 model_id: Optional[str] = None81 """Id of the model to call, e.g., cohere.embed-english-light-v2.0"""82 83 model_kwargs: Optional[Dict] = None84 """Keyword arguments to pass to the model"""85 86 service_endpoint: Optional[str] = None87 """service endpoint url"""88 89 compartment_id: Optional[str] = None90 """OCID of compartment"""91 92 truncate: Optional[str] = "END"93 """Truncate embeddings that are too long from start or end ("NONE"|"START"|"END")"""94 95 batch_size: int = 9696 """Batch size of OCI GenAI embedding requests. OCI GenAI may handle up to 96 texts97 per request"""98 99 model_config = ConfigDict(extra="forbid", protected_namespaces=())100 101 @pre_init102 def validate_environment(cls, values: Dict) -> Dict: # pylint: disable=no-self-argument103 """Validate that OCI config and python package exists in environment."""104 105 # Skip creating new client if passed in constructor106 if values["client"] is not None:107 return values108 109 try:110 import oci111 112 client_kwargs = {113 "config": {},114 "signer": None,115 "service_endpoint": values["service_endpoint"],116 "retry_strategy": oci.retry.DEFAULT_RETRY_STRATEGY,117 "timeout": (10, 240), # default timeout config for OCI Gen AI service118 }119 120 if values["auth_type"] == OCIAuthType(1).name:121 client_kwargs["config"] = oci.config.from_file(122 file_location=values["auth_file_location"],123 profile_name=values["auth_profile"],124 )125 client_kwargs.pop("signer", None)126 elif values["auth_type"] == OCIAuthType(2).name:127 128 def make_security_token_signer(129 oci_config: dict[str, Any],130 ) -> "oci.auth.signers.SecurityTokenSigner":131 pk = oci.signer.load_private_key_from_file(132 oci_config.get("key_file"), None133 )134 with open(135 str(oci_config.get("security_token_file")), encoding="utf-8"136 ) as f:137 st_string = f.read()138 return oci.auth.signers.SecurityTokenSigner(st_string, pk)139 140 client_kwargs["config"] = oci.config.from_file(141 file_location=values["auth_file_location"],142 profile_name=values["auth_profile"],143 )144 client_kwargs["signer"] = make_security_token_signer(145 oci_config=client_kwargs["config"]146 )147 elif values["auth_type"] == OCIAuthType(3).name:148 client_kwargs["signer"] = (149 oci.auth.signers.InstancePrincipalsSecurityTokenSigner()150 )151 elif values["auth_type"] == OCIAuthType(4).name:152 client_kwargs["signer"] = (153 oci.auth.signers.get_resource_principals_signer()154 )155 else:156 raise ValueError("Please provide valid value to auth_type")157 158 values["client"] = oci.generative_ai_inference.GenerativeAiInferenceClient(159 **client_kwargs160 )161 162 except ImportError as ex:163 raise ImportError(164 "Could not import oci python package. "165 "Please make sure you have the oci package installed."166 ) from ex167 except Exception as e:168 raise ValueError(169 """Could not authenticate with OCI client.170 If INSTANCE_PRINCIPAL or RESOURCE_PRINCIPAL is used,171 please check the specified172 auth_profile, auth_file_location and auth_type are valid.""",173 e,174 ) from e175 176 return values177 178 @property179 def _identifying_params(self) -> Mapping[str, Any]:180 """Get the identifying parameters."""181 _model_kwargs = self.model_kwargs or {}182 return {183 **{"model_kwargs": _model_kwargs},184 }185 186 def embed_documents(self, texts: List[str]) -> List[List[float]]:187 """Call out to OCIGenAI's embedding endpoint.188 189 Args:190 texts: The list of texts to embed.191 192 Returns:193 List of embeddings, one for each text.194 """195 from oci.generative_ai_inference import models196 197 if not self.model_id:198 raise ValueError("Model ID is required to embed documents")199 200 if self.model_id.startswith(CUSTOM_ENDPOINT_PREFIX):201 serving_mode = models.DedicatedServingMode(endpoint_id=self.model_id)202 else:203 serving_mode = models.OnDemandServingMode(model_id=self.model_id)204 205 embeddings = []206 207 def split_texts() -> Iterator[List[str]]:208 for i in range(0, len(texts), self.batch_size):209 yield texts[i : i + self.batch_size]210 211 for chunk in split_texts():212 invocation_obj = models.EmbedTextDetails(213 serving_mode=serving_mode,214 compartment_id=self.compartment_id,215 truncate=self.truncate,216 inputs=chunk,217 )218 response = self.client.embed_text(invocation_obj)219 embeddings.extend(response.data.embeddings)220 221 return embeddings222 223 def embed_query(self, text: str) -> List[float]:224 """Call out to OCIGenAI's embedding endpoint.225 226 Args:227 text: The text to embed.228 229 Returns:230 Embeddings for the text.231 """232 return self.embed_documents([text])[0]233 