From 5c796b436512d1f9addb303c634db1842d9116a3 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Sat, 6 Apr 2024 17:53:06 -0700 Subject: [PATCH] async streaming anthropic --- litellm/llms/anthropic.py | 78 +++++++++++++++++++++-- litellm/llms/custom_httpx/http_handler.py | 19 ++++-- 2 files changed, 87 insertions(+), 10 deletions(-) diff --git a/litellm/llms/anthropic.py b/litellm/llms/anthropic.py index db41ae6e358..47e485ecb86 100644 --- a/litellm/llms/anthropic.py +++ b/litellm/llms/anthropic.py @@ -9,8 +9,6 @@ import litellm from .prompt_templates.factory import prompt_factory, custom_prompt from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler -async_handler = AsyncHTTPHandler() - import httpx @@ -18,6 +16,11 @@ class AnthropicConstants(Enum): HUMAN_PROMPT = "\n\nHuman: " AI_PROMPT = "\n\nAssistant: " + # constants from https://github.com/anthropics/anthropic-sdk-python/blob/main/src/anthropic/_constants.py + + +async_handler = AsyncHTTPHandler(timeout=httpx.Timeout(timeout=600.0, connect=5.0)) + class AnthropicError(Exception): def __init__(self, status_code, message): @@ -230,6 +233,42 @@ def process_response( return model_response +async def acompletion_stream_function( + model: str, + messages: list, + api_base: str, + custom_prompt_dict: dict, + model_response: ModelResponse, + print_verbose: Callable, + encoding, + api_key, + logging_obj, + stream, + _is_function_call, + data=None, + optional_params=None, + litellm_params=None, + logger_fn=None, + headers={}, +): + response = await async_handler.post( + api_base, headers=headers, data=json.dumps(data) + ) + + if response.status_code != 200: + raise AnthropicError(status_code=response.status_code, message=response.text) + + completion_stream = response.aiter_lines() + + streamwrapper = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider="anthropic", + logging_obj=logging_obj, + ) + return streamwrapper + + async def acompletion_function( model: str, messages: list, @@ -356,8 +395,29 @@ def completion( ) print_verbose(f"_is_function_call: {_is_function_call}") if acompletion == True: - if optional_params.get("stream", False): - pass + if ( + stream and not _is_function_call + ): # 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 + return acompletion_stream_function( + model=model, + messages=messages, + data=data, + api_base=api_base, + custom_prompt_dict=custom_prompt_dict, + model_response=model_response, + print_verbose=print_verbose, + encoding=encoding, + api_key=api_key, + logging_obj=logging_obj, + optional_params=optional_params, + stream=stream, + _is_function_call=_is_function_call, + litellm_params=litellm_params, + logger_fn=logger_fn, + headers=headers, + ) else: return acompletion_function( model=model, @@ -396,7 +456,15 @@ def completion( status_code=response.status_code, message=response.text ) - return response.iter_lines() + completion_stream = response.iter_lines() + streaming_response = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider="anthropic", + logging_obj=logging_obj, + ) + return streaming_response + else: response = requests.post(api_base, headers=headers, data=json.dumps(data)) if response.status_code != 200: diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 51723a2f995..c008b059330 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1,15 +1,21 @@ import httpx, asyncio -from typing import Optional, Union +from typing import Optional, Union, Mapping, Any + +# https://www.python-httpx.org/advanced/timeouts +_DEFAULT_TIMEOUT = httpx.Timeout(timeout=5.0, connect=5.0) class AsyncHTTPHandler: - def __init__(self, concurrent_limit=1000): + def __init__( + self, timeout: httpx.Timeout = _DEFAULT_TIMEOUT, concurrent_limit=1000 + ): # Create a client with a connection pool self.client = httpx.AsyncClient( + timeout=timeout, limits=httpx.Limits( max_connections=concurrent_limit, max_keepalive_connections=concurrent_limit, - ) + ), ) async def close(self): @@ -25,12 +31,15 @@ class AsyncHTTPHandler: async def post( self, url: str, - data: Optional[Union[dict, str]] = None, + data: Optional[Union[dict, str]] = None, # type: ignore params: Optional[dict] = None, headers: Optional[dict] = None, ): response = await self.client.post( - url, data=data, params=params, headers=headers + url, + data=data, # type: ignore + params=params, + headers=headers, ) return response