Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
bedrock.py918 linesDownload Raw Back to llms
1import asyncio2import json3import warnings4from abc import ABC5from typing import (6    Any,7    AsyncGenerator,8    AsyncIterator,9    Dict,10    Iterator,11    List,12    Mapping,13    Optional,14    Tuple,15)16 17from langchain_core._api.deprecation import deprecated18from langchain_core.callbacks import (19    AsyncCallbackManagerForLLMRun,20    CallbackManagerForLLMRun,21)22from langchain_core.language_models.llms import LLM23from langchain_core.outputs import GenerationChunk24from langchain_core.utils import get_from_dict_or_env, pre_init25from pydantic import BaseModel, ConfigDict, Field26 27from langchain_community.llms.utils import enforce_stop_tokens28from langchain_community.utilities.anthropic import (29    get_num_tokens_anthropic,30    get_token_ids_anthropic,31)32 33AMAZON_BEDROCK_TRACE_KEY = "amazon-bedrock-trace"34GUARDRAILS_BODY_KEY = "amazon-bedrock-guardrailAssessment"35HUMAN_PROMPT = "\n\nHuman:"36ASSISTANT_PROMPT = "\n\nAssistant:"37ALTERNATION_ERROR = (38    "Error: Prompt must alternate between '\n\nHuman:' and '\n\nAssistant:'."39)40 41 42def _add_newlines_before_ha(input_text: str) -> str:43    new_text = input_text44    for word in ["Human:", "Assistant:"]:45        new_text = new_text.replace(word, "\n\n" + word)46        for i in range(2):47            new_text = new_text.replace("\n\n\n" + word, "\n\n" + word)48    return new_text49 50 51def _human_assistant_format(input_text: str) -> str:52    if input_text.count("Human:") == 0 or (53        input_text.find("Human:") > input_text.find("Assistant:")54        and "Assistant:" in input_text55    ):56        input_text = HUMAN_PROMPT + " " + input_text  # SILENT CORRECTION57    if input_text.count("Assistant:") == 0:58        input_text = input_text + ASSISTANT_PROMPT  # SILENT CORRECTION59    if input_text[: len("Human:")] == "Human:":60        input_text = "\n\n" + input_text61    input_text = _add_newlines_before_ha(input_text)62    count = 063    # track alternation64    for i in range(len(input_text)):65        if input_text[i : i + len(HUMAN_PROMPT)] == HUMAN_PROMPT:66            if count % 2 == 0:67                count += 168            else:69                warnings.warn(ALTERNATION_ERROR + f" Received {input_text}")70        if input_text[i : i + len(ASSISTANT_PROMPT)] == ASSISTANT_PROMPT:71            if count % 2 == 1:72                count += 173            else:74                warnings.warn(ALTERNATION_ERROR + f" Received {input_text}")75 76    if count % 2 == 1:  # Only saw Human, no Assistant77        input_text = input_text + ASSISTANT_PROMPT  # SILENT CORRECTION78 79    return input_text80 81 82def _stream_response_to_generation_chunk(83    stream_response: Dict[str, Any],84) -> GenerationChunk:85    """Convert a stream response to a generation chunk."""86    if not stream_response["delta"]:87        return GenerationChunk(text="")88    return GenerationChunk(89        text=stream_response["delta"]["text"],90        generation_info=dict(91            finish_reason=stream_response.get("stop_reason", None),92        ),93    )94 95 96class LLMInputOutputAdapter:97    """Adapter class to prepare the inputs from Langchain to a format98    that LLM model expects.99 100    It also provides helper function to extract101    the generated text from the model response."""102 103    provider_to_output_key_map = {104        "anthropic": "completion",105        "amazon": "outputText",106        "cohere": "text",107        "meta": "generation",108        "mistral": "outputs",109    }110 111    @classmethod112    def prepare_input(113        cls,114        provider: str,115        model_kwargs: Dict[str, Any],116        prompt: Optional[str] = None,117        system: Optional[str] = None,118        messages: Optional[List[Dict]] = None,119    ) -> Dict[str, Any]:120        input_body = {**model_kwargs}121        if provider == "anthropic":122            if messages:123                input_body["anthropic_version"] = "bedrock-2023-05-31"124                input_body["messages"] = messages125                if system:126                    input_body["system"] = system127                if "max_tokens" not in input_body:128                    input_body["max_tokens"] = 1024129            if prompt:130                input_body["prompt"] = _human_assistant_format(prompt)131                if "max_tokens_to_sample" not in input_body:132                    input_body["max_tokens_to_sample"] = 1024133        elif provider in ("ai21", "cohere", "meta", "mistral"):134            input_body["prompt"] = prompt135        elif provider == "amazon":136            input_body = dict()137            input_body["inputText"] = prompt138            input_body["textGenerationConfig"] = {**model_kwargs}139        else:140            input_body["inputText"] = prompt141 142        return input_body143 144    @classmethod145    def prepare_output(cls, provider: str, response: Any) -> dict:146        text = ""147        if provider == "anthropic":148            response_body = json.loads(response.get("body").read().decode())149            if "completion" in response_body:150                text = response_body.get("completion")151            elif "content" in response_body:152                content = response_body.get("content")153                text = content[0].get("text")154        else:155            response_body = json.loads(response.get("body").read())156 157            if provider == "ai21":158                text = response_body.get("completions")[0].get("data").get("text")159            elif provider == "cohere":160                text = response_body.get("generations")[0].get("text")161            elif provider == "meta":162                text = response_body.get("generation")163            elif provider == "mistral":164                text = response_body.get("outputs")[0].get("text")165            else:166                text = response_body.get("results")[0].get("outputText")167 168        headers = response.get("ResponseMetadata", {}).get("HTTPHeaders", {})169        prompt_tokens = int(headers.get("x-amzn-bedrock-input-token-count", 0))170        completion_tokens = int(headers.get("x-amzn-bedrock-output-token-count", 0))171        return {172            "text": text,173            "body": response_body,174            "usage": {175                "prompt_tokens": prompt_tokens,176                "completion_tokens": completion_tokens,177                "total_tokens": prompt_tokens + completion_tokens,178            },179        }180 181    @classmethod182    def prepare_output_stream(183        cls,184        provider: str,185        response: Any,186        stop: Optional[List[str]] = None,187        messages_api: bool = False,188    ) -> Iterator[GenerationChunk]:189        stream = response.get("body")190 191        if not stream:192            return193 194        if messages_api:195            output_key = "message"196        else:197            output_key = cls.provider_to_output_key_map.get(provider, "")198 199        if not output_key:200            raise ValueError(201                f"Unknown streaming response output key for provider: {provider}"202            )203 204        for event in stream:205            chunk = event.get("chunk")206            if not chunk:207                continue208 209            chunk_obj = json.loads(chunk.get("bytes").decode())210 211            if provider == "cohere" and (212                chunk_obj["is_finished"] or chunk_obj[output_key] == "<EOS_TOKEN>"213            ):214                return215 216            elif (217                provider == "mistral"218                and chunk_obj.get(output_key, [{}])[0].get("stop_reason", "") == "stop"219            ):220                return221 222            elif messages_api and (chunk_obj.get("type") == "content_block_stop"):223                return224 225            if messages_api and chunk_obj.get("type") in (226                "message_start",227                "content_block_start",228                "content_block_delta",229            ):230                if chunk_obj.get("type") == "content_block_delta":231                    chk = _stream_response_to_generation_chunk(chunk_obj)232                    yield chk233                else:234                    continue235            else:236                # chunk obj format varies with provider237                yield GenerationChunk(238                    text=(239                        chunk_obj[output_key]240                        if provider != "mistral"241                        else chunk_obj[output_key][0]["text"]242                    ),243                    generation_info={244                        GUARDRAILS_BODY_KEY: (245                            chunk_obj.get(GUARDRAILS_BODY_KEY)246                            if GUARDRAILS_BODY_KEY in chunk_obj247                            else None248                        ),249                    },250                )251 252    @classmethod253    async def aprepare_output_stream(254        cls, provider: str, response: Any, stop: Optional[List[str]] = None255    ) -> AsyncIterator[GenerationChunk]:256        stream = response.get("body")257 258        if not stream:259            return260 261        output_key = cls.provider_to_output_key_map.get(provider, None)262 263        if not output_key:264            raise ValueError(265                f"Unknown streaming response output key for provider: {provider}"266            )267 268        for event in stream:269            chunk = event.get("chunk")270            if not chunk:271                continue272 273            chunk_obj = json.loads(chunk.get("bytes").decode())274 275            if provider == "cohere" and (276                chunk_obj["is_finished"] or chunk_obj[output_key] == "<EOS_TOKEN>"277            ):278                return279 280            if (281                provider == "mistral"282                and chunk_obj.get(output_key, [{}])[0].get("stop_reason", "") == "stop"283            ):284                return285 286            yield GenerationChunk(287                text=(288                    chunk_obj[output_key]289                    if provider != "mistral"290                    else chunk_obj[output_key][0]["text"]291                )292            )293 294 295class BedrockBase(BaseModel, ABC):296    """Base class for Bedrock models."""297 298    model_config = ConfigDict(protected_namespaces=())299 300    client: Any = Field(exclude=True)  #: :meta private:301 302    region_name: Optional[str] = None303    """The aws region e.g., `us-west-2`. Fallsback to AWS_DEFAULT_REGION env variable304    or region specified in ~/.aws/config in case it is not provided here.305    """306 307    credentials_profile_name: Optional[str] = Field(default=None, exclude=True)308    """The name of the profile in the ~/.aws/credentials or ~/.aws/config files, which309    has either access keys or role information specified.310    If not specified, the default credential profile or, if on an EC2 instance,311    credentials from IMDS will be used.312    See: https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html313    """314 315    config: Any = None316    """An optional botocore.config.Config instance to pass to the client."""317 318    provider: Optional[str] = None319    """The model provider, e.g., amazon, cohere, ai21, etc. When not supplied, provider320    is extracted from the first part of the model_id e.g. 'amazon' in 321    'amazon.titan-text-express-v1'. This value should be provided for model ids that do322    not have the provider in them, e.g., custom and provisioned models that have an ARN323    associated with them."""324 325    model_id: str326    """Id of the model to call, e.g., amazon.titan-text-express-v1, this is327    equivalent to the modelId property in the list-foundation-models api. For custom and328    provisioned models, an ARN value is expected."""329 330    model_kwargs: Optional[Dict] = None331    """Keyword arguments to pass to the model."""332 333    endpoint_url: Optional[str] = None334    """Needed if you don't want to default to us-east-1 endpoint"""335 336    streaming: bool = False337    """Whether to stream the results."""338 339    provider_stop_sequence_key_name_map: Mapping[str, str] = {340        "anthropic": "stop_sequences",341        "amazon": "stopSequences",342        "ai21": "stop_sequences",343        "cohere": "stop_sequences",344        "mistral": "stop",345    }346 347    guardrails: Optional[Mapping[str, Any]] = {348        "id": None,349        "version": None,350        "trace": False,351    }352    """353    An optional dictionary to configure guardrails for Bedrock.354 355    This field 'guardrails' consists of two keys: 'id' and 'version',356    which should be strings, but are initialized to None. It's used to357    determine if specific guardrails are enabled and properly set.358 359    Type:360        Optional[Mapping[str, str]]: A mapping with 'id' and 'version' keys.361 362    Example:363    llm = Bedrock(model_id="<model_id>", client=<bedrock_client>,364                  model_kwargs={},365                  guardrails={366                        "id": "<guardrail_id>",367                        "version": "<guardrail_version>"})368 369    To enable tracing for guardrails, set the 'trace' key to True and pass a callback handler to the370    'run_manager' parameter of the 'generate', '_call' methods.371 372    Example:373    llm = Bedrock(model_id="<model_id>", client=<bedrock_client>,374                  model_kwargs={},375                  guardrails={376                        "id": "<guardrail_id>",377                        "version": "<guardrail_version>",378                        "trace": True},379                callbacks=[BedrockAsyncCallbackHandler()])380 381    [https://python.langchain.com/docs/modules/callbacks/] for more information on callback handlers.382 383    class BedrockAsyncCallbackHandler(AsyncCallbackHandler):384        async def on_llm_error(385            self,386            error: BaseException,387            **kwargs: Any,388        ) -> Any:389            reason = kwargs.get("reason")390            if reason == "GUARDRAIL_INTERVENED":391                ...Logic to handle guardrail intervention...392    """  # noqa: E501393 394    @pre_init395    def validate_environment(cls, values: Dict) -> Dict:396        """Validate that AWS credentials to and python package exists in environment."""397 398        # Skip creating new client if passed in constructor399        if values.get("client") is not None:400            return values401 402        try:403            import boto3404 405            if values["credentials_profile_name"] is not None:406                session = boto3.Session(profile_name=values["credentials_profile_name"])407            else:408                # use default credentials409                session = boto3.Session()410 411            values["region_name"] = get_from_dict_or_env(412                values,413                "region_name",414                "AWS_DEFAULT_REGION",415                default=session.region_name,416            )417 418            client_params = {}419            if values["region_name"]:420                client_params["region_name"] = values["region_name"]421            if values["endpoint_url"]:422                client_params["endpoint_url"] = values["endpoint_url"]423            if values["config"]:424                client_params["config"] = values["config"]425 426            values["client"] = session.client("bedrock-runtime", **client_params)427 428        except ImportError:429            raise ImportError(430                "Could not import boto3 python package. "431                "Please install it with `pip install boto3`."432            )433        except ValueError as e:434            raise ValueError(f"Error raised by bedrock service: {e}")435        except Exception as e:436            raise ValueError(437                "Could not load credentials to authenticate with AWS client. "438                "Please check that credentials in the specified "439                f"profile name are valid. Bedrock error: {e}"440            ) from e441 442        return values443 444    @property445    def _identifying_params(self) -> Mapping[str, Any]:446        """Get the identifying parameters."""447        _model_kwargs = self.model_kwargs or {}448        return {449            **{"model_kwargs": _model_kwargs},450        }451 452    def _get_provider(self) -> str:453        if self.provider:454            return self.provider455        if self.model_id.startswith("arn"):456            raise ValueError(457                "Model provider should be supplied when passing a model ARN as model_id"458            )459 460        return self.model_id.split(".")[0]461 462    @property463    def _model_is_anthropic(self) -> bool:464        return self._get_provider() == "anthropic"465 466    @property467    def _guardrails_enabled(self) -> bool:468        """469        Determines if guardrails are enabled and correctly configured.470        Checks if 'guardrails' is a dictionary with non-empty 'id' and 'version' keys.471        Checks if 'guardrails.trace' is true.472 473        Returns:474            bool: True if guardrails are correctly configured, False otherwise.475        Raises:476            TypeError: If 'guardrails' lacks 'id' or 'version' keys.477        """478        try:479            return (480                isinstance(self.guardrails, dict)481                and bool(self.guardrails["id"])482                and bool(self.guardrails["version"])483            )484 485        except KeyError as e:486            raise TypeError(487                "Guardrails must be a dictionary with 'id' and 'version' keys."488            ) from e489 490    def _get_guardrails_canonical(self) -> Dict[str, Any]:491        """492        The canonical way to pass in guardrails to the bedrock service493        adheres to the following format:494 495        "amazon-bedrock-guardrailDetails": {496            "guardrailId": "string",497            "guardrailVersion": "string"498        }499        """500        return {501            "amazon-bedrock-guardrailDetails": {502                "guardrailId": self.guardrails.get("id"),  # type: ignore[union-attr]503                "guardrailVersion": self.guardrails.get("version"),  # type: ignore[union-attr]504            }505        }506 507    def _prepare_input_and_invoke(508        self,509        prompt: Optional[str] = None,510        system: Optional[str] = None,511        messages: Optional[List[Dict]] = None,512        stop: Optional[List[str]] = None,513        run_manager: Optional[CallbackManagerForLLMRun] = None,514        **kwargs: Any,515    ) -> Tuple[str, Dict[str, Any]]:516        _model_kwargs = self.model_kwargs or {}517 518        provider = self._get_provider()519        params = {**_model_kwargs, **kwargs}520        if self._guardrails_enabled:521            params.update(self._get_guardrails_canonical())522        input_body = LLMInputOutputAdapter.prepare_input(523            provider=provider,524            model_kwargs=params,525            prompt=prompt,526            system=system,527            messages=messages,528        )529        body = json.dumps(input_body)530        accept = "application/json"531        contentType = "application/json"532 533        request_options = {534            "body": body,535            "modelId": self.model_id,536            "accept": accept,537            "contentType": contentType,538        }539 540        if self._guardrails_enabled:541            request_options["guardrail"] = "ENABLED"542            if self.guardrails.get("trace"):  # type: ignore[union-attr]543                request_options["trace"] = "ENABLED"544 545        try:546            response = self.client.invoke_model(**request_options)547 548            text, body, usage_info = LLMInputOutputAdapter.prepare_output(549                provider, response550            ).values()551 552        except Exception as e:553            raise ValueError(f"Error raised by bedrock service: {e}")554 555        if stop is not None:556            text = enforce_stop_tokens(text, stop)557 558        # Verify and raise a callback error if any intervention occurs or a signal is559        # sent from a Bedrock service,560        # such as when guardrails are triggered.561        services_trace = self._get_bedrock_services_signal(body)562 563        if services_trace.get("signal") and run_manager is not None:564            run_manager.on_llm_error(565                Exception(566                    f"Error raised by bedrock service: {services_trace.get('reason')}"567                ),568                **services_trace,569            )570 571        return text, usage_info572 573    def _get_bedrock_services_signal(self, body: dict) -> dict:574        """575        This function checks the response body for an interrupt flag or message that indicates576        whether any of the Bedrock services have intervened in the processing flow. It is577        primarily used to identify modifications or interruptions imposed by these services578        during the request-response cycle with a Large Language Model (LLM).579        """  # noqa: E501580 581        if (582            self._guardrails_enabled583            and self.guardrails.get("trace")  # type: ignore[union-attr]584            and self._is_guardrails_intervention(body)585        ):586            return {587                "signal": True,588                "reason": "GUARDRAIL_INTERVENED",589                "trace": body.get(AMAZON_BEDROCK_TRACE_KEY),590            }591 592        return {593            "signal": False,594            "reason": None,595            "trace": None,596        }597 598    def _is_guardrails_intervention(self, body: dict) -> bool:599        return body.get(GUARDRAILS_BODY_KEY) == "GUARDRAIL_INTERVENED"600 601    def _prepare_input_and_invoke_stream(602        self,603        prompt: Optional[str] = None,604        system: Optional[str] = None,605        messages: Optional[List[Dict]] = None,606        stop: Optional[List[str]] = None,607        run_manager: Optional[CallbackManagerForLLMRun] = None,608        **kwargs: Any,609    ) -> Iterator[GenerationChunk]:610        _model_kwargs = self.model_kwargs or {}611        provider = self._get_provider()612 613        if stop:614            if provider not in self.provider_stop_sequence_key_name_map:615                raise ValueError(616                    f"Stop sequence key name for {provider} is not supported."617                )618 619            # stop sequence from _generate() overrides620            # stop sequences in the class attribute621            _model_kwargs[self.provider_stop_sequence_key_name_map.get(provider)] = stop622 623        if provider == "cohere":624            _model_kwargs["stream"] = True625 626        params = {**_model_kwargs, **kwargs}627 628        if self._guardrails_enabled:629            params.update(self._get_guardrails_canonical())630 631        input_body = LLMInputOutputAdapter.prepare_input(632            provider=provider,633            prompt=prompt,634            system=system,635            messages=messages,636            model_kwargs=params,637        )638        body = json.dumps(input_body)639 640        request_options = {641            "body": body,642            "modelId": self.model_id,643            "accept": "application/json",644            "contentType": "application/json",645        }646 647        if self._guardrails_enabled:648            request_options["guardrail"] = "ENABLED"649            if self.guardrails.get("trace"):  # type: ignore[union-attr]650                request_options["trace"] = "ENABLED"651 652        try:653            response = self.client.invoke_model_with_response_stream(**request_options)654 655        except Exception as e:656            raise ValueError(f"Error raised by bedrock service: {e}")657 658        for chunk in LLMInputOutputAdapter.prepare_output_stream(659            provider, response, stop, True if messages else False660        ):661            # verify and raise callback error if any middleware intervened662            self._get_bedrock_services_signal(chunk.generation_info)  # type: ignore[arg-type]663 664            if run_manager is not None:665                run_manager.on_llm_new_token(chunk.text, chunk=chunk)666            yield chunk667 668    async def _aprepare_input_and_invoke_stream(669        self,670        prompt: str,671        stop: Optional[List[str]] = None,672        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,673        **kwargs: Any,674    ) -> AsyncIterator[GenerationChunk]:675        _model_kwargs = self.model_kwargs or {}676        provider = self._get_provider()677 678        if stop:679            if provider not in self.provider_stop_sequence_key_name_map:680                raise ValueError(681                    f"Stop sequence key name for {provider} is not supported."682                )683            _model_kwargs[self.provider_stop_sequence_key_name_map.get(provider)] = stop684 685        if provider == "cohere":686            _model_kwargs["stream"] = True687 688        params = {**_model_kwargs, **kwargs}689        input_body = LLMInputOutputAdapter.prepare_input(690            provider=provider, prompt=prompt, model_kwargs=params691        )692        body = json.dumps(input_body)693 694        response = await asyncio.get_running_loop().run_in_executor(695            None,696            lambda: self.client.invoke_model_with_response_stream(697                body=body,698                modelId=self.model_id,699                accept="application/json",700                contentType="application/json",701            ),702        )703 704        async for chunk in LLMInputOutputAdapter.aprepare_output_stream(705            provider, response, stop706        ):707            if run_manager is not None and asyncio.iscoroutinefunction(708                run_manager.on_llm_new_token709            ):710                await run_manager.on_llm_new_token(chunk.text, chunk=chunk)711            elif run_manager is not None:712                run_manager.on_llm_new_token(chunk.text, chunk=chunk)  # type: ignore[unused-coroutine]713            yield chunk714 715 716@deprecated(717    since="0.0.34", removal="1.0", alternative_import="langchain_aws.BedrockLLM"718)719class Bedrock(LLM, BedrockBase):720    """Bedrock models.721 722    To authenticate, the AWS client uses the following methods to723    automatically load credentials:724    https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html725 726    If a specific credential profile should be used, you must pass727    the name of the profile from the ~/.aws/credentials file that is to be used.728 729    Make sure the credentials / roles used have the required policies to730    access the Bedrock service.731    """732 733    """734    Example:735        .. code-block:: python736 737            from bedrock_langchain.bedrock_llm import BedrockLLM738 739            llm = BedrockLLM(740                credentials_profile_name="default",741                model_id="amazon.titan-text-express-v1",742                streaming=True743            )744 745    """746 747    @pre_init748    def validate_environment(cls, values: Dict) -> Dict:749        model_id = values["model_id"]750        if model_id.startswith("anthropic.claude-3"):751            raise ValueError(752                "Claude v3 models are not supported by this LLM."753                "Please use `from langchain_community.chat_models import BedrockChat` "754                "instead."755            )756        return super().validate_environment(values)757 758    @property759    def _llm_type(self) -> str:760        """Return type of llm."""761        return "amazon_bedrock"762 763    @classmethod764    def is_lc_serializable(cls) -> bool:765        """Return whether this model can be serialized by Langchain."""766        return True767 768    @classmethod769    def get_lc_namespace(cls) -> List[str]:770        """Get the namespace of the langchain object."""771        return ["langchain", "llms", "bedrock"]772 773    @property774    def lc_attributes(self) -> Dict[str, Any]:775        attributes: Dict[str, Any] = {}776 777        if self.region_name:778            attributes["region_name"] = self.region_name779 780        return attributes781 782    model_config = ConfigDict(783        extra="forbid",784    )785 786    def _stream(787        self,788        prompt: str,789        stop: Optional[List[str]] = None,790        run_manager: Optional[CallbackManagerForLLMRun] = None,791        **kwargs: Any,792    ) -> Iterator[GenerationChunk]:793        """Call out to Bedrock service with streaming.794 795        Args:796            prompt (str): The prompt to pass into the model797            stop (Optional[List[str]], optional): Stop sequences. These will798                override any stop sequences in the `model_kwargs` attribute.799                Defaults to None.800            run_manager (Optional[CallbackManagerForLLMRun], optional): Callback801                run managers used to process the output. Defaults to None.802 803        Returns:804            Iterator[GenerationChunk]: Generator that yields the streamed responses.805 806        Yields:807            Iterator[GenerationChunk]: Responses from the model.808        """809        return self._prepare_input_and_invoke_stream(810            prompt=prompt, stop=stop, run_manager=run_manager, **kwargs811        )812 813    def _call(814        self,815        prompt: str,816        stop: Optional[List[str]] = None,817        run_manager: Optional[CallbackManagerForLLMRun] = None,818        **kwargs: Any,819    ) -> str:820        """Call out to Bedrock service model.821 822        Args:823            prompt: The prompt to pass into the model.824            stop: Optional list of stop words to use when generating.825 826        Returns:827            The string generated by the model.828 829        Example:830            .. code-block:: python831 832                response = llm.invoke("Tell me a joke.")833        """834 835        if self.streaming:836            completion = ""837            for chunk in self._stream(838                prompt=prompt, stop=stop, run_manager=run_manager, **kwargs839            ):840                completion += chunk.text841            return completion842 843        text, _ = self._prepare_input_and_invoke(844            prompt=prompt, stop=stop, run_manager=run_manager, **kwargs845        )846        return text847 848    async def _astream(849        self,850        prompt: str,851        stop: Optional[List[str]] = None,852        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,853        **kwargs: Any,854    ) -> AsyncGenerator[GenerationChunk, None]:855        """Call out to Bedrock service with streaming.856 857        Args:858            prompt (str): The prompt to pass into the model859            stop (Optional[List[str]], optional): Stop sequences. These will860                override any stop sequences in the `model_kwargs` attribute.861                Defaults to None.862            run_manager (Optional[CallbackManagerForLLMRun], optional): Callback863                run managers used to process the output. Defaults to None.864 865        Yields:866            AsyncGenerator[GenerationChunk, None]: Generator that asynchronously yields867            the streamed responses.868        """869        async for chunk in self._aprepare_input_and_invoke_stream(870            prompt=prompt, stop=stop, run_manager=run_manager, **kwargs871        ):872            yield chunk873 874    async def _acall(875        self,876        prompt: str,877        stop: Optional[List[str]] = None,878        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,879        **kwargs: Any,880    ) -> str:881        """Call out to Bedrock service model.882 883        Args:884            prompt: The prompt to pass into the model.885            stop: Optional list of stop words to use when generating.886 887        Returns:888            The string generated by the model.889 890        Example:891            .. code-block:: python892 893                response = await llm._acall("Tell me a joke.")894        """895 896        if not self.streaming:897            raise ValueError("Streaming must be set to True for async operations. ")898 899        chunks = [900            chunk.text901            async for chunk in self._astream(902                prompt=prompt, stop=stop, run_manager=run_manager, **kwargs903            )904        ]905        return "".join(chunks)906 907    def get_num_tokens(self, text: str) -> int:908        if self._model_is_anthropic:909            return get_num_tokens_anthropic(text)910        else:911            return super().get_num_tokens(text)912 913    def get_token_ids(self, text: str) -> List[int]:914        if self._model_is_anthropic:915            return get_token_ids_anthropic(text)916        else:917            return super().get_token_ids(text)918 
codekingpro/portable-devtools · Team Ai