async streaming anthropic

This commit is contained in:
Ishaan Jaff 2024-04-06 17:53:06 -07:00
parent 7849c29f70
commit 5c796b4365
2 changed files with 87 additions and 10 deletions

View file

@ -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:

View file

@ -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