fix(vertex_httpx.py): fix streaming for cloudflare proxy calls

This commit is contained in:
Krrish Dholakia 2024-06-29 09:09:23 -07:00
parent 2a7592d026
commit 4f32f283a3
3 changed files with 92 additions and 22 deletions

View file

@ -467,7 +467,7 @@ async def make_call(
raise VertexAIError(status_code=response.status_code, message=response.text)
completion_stream = ModelResponseIterator(
streaming_response=response.aiter_bytes(), sync_stream=False
streaming_response=response.aiter_lines(), sync_stream=False
)
# LOGGING
logging_obj.post_call(
@ -498,7 +498,7 @@ def make_sync_call(
raise VertexAIError(status_code=response.status_code, message=response.read())
completion_stream = ModelResponseIterator(
streaming_response=response.iter_bytes(), sync_stream=True
streaming_response=response.iter_lines(), sync_stream=True
)
# LOGGING
@ -1028,7 +1028,7 @@ class VertexLLM(BaseLLM):
data["generationConfig"] = generation_config
headers = {
"Content-Type": "application/json; charset=utf-8",
"Content-Type": "application/json",
}
if auth_header is not None:
headers["Authorization"] = f"Bearer {auth_header}"
@ -1310,9 +1310,9 @@ class ModelResponseIterator:
if "usageMetadata" in processed_chunk:
usage = ChatCompletionUsageBlock(
prompt_tokens=processed_chunk["usageMetadata"]["promptTokenCount"],
completion_tokens=processed_chunk["usageMetadata"][
"candidatesTokenCount"
],
completion_tokens=processed_chunk["usageMetadata"].get(
"candidatesTokenCount", 0
),
total_tokens=processed_chunk["usageMetadata"]["totalTokenCount"],
)
@ -1336,15 +1336,30 @@ class ModelResponseIterator:
def __next__(self):
try:
chunk = self.response_iterator.__next__()
chunk = chunk.decode()
chunk = chunk.replace("data:", "")
chunk = chunk.strip()
json_chunk = json.loads(chunk)
return self.chunk_parser(chunk=json_chunk)
except StopIteration:
raise StopIteration
except ValueError as e:
raise RuntimeError(f"Error parsing chunk: {e}")
raise RuntimeError(f"Error receiving chunk from stream: {e}")
try:
chunk = chunk.replace("data:", "")
chunk = chunk.strip()
if len(chunk) > 0:
json_chunk = json.loads(chunk)
return self.chunk_parser(chunk=json_chunk)
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):
@ -1354,12 +1369,27 @@ class ModelResponseIterator:
async def __anext__(self):
try:
chunk = await self.async_response_iterator.__anext__()
chunk = chunk.decode()
chunk = chunk.replace("data:", "")
chunk = chunk.strip()
json_chunk = json.loads(chunk)
return self.chunk_parser(chunk=json_chunk)
except StopAsyncIteration:
raise StopAsyncIteration
except ValueError as e:
raise RuntimeError(f"Error parsing chunk: {e}")
raise RuntimeError(f"Error receiving chunk from stream: {e}")
try:
chunk = chunk.replace("data:", "")
chunk = chunk.strip()
if len(chunk) > 0:
json_chunk = json.loads(chunk)
return self.chunk_parser(chunk=json_chunk)
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}")

View file

@ -4,10 +4,8 @@ model_list:
model: anthropic/claude-3-5-sonnet
- model_name: gemini-1.5-flash-gemini
litellm_params:
model: gemini/gemini-1.5-flash
- model_name: gemini-1.5-flash-gemini
litellm_params:
model: gemini/gemini-1.5-flash
model: vertex_ai_beta/gemini-1.5-flash
api_base: https://gateway.ai.cloudflare.com/v1/fa4cdcab1f32b95ca3b53fd36043d691/test/google-vertex-ai/v1/projects/adroit-crow-413218/locations/us-central1/publishers/google/models/gemini-1.5-flash
- litellm_params:
api_base: http://0.0.0.0:8080
api_key: ''

View file

@ -914,6 +914,48 @@ async def test_gemini_pro_httpx_custom_api_base(provider):
assert "hello" in mock_call.call_args.kwargs["headers"]
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.parametrize("provider", ["vertex_ai_beta"]) # "vertex_ai",
@pytest.mark.asyncio
async def test_gemini_pro_httpx_custom_api_base_streaming_real_call(
provider, sync_mode
):
load_vertex_ai_credentials()
import random
litellm.set_verbose = True
messages = [
{
"role": "user",
"content": "Hey, how's it going?",
}
]
vertex_region = random.sample(["asia-southeast1", "us-central1"], k=1)[0]
if sync_mode is True:
response = completion(
model="vertex_ai_beta/gemini-1.5-flash",
messages=messages,
api_base="https://gateway.ai.cloudflare.com/v1/fa4cdcab1f32b95ca3b53fd36043d691/test/google-vertex-ai/v1/projects/adroit-crow-413218/locations/us-central1/publishers/google/models/gemini-1.5-flash",
stream=True,
vertex_region=vertex_region,
)
for chunk in response:
print(chunk)
else:
response = await litellm.acompletion(
model="vertex_ai_beta/gemini-1.5-flash",
messages=messages,
api_base="https://gateway.ai.cloudflare.com/v1/fa4cdcab1f32b95ca3b53fd36043d691/test/google-vertex-ai/v1/projects/adroit-crow-413218/locations/us-central1/publishers/google/models/gemini-1.5-flash",
stream=True,
vertex_region=vertex_region,
)
async for chunk in response:
print(chunk)
@pytest.mark.skip(reason="exhausted vertex quota. need to refactor to mock the call")
@pytest.mark.parametrize("sync_mode", [True])
@pytest.mark.parametrize("provider", ["vertex_ai"])