Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
reka.py441 linesDownload Raw Back to chat_models
1import json2from typing import (3    Any,4    AsyncIterator,5    Callable,6    Dict,7    Iterator,8    List,9    Literal,10    Mapping,11    Optional,12    Sequence,13    Type,14    Union,15)16 17from langchain_core.callbacks import (18    AsyncCallbackManagerForLLMRun,19    CallbackManagerForLLMRun,20)21from langchain_core.language_models import LanguageModelInput22from langchain_core.language_models.chat_models import (23    BaseChatModel,24    agenerate_from_stream,25    generate_from_stream,26)27from langchain_core.messages import (28    AIMessage,29    AIMessageChunk,30    BaseMessage,31    HumanMessage,32    SystemMessage,33    ToolMessage,34)35from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult36from langchain_core.runnables import Runnable37from langchain_core.tools import BaseTool38from langchain_core.utils import get_from_dict_or_env39from langchain_core.utils.function_calling import convert_to_openai_tool40from pydantic import BaseModel, ConfigDict, Field, model_validator41 42DEFAULT_REKA_MODEL = "reka-flash"43 44ContentType = Union[str, List[Union[str, Dict[str, Any]]]]45 46 47def process_content_item(item: Dict[str, Any]) -> Dict[str, Any]:48    """Process a single content item."""49    if item["type"] == "image_url":50        image_url = item["image_url"]51        if isinstance(image_url, dict) and "url" in image_url:52            # If it's in LangChain format, extract the URL value53            item["image_url"] = image_url["url"]54    return item55 56 57def process_content(content: ContentType) -> List[Dict[str, Any]]:58    """Process content to handle both text and media inputs,59    returning a list of content items."""60    if isinstance(content, str):61        return [{"type": "text", "text": content}]62    elif isinstance(content, list):63        result = []64        for item in content:65            if isinstance(item, str):66                result.append({"type": "text", "text": item})67            elif isinstance(item, dict):68                result.append(process_content_item(item))69            else:70                raise ValueError(f"Invalid content item format: {item}")71        return result72    else:73        raise ValueError("Invalid content format")74 75 76def convert_to_reka_messages(messages: List[BaseMessage]) -> List[Dict[str, Any]]:77    """Convert LangChain messages to Reka message format."""78    reka_messages: List[Dict[str, Any]] = []79    system_message: Optional[str] = None80 81    for message in messages:82        if isinstance(message, SystemMessage):83            if system_message is None:84                if isinstance(message.content, str):85                    system_message = message.content86                else:87                    raise TypeError("SystemMessage content must be a string.")88            else:89                raise ValueError("Multiple system messages are not supported.")90        elif isinstance(message, HumanMessage):91            processed_content = process_content(message.content)92            if system_message:93                if (94                    processed_content95                    and isinstance(processed_content[0], dict)96                    and processed_content[0].get("type") == "text"97                    and "text" in processed_content[0]98                ):99                    processed_content[0]["text"] = (100                        f"{system_message}\n{processed_content[0]['text']}"101                    )102                else:103                    processed_content.insert(104                        0, {"type": "text", "text": system_message}105                    )106                system_message = None107            reka_messages.append({"role": "user", "content": processed_content})108        elif isinstance(message, AIMessage):109            reka_message: Dict[str, Any] = {"role": "assistant"}110            if message.content:111                processed_content = process_content(message.content)112                reka_message["content"] = processed_content113            if "tool_calls" in message.additional_kwargs:114                tool_calls = message.additional_kwargs["tool_calls"]115                formatted_tool_calls = []116                for tool_call in tool_calls:117                    formatted_tool_call = {118                        "id": tool_call["id"],119                        "name": tool_call["function"]["name"],120                        "parameters": json.loads(tool_call["function"]["arguments"]),121                    }122                    formatted_tool_calls.append(formatted_tool_call)123                reka_message["tool_calls"] = formatted_tool_calls124            reka_messages.append(reka_message)125        elif isinstance(message, ToolMessage):126            content_list: List[Dict[str, Any]] = []127            content_list.append(128                {129                    "tool_call_id": message.tool_call_id,130                    "output": json.dumps({"status": message.content}),131                }132            )133            reka_messages.append(134                {135                    "role": "tool_output",136                    "content": content_list,137                }138            )139        else:140            raise ValueError(f"Unsupported message type: {type(message)}")141 142    return reka_messages143 144 145class ChatReka(BaseChatModel):146    """Reka chat large language models."""147 148    client: Any = None  #: :meta private:149    async_client: Any = None  #: :meta private:150    model: str = Field(default=DEFAULT_REKA_MODEL)151    max_tokens: int = Field(default=256)152    temperature: Optional[float] = None153    streaming: bool = False154    default_request_timeout: Optional[float] = None155    max_retries: int = 2156    reka_api_key: Optional[str] = None157    model_kwargs: Dict[str, Any] = Field(default_factory=dict)158    model_config = ConfigDict(extra="forbid")159    token_counter: Optional[160        Callable[[Union[str, BaseMessage, List[BaseMessage]]], int]161    ] = None162 163    @model_validator(mode="before")164    @classmethod165    def validate_environment(cls, values: Dict[str, Any]) -> Dict[str, Any]:166        """Validate that API key and Python package exist in the environment."""167        reka_api_key = values.get("reka_api_key")168        reka_api_key = get_from_dict_or_env(169            {"reka_api_key": reka_api_key}, "reka_api_key", "REKA_API_KEY"170        )171        values["reka_api_key"] = reka_api_key172 173        try:174            # Import reka libraries here175            from reka.client import AsyncReka, Reka176 177            values["client"] = Reka(178                api_key=reka_api_key,179            )180            values["async_client"] = AsyncReka(181                api_key=reka_api_key,182            )183        except ImportError:184            raise ImportError(185                "Could not import Reka Python package. "186                "Please install it with `pip install reka-api`."187            )188        return values189 190    @property191    def _default_params(self) -> Mapping[str, Any]:192        """Get the default parameters for calling Reka API."""193        params = {194            "model": self.model,195            "max_tokens": self.max_tokens,196        }197        if self.temperature is not None:198            params["temperature"] = self.temperature199        return {**params, **self.model_kwargs}200 201    @property202    def _llm_type(self) -> str:203        """Return type of chat model."""204        return "reka-chat"205 206    def _stream(207        self,208        messages: List[BaseMessage],209        stop: Optional[List[str]] = None,210        run_manager: Optional[CallbackManagerForLLMRun] = None,211        **kwargs: Any,212    ) -> Iterator[ChatGenerationChunk]:213        reka_messages = convert_to_reka_messages(messages)214        params = {**self._default_params, **kwargs}215        if stop:216            params["stop"] = stop217 218        stream = self.client.chat.create_stream(messages=reka_messages, **params)219 220        for chunk in stream:221            content = chunk.responses[0].chunk.content222            chat_chunk = ChatGenerationChunk(message=AIMessageChunk(content=content))223            if run_manager:224                run_manager.on_llm_new_token(content, chunk=chat_chunk)225            yield chat_chunk226 227    async def _astream(228        self,229        messages: List[BaseMessage],230        stop: Optional[List[str]] = None,231        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,232        **kwargs: Any,233    ) -> AsyncIterator[ChatGenerationChunk]:234        reka_messages = convert_to_reka_messages(messages)235        params = {**self._default_params, **kwargs}236        if stop:237            params["stop"] = stop238 239        stream = self.async_client.chat.create_stream(messages=reka_messages, **params)240 241        async for chunk in stream:242            content = chunk.responses[0].chunk.content243            chat_chunk = ChatGenerationChunk(message=AIMessageChunk(content=content))244            if run_manager:245                await run_manager.on_llm_new_token(content, chunk=chat_chunk)246            yield chat_chunk247 248    def _generate(249        self,250        messages: List[BaseMessage],251        stop: Optional[List[str]] = None,252        run_manager: Optional[CallbackManagerForLLMRun] = None,253        **kwargs: Any,254    ) -> ChatResult:255        if self.streaming:256            return generate_from_stream(257                self._stream(messages, stop=stop, run_manager=run_manager, **kwargs)258            )259 260        reka_messages = convert_to_reka_messages(messages)261        params = {**self._default_params, **kwargs}262        if stop:263            params["stop"] = stop264        response = self.client.chat.create(messages=reka_messages, **params)265 266        if response.responses[0].message.tool_calls:267            tool_calls = response.responses[0].message.tool_calls268            message = AIMessage(269                content="",  # Empty string instead of None270                additional_kwargs={271                    "tool_calls": [272                        {273                            "id": tc.id,274                            "type": "function",275                            "function": {276                                "name": tc.name,277                                "arguments": json.dumps(tc.parameters),278                            },279                        }280                        for tc in tool_calls281                    ]282                },283            )284        else:285            content = response.responses[0].message.content286            # Ensure content is never None287            message = AIMessage(content=content if content is not None else "")288 289        return ChatResult(generations=[ChatGeneration(message=message)])290 291    async def _agenerate(292        self,293        messages: List[BaseMessage],294        stop: Optional[List[str]] = None,295        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,296        **kwargs: Any,297    ) -> ChatResult:298        if self.streaming:299            return await agenerate_from_stream(300                self._astream(messages, stop=stop, run_manager=run_manager, **kwargs)301            )302 303        reka_messages = convert_to_reka_messages(messages)304        params = {**self._default_params, **kwargs}305        if stop:306            params["stop"] = stop307        response = await self.async_client.chat.create(messages=reka_messages, **params)308 309        if response.responses[0].message.tool_calls:310            tool_calls = response.responses[0].message.tool_calls311            message = AIMessage(312                content="",  # Empty string instead of None313                additional_kwargs={314                    "tool_calls": [315                        {316                            "id": tc.id,317                            "type": "function",318                            "function": {319                                "name": tc.name,320                                "arguments": json.dumps(tc.parameters),321                            },322                        }323                        for tc in tool_calls324                    ]325                },326            )327        else:328            content = response.responses[0].message.content329            # Ensure content is never None330            message = AIMessage(content=content if content is not None else "")331 332        return ChatResult(generations=[ChatGeneration(message=message)])333 334    def get_num_tokens(self, input: Union[str, BaseMessage, List[BaseMessage]]) -> int:335        """Calculate number of tokens.336 337        Args:338            input: Either a string, a single BaseMessage, or a list of BaseMessages.339 340        Returns:341            int: Number of tokens in the input.342 343        Raises:344            ImportError: If tiktoken is not installed.345            ValueError: If message content is not a string.346        """347        if self.token_counter is not None:348            return self.token_counter(input)349 350        try:351            import tiktoken352        except ImportError:353            raise ImportError(354                "Could not import tiktoken python package. "355                "Please install it with `pip install tiktoken`."356            )357 358        encoding = tiktoken.get_encoding("cl100k_base")359 360        if isinstance(input, str):361            return len(encoding.encode(input))362        elif isinstance(input, BaseMessage):363            content = input.content364            if not isinstance(content, str):365                raise ValueError(366                    f"Message content must be a string, got {type(content)}"367                )368            return len(encoding.encode(content))369        elif isinstance(input, list):370            total = 0371            for msg in input:372                content = msg.content373                if not isinstance(content, str):374                    raise ValueError(375                        f"Message content must be a string, got {type(content)}"376                    )377                total += len(encoding.encode(content))378            return total379        else:380            raise TypeError(f"Unsupported input type: {type(input)}")381 382    def bind_tools(383        self,384        tools: Sequence[Union[Dict[str, Any], Type[BaseModel], Callable, BaseTool]],385        *,386        tool_choice: Optional[Union[str, Literal["any"]]] = "auto",387        strict: Optional[bool] = None,388        **kwargs: Any,389    ) -> Runnable[LanguageModelInput, AIMessage]:390        """Bind tool-like objects to this chat model.391 392        The `tool_choice` parameter controls how the model uses the tools you pass.393        There are three available options:394 395        - `"auto"`: Lets the model decide whether or not to invoke a tool. This is the396          recommended way to do function calling with our models.397        - `"none"`: Disables tool calling. In this case, even if you pass tools to398          the model, the model will not invoke any tools.399        - `"tool"`: Forces the model to invoke one or more of the tools it has400          been passed.401 402        Args:403            tools: A list of tool definitions to bind to this chat model.404                Supports any tool definition handled by405                :meth:`langchain_core.utils.function_calling.convert_to_openai_tool`.406            tool_choice: Controls how the model uses the tools you pass.407                Options are "auto", "none", or "tool". Defaults to "auto".408            strict:409            If True, model output is guaranteed to exactly match the JSON Schema410                provided in the tool definition.411                If False, input schema will not be validated412                and model output will not be validated.413                If None, ``strict`` argument will not414                be passed to the model.415            kwargs: Any additional parameters are passed directly to the model.416 417        Returns:418            Runnable: An executable chain or component.419        """420        formatted_tools = [421            convert_to_openai_tool(tool, strict=strict) for tool in tools422        ]423 424        # Ensure tool_choice is one of the allowed options425        if tool_choice is None:426            tool_choice = "auto"427        if tool_choice == "any":428            tool_choice = "tool"429        if tool_choice not in ("auto", "none", "tool"):430            raise ValueError(431                f"Invalid tool_choice '{tool_choice}' provided. "432                "Tool choice must be one of: 'auto', 'none', or 'tool'."433            )434 435        # Map tool_choice to the parameter expected by the Reka API436        kwargs["tool_choice"] = tool_choice437 438        # Pass the tools and updated kwargs to the model439        formatted_tools = [tool["function"] for tool in formatted_tools]440        return super().bind(tools=formatted_tools, **kwargs)441 
codekingpro/portable-devtools · Team Ai