From 344268e053a5b39509c3df6cd5bc03c4e86e7248 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 3 Jul 2024 19:48:35 -0700 Subject: [PATCH] fix(anthropic.py): support *real* anthropic tool calling + streaming Parses each chunk and translates to openai format --- litellm/llms/anthropic.py | 466 +++++++++++++++++++++----------- litellm/tests/test_streaming.py | 6 +- litellm/types/llms/anthropic.py | 114 +++++++- litellm/utils.py | 69 ++--- 4 files changed, 453 insertions(+), 202 deletions(-) diff --git a/litellm/llms/anthropic.py b/litellm/llms/anthropic.py index ce15dd359c5..6e3b246bff5 100644 --- a/litellm/llms/anthropic.py +++ b/litellm/llms/anthropic.py @@ -18,7 +18,20 @@ from litellm.llms.custom_httpx.http_handler import ( _get_async_httpx_client, _get_httpx_client, ) -from litellm.types.llms.anthropic import AnthropicMessagesToolChoice +from litellm.types.llms.anthropic import ( + AnthropicMessagesToolChoice, + ContentBlockDelta, + ContentBlockStart, + MessageBlockDelta, + MessageStartBlock, +) +from litellm.types.llms.openai import ( + ChatCompletionResponseMessage, + ChatCompletionToolCallChunk, + ChatCompletionToolCallFunctionChunk, + ChatCompletionUsageBlock, +) +from litellm.types.utils import GenericStreamingChunk from litellm.utils import CustomStreamWrapper, ModelResponse, Usage from .base import BaseLLM @@ -198,7 +211,9 @@ async def make_call( status_code=response.status_code, message=await response.aread() ) - completion_stream = response.aiter_lines() + completion_stream = ModelResponseIterator( + streaming_response=response.aiter_lines(), sync_stream=False + ) # LOGGING logging_obj.post_call( @@ -215,120 +230,120 @@ class AnthropicChatCompletion(BaseLLM): def __init__(self) -> None: super().__init__() - def process_streaming_response( - self, - model: str, - response: Union[requests.Response, httpx.Response], - model_response: ModelResponse, - stream: bool, - logging_obj: litellm.litellm_core_utils.litellm_logging.Logging, - optional_params: dict, - api_key: str, - data: Union[dict, str], - messages: List, - print_verbose, - encoding, - ) -> CustomStreamWrapper: - """ - Return stream object for tool-calling + streaming - """ - ## LOGGING - logging_obj.post_call( - input=messages, - api_key=api_key, - original_response=response.text, - additional_args={"complete_input_dict": data}, - ) - print_verbose(f"raw model_response: {response.text}") - ## RESPONSE OBJECT - try: - completion_response = response.json() - except: - raise AnthropicError( - message=response.text, status_code=response.status_code - ) - text_content = "" - tool_calls = [] - for content in completion_response["content"]: - if content["type"] == "text": - text_content += content["text"] - ## TOOL CALLING - elif content["type"] == "tool_use": - tool_calls.append( - { - "id": content["id"], - "type": "function", - "function": { - "name": content["name"], - "arguments": json.dumps(content["input"]), - }, - } - ) - if "error" in completion_response: - raise AnthropicError( - message=str(completion_response["error"]), - status_code=response.status_code, - ) - _message = litellm.Message( - tool_calls=tool_calls, - content=text_content or None, - ) - model_response.choices[0].message = _message # type: ignore - model_response._hidden_params["original_response"] = completion_response[ - "content" - ] # allow user to access raw anthropic tool calling response + # def process_streaming_response( + # self, + # model: str, + # response: Union[requests.Response, httpx.Response], + # model_response: ModelResponse, + # stream: bool, + # logging_obj: litellm.litellm_core_utils.litellm_logging.Logging, + # optional_params: dict, + # api_key: str, + # data: Union[dict, str], + # messages: List, + # print_verbose, + # encoding, + # ) -> CustomStreamWrapper: + # """ + # Return stream object for tool-calling + streaming + # """ + # ## LOGGING + # logging_obj.post_call( + # input=messages, + # api_key=api_key, + # original_response=response.text, + # additional_args={"complete_input_dict": data}, + # ) + # print_verbose(f"raw model_response: {response.text}") + # ## RESPONSE OBJECT + # try: + # completion_response = response.json() + # except: + # raise AnthropicError( + # message=response.text, status_code=response.status_code + # ) + # text_content = "" + # tool_calls = [] + # for content in completion_response["content"]: + # if content["type"] == "text": + # text_content += content["text"] + # ## TOOL CALLING + # elif content["type"] == "tool_use": + # tool_calls.append( + # { + # "id": content["id"], + # "type": "function", + # "function": { + # "name": content["name"], + # "arguments": json.dumps(content["input"]), + # }, + # } + # ) + # if "error" in completion_response: + # raise AnthropicError( + # message=str(completion_response["error"]), + # status_code=response.status_code, + # ) + # _message = litellm.Message( + # tool_calls=tool_calls, + # content=text_content or None, + # ) + # model_response.choices[0].message = _message # type: ignore + # model_response._hidden_params["original_response"] = completion_response[ + # "content" + # ] # allow user to access raw anthropic tool calling response - model_response.choices[0].finish_reason = map_finish_reason( - completion_response["stop_reason"] - ) + # model_response.choices[0].finish_reason = map_finish_reason( + # completion_response["stop_reason"] + # ) - print_verbose("INSIDE ANTHROPIC STREAMING TOOL CALLING CONDITION BLOCK") - # return an iterator - streaming_model_response = ModelResponse(stream=True) - streaming_model_response.choices[0].finish_reason = model_response.choices[ # type: ignore - 0 - ].finish_reason - # streaming_model_response.choices = [litellm.utils.StreamingChoices()] - streaming_choice = litellm.utils.StreamingChoices() - streaming_choice.index = model_response.choices[0].index - _tool_calls = [] - print_verbose( - f"type of model_response.choices[0]: {type(model_response.choices[0])}" - ) - print_verbose(f"type of streaming_choice: {type(streaming_choice)}") - if isinstance(model_response.choices[0], litellm.Choices): - if getattr( - model_response.choices[0].message, "tool_calls", None - ) is not None and isinstance( - model_response.choices[0].message.tool_calls, list - ): - for tool_call in model_response.choices[0].message.tool_calls: - _tool_call = {**tool_call.dict(), "index": 0} - _tool_calls.append(_tool_call) - delta_obj = litellm.utils.Delta( - content=getattr(model_response.choices[0].message, "content", None), - role=model_response.choices[0].message.role, - tool_calls=_tool_calls, - ) - streaming_choice.delta = delta_obj - streaming_model_response.choices = [streaming_choice] - completion_stream = ModelResponseIterator( - model_response=streaming_model_response - ) - print_verbose( - "Returns anthropic CustomStreamWrapper with 'cached_response' streaming object" - ) - return CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="cached_response", - logging_obj=logging_obj, - ) - else: - raise AnthropicError( - status_code=422, - message="Unprocessable response object - {}".format(response.text), - ) + # print_verbose("INSIDE ANTHROPIC STREAMING TOOL CALLING CONDITION BLOCK") + # # return an iterator + # streaming_model_response = ModelResponse(stream=True) + # streaming_model_response.choices[0].finish_reason = model_response.choices[ # type: ignore + # 0 + # ].finish_reason + # # streaming_model_response.choices = [litellm.utils.StreamingChoices()] + # streaming_choice = litellm.utils.StreamingChoices() + # streaming_choice.index = model_response.choices[0].index + # _tool_calls = [] + # print_verbose( + # f"type of model_response.choices[0]: {type(model_response.choices[0])}" + # ) + # print_verbose(f"type of streaming_choice: {type(streaming_choice)}") + # if isinstance(model_response.choices[0], litellm.Choices): + # if getattr( + # model_response.choices[0].message, "tool_calls", None + # ) is not None and isinstance( + # model_response.choices[0].message.tool_calls, list + # ): + # for tool_call in model_response.choices[0].message.tool_calls: + # _tool_call = {**tool_call.dict(), "index": 0} + # _tool_calls.append(_tool_call) + # delta_obj = litellm.utils.Delta( + # content=getattr(model_response.choices[0].message, "content", None), + # role=model_response.choices[0].message.role, + # tool_calls=_tool_calls, + # ) + # streaming_choice.delta = delta_obj + # streaming_model_response.choices = [streaming_choice] + # completion_stream = ModelResponseIterator( + # model_response=streaming_model_response + # ) + # print_verbose( + # "Returns anthropic CustomStreamWrapper with 'cached_response' streaming object" + # ) + # return CustomStreamWrapper( + # completion_stream=completion_stream, + # model=model, + # custom_llm_provider="cached_response", + # logging_obj=logging_obj, + # ) + # else: + # raise AnthropicError( + # status_code=422, + # message="Unprocessable response object - {}".format(response.text), + # ) def process_response( self, @@ -481,20 +496,6 @@ class AnthropicChatCompletion(BaseLLM): additional_args={"complete_input_dict": data}, ) raise e - if stream and _is_function_call: - return self.process_streaming_response( - model=model, - response=response, - model_response=model_response, - stream=stream, - logging_obj=logging_obj, - api_key=api_key, - data=data, - messages=messages, - print_verbose=print_verbose, - optional_params=optional_params, - encoding=encoding, - ) return self.process_response( model=model, response=response, @@ -607,7 +608,7 @@ class AnthropicChatCompletion(BaseLLM): print_verbose(f"_is_function_call: {_is_function_call}") if acompletion == True: if ( - stream and not _is_function_call + stream is True ): # if function call - fake the streaming (need complete blocks for output parsing in openai format) print_verbose("makes async anthropic streaming POST request") data["stream"] = stream @@ -651,7 +652,7 @@ class AnthropicChatCompletion(BaseLLM): else: ## COMPLETION CALL if ( - stream and not _is_function_call + stream is True ): # if function call - fake the streaming (need complete blocks for output parsing in openai format) print_verbose("makes anthropic streaming POST request") data["stream"] = stream @@ -667,7 +668,9 @@ class AnthropicChatCompletion(BaseLLM): status_code=response.status_code, message=response.text ) - completion_stream = response.iter_lines() + completion_stream = ModelResponseIterator( + streaming_response=response.iter_lines(), sync_stream=True + ) streaming_response = CustomStreamWrapper( completion_stream=completion_stream, model=model, @@ -702,20 +705,6 @@ class AnthropicChatCompletion(BaseLLM): status_code=response.status_code, message=response.text ) - if stream and _is_function_call: - return self.process_streaming_response( - model=model, - response=response, - model_response=model_response, - stream=stream, - logging_obj=logging_obj, - api_key=api_key, - data=data, - messages=messages, - print_verbose=print_verbose, - optional_params=optional_params, - encoding=encoding, - ) return self.process_response( model=model, response=response, @@ -736,26 +725,195 @@ class AnthropicChatCompletion(BaseLLM): class ModelResponseIterator: - def __init__(self, model_response): - self.model_response = model_response - self.is_done = False + def __init__(self, streaming_response, sync_stream: bool): + self.streaming_response = streaming_response + self.response_iterator = self.streaming_response + + def chunk_parser(self, chunk: dict) -> GenericStreamingChunk: + try: + type_chunk = chunk.get("type", "") or "" + + text = "" + tool_use: Optional[ChatCompletionToolCallChunk] = None + is_finished = False + finish_reason = "" + usage: Optional[ChatCompletionUsageBlock] = None + + index = int(chunk.get("index", 0)) + if type_chunk == "content_block_delta": + """ + Anthropic content chunk + chunk = {'type': 'content_block_delta', 'index': 0, 'delta': {'type': 'text_delta', 'text': 'Hello'}} + """ + content_block = ContentBlockDelta(**chunk) # type: ignore + if "text" in content_block["delta"]: + text = content_block["delta"]["text"] + elif "partial_json" in content_block["delta"]: + tool_use = { + "id": None, + "type": "function", + "function": { + "name": None, + "arguments": content_block["delta"]["partial_json"], + }, + } + elif type_chunk == "content_block_start": + """ + event: content_block_start + data: {"type":"content_block_start","index":1,"content_block":{"type":"tool_use","id":"toolu_01T1x1fJ34qAmk2tNTrN7Up6","name":"get_weather","input":{}}} + """ + content_block_start = ContentBlockStart(**chunk) # type: ignore + if content_block_start["content_block"]["type"] == "text": + text = content_block_start["content_block"]["text"] + elif content_block_start["content_block"]["type"] == "tool_use": + tool_use = { + "id": content_block_start["content_block"]["id"], + "type": "function", + "function": { + "name": content_block_start["content_block"]["name"], + "arguments": json.dumps( + content_block_start["content_block"]["input"] + ), + }, + } + elif type_chunk == "message_delta": + """ + Anthropic + chunk = {'type': 'message_delta', 'delta': {'stop_reason': 'max_tokens', 'stop_sequence': None}, 'usage': {'output_tokens': 10}} + """ + # TODO - get usage from this chunk, set in response + message_delta = MessageBlockDelta(**chunk) # type: ignore + finish_reason = map_finish_reason( + finish_reason=message_delta["delta"].get("stop_reason", "stop") + or "stop" + ) + usage = ChatCompletionUsageBlock( + prompt_tokens=message_delta["usage"].get("input_tokens", 0), + completion_tokens=message_delta["usage"].get("output_tokens", 0), + total_tokens=message_delta["usage"].get("input_tokens", 0) + + message_delta["usage"].get("output_tokens", 0), + ) + is_finished = True + elif type_chunk == "message_start": + """ + Anthropic + chunk = { + "type": "message_start", + "message": { + "id": "msg_vrtx_011PqREFEMzd3REdCoUFAmdG", + "type": "message", + "role": "assistant", + "model": "claude-3-sonnet-20240229", + "content": [], + "stop_reason": null, + "stop_sequence": null, + "usage": { + "input_tokens": 270, + "output_tokens": 1 + } + } + } + """ + message_start_block = MessageStartBlock(**chunk) # type: ignore + usage = ChatCompletionUsageBlock( + prompt_tokens=message_start_block["message"] + .get("usage", {}) + .get("input_tokens", 0), + completion_tokens=message_start_block["message"] + .get("usage", {}) + .get("output_tokens", 0), + total_tokens=message_start_block["message"] + .get("usage", {}) + .get("input_tokens", 0) + + message_start_block["message"] + .get("usage", {}) + .get("output_tokens", 0), + ) + returned_chunk = GenericStreamingChunk( + text=text, + tool_use=tool_use, + is_finished=is_finished, + finish_reason=finish_reason, + usage=usage, + index=index, + ) + + return returned_chunk + + except json.JSONDecodeError: + raise ValueError(f"Failed to decode JSON from chunk: {chunk}") # Sync iterator def __iter__(self): return self def __next__(self): - if self.is_done: + try: + chunk = self.response_iterator.__next__() + except StopIteration: raise StopIteration - self.is_done = True - return self.model_response + except ValueError as e: + raise RuntimeError(f"Error receiving chunk from stream: {e}") + + try: + str_line = chunk + if isinstance(chunk, bytes): # Handle binary data + str_line = chunk.decode("utf-8") # Convert bytes to string + index = str_line.find("data:") + if index != -1: + str_line = str_line[index:] + + if str_line.startswith("data:"): + data_json = json.loads(str_line[5:]) + return self.chunk_parser(chunk=data_json) + else: + return GenericStreamingChunk( + text="", + is_finished=False, + finish_reason="", + usage=None, + index=0, + tool_use=None, + ) + except StopIteration: + raise StopIteration + except ValueError as e: + raise RuntimeError(f"Error parsing chunk: {e},\nReceived chunk: {chunk}") # Async iterator def __aiter__(self): + self.async_response_iterator = self.streaming_response.__aiter__() return self async def __anext__(self): - if self.is_done: + try: + chunk = await self.async_response_iterator.__anext__() + except StopAsyncIteration: raise StopAsyncIteration - self.is_done = True - return self.model_response + except ValueError as e: + raise RuntimeError(f"Error receiving chunk from stream: {e}") + + try: + str_line = chunk + if isinstance(chunk, bytes): # Handle binary data + str_line = chunk.decode("utf-8") # Convert bytes to string + index = str_line.find("data:") + if index != -1: + str_line = str_line[index:] + + if str_line.startswith("data:"): + data_json = json.loads(str_line[5:]) + return self.chunk_parser(chunk=data_json) + else: + return GenericStreamingChunk( + text="", + is_finished=False, + finish_reason="", + usage=None, + index=0, + tool_use=None, + ) + except StopAsyncIteration: + raise StopAsyncIteration + except ValueError as e: + raise RuntimeError(f"Error parsing chunk: {e},\nReceived chunk: {chunk}") diff --git a/litellm/tests/test_streaming.py b/litellm/tests/test_streaming.py index 0dd81e3b34b..880059596a3 100644 --- a/litellm/tests/test_streaming.py +++ b/litellm/tests/test_streaming.py @@ -2746,7 +2746,7 @@ class Chunk2(BaseModel): object: str created: int model: str - system_fingerprint: str + system_fingerprint: Optional[str] choices: List[Choices2] @@ -3001,7 +3001,7 @@ def test_completion_claude_3_function_call_with_streaming(): model="claude-3-opus-20240229", messages=messages, tools=tools, - tool_choice="auto", + tool_choice="required", stream=True, ) idx = 0 @@ -3060,7 +3060,7 @@ async def test_acompletion_claude_3_function_call_with_streaming(): model="claude-3-opus-20240229", messages=messages, tools=tools, - tool_choice="auto", + tool_choice="required", stream=True, ) idx = 0 diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index ffe403f8bf6..8d8280ea794 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -1,7 +1,6 @@ -from typing import List, Optional, Union, Iterable +from typing import Iterable, List, Optional, Union from pydantic import BaseModel, validator - from typing_extensions import Literal, Required, TypedDict @@ -45,3 +44,114 @@ class AnthopicMessagesAssistantMessageParam(TypedDict, total=False): Provides the model information to differentiate between participants of the same role. """ + + +class ContentTextBlockDelta(TypedDict): + """ + 'delta': {'type': 'text_delta', 'text': 'Hello'} + """ + + type: str + text: str + + +class ContentJsonBlockDelta(TypedDict): + """ + "delta": {"type": "input_json_delta","partial_json": "{\"location\": \"San Fra"}} + """ + + type: str + partial_json: str + + +class ContentBlockDelta(TypedDict): + type: str + index: int + delta: Union[ContentTextBlockDelta, ContentJsonBlockDelta] + + +class ToolUseBlock(TypedDict): + """ + "content_block":{"type":"tool_use","id":"toolu_01T1x1fJ34qAmk2tNTrN7Up6","name":"get_weather","input":{}} + """ + + id: str + + input: dict + + name: str + + type: Literal["tool_use"] + + +class TextBlock(TypedDict): + text: str + + type: Literal["text"] + + +class ContentBlockStart(TypedDict): + """ + event: content_block_start + data: {"type":"content_block_start","index":1,"content_block":{"type":"tool_use","id":"toolu_01T1x1fJ34qAmk2tNTrN7Up6","name":"get_weather","input":{}}} + """ + + type: str + index: int + content_block: Union[ToolUseBlock, TextBlock] + + +class MessageDelta(TypedDict, total=False): + stop_reason: Optional[str] + + +class UsageDelta(TypedDict, total=False): + input_tokens: int + output_tokens: int + + +class MessageBlockDelta(TypedDict): + """ + Anthropic + chunk = {'type': 'message_delta', 'delta': {'stop_reason': 'max_tokens', 'stop_sequence': None}, 'usage': {'output_tokens': 10}} + """ + + type: Literal["message_delta"] + delta: MessageDelta + usage: UsageDelta + + +class MessageChunk(TypedDict, total=False): + id: str + type: str + role: str + model: str + content: List + stop_reason: Optional[str] + stop_sequence: Optional[str] + usage: UsageDelta + + +class MessageStartBlock(TypedDict): + """ + Anthropic + chunk = { + "type": "message_start", + "message": { + "id": "msg_vrtx_011PqREFEMzd3REdCoUFAmdG", + "type": "message", + "role": "assistant", + "model": "claude-3-sonnet-20240229", + "content": [], + "stop_reason": null, + "stop_sequence": null, + "usage": { + "input_tokens": 270, + "output_tokens": 1 + } + } + } + """ + + type: Literal["message_start"] + message: MessageChunk diff --git a/litellm/utils.py b/litellm/utils.py index 26f90fa57a7..a5f11937be1 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8003,6 +8003,11 @@ class CustomStreamWrapper: return hold, curr_chunk def handle_anthropic_text_chunk(self, chunk): + """ + For old anthropic models - claude-1, claude-2. + + Claude-3 is handled from within Anthropic.py VIA ModelResponseIterator() + """ str_line = chunk if isinstance(chunk, bytes): # Handle binary data str_line = chunk.decode("utf-8") # Convert bytes to string @@ -8031,48 +8036,6 @@ class CustomStreamWrapper: "finish_reason": finish_reason, } - def handle_anthropic_chunk(self, chunk): - str_line = chunk - if isinstance(chunk, bytes): # Handle binary data - str_line = chunk.decode("utf-8") # Convert bytes to string - index = str_line.find("data:") - if index != -1: - str_line = str_line[index:] - - text = "" - is_finished = False - finish_reason = None - if str_line.startswith("data:"): - data_json = json.loads(str_line[5:]) - type_chunk = data_json.get("type", None) - if type_chunk == "content_block_delta": - """ - Anthropic content chunk - chunk = {'type': 'content_block_delta', 'index': 0, 'delta': {'type': 'text_delta', 'text': 'Hello'}} - """ - text = data_json.get("delta", {}).get("text", "") - elif type_chunk == "message_delta": - """ - Anthropic - chunk = {'type': 'message_delta', 'delta': {'stop_reason': 'max_tokens', 'stop_sequence': None}, 'usage': {'output_tokens': 10}} - """ - # TODO - get usage from this chunk, set in response - finish_reason = data_json.get("delta", {}).get("stop_reason", None) - is_finished = True - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - elif "error" in str_line: - raise ValueError(f"Unable to parse response. Original response: {str_line}") - else: - return { - "text": text, - "is_finished": is_finished, - "finish_reason": finish_reason, - } - def handle_vertexai_anthropic_chunk(self, chunk): """ - MessageStartEvent(message=Message(id='msg_01LeRRgvX4gwkX3ryBVgtuYZ', content=[], model='claude-3-sonnet-20240229', role='assistant', stop_reason=None, stop_sequence=None, type='message', usage=Usage(input_tokens=8, output_tokens=1)), type='message_start'); custom_llm_provider: vertex_ai @@ -8823,10 +8786,30 @@ class CustomStreamWrapper: # return this for all models completion_obj = {"content": ""} if self.custom_llm_provider and self.custom_llm_provider == "anthropic": - response_obj = self.handle_anthropic_chunk(chunk) + from litellm.types.llms.bedrock import GenericStreamingChunk + + if self.received_finish_reason is not None: + raise StopIteration + response_obj: GenericStreamingChunk = chunk completion_obj["content"] = response_obj["text"] if response_obj["is_finished"]: self.received_finish_reason = response_obj["finish_reason"] + + if ( + self.stream_options + and self.stream_options.get("include_usage", False) is True + and response_obj["usage"] is not None + ): + self.sent_stream_usage = True + model_response.usage = litellm.Usage( + prompt_tokens=response_obj["usage"]["inputTokens"], + completion_tokens=response_obj["usage"]["outputTokens"], + total_tokens=response_obj["usage"]["totalTokens"], + ) + + if "tool_use" in response_obj and response_obj["tool_use"] is not None: + completion_obj["tool_calls"] = [response_obj["tool_use"]] + elif ( self.custom_llm_provider and self.custom_llm_provider == "anthropic_text"