mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(vertex_httpx.py): fix streaming for cloudflare proxy calls
This commit is contained in:
parent
2a7592d026
commit
4f32f283a3
3 changed files with 92 additions and 22 deletions
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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: ''
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue