From 4f32f283a3442b4abe73469f250a6a85bc517c68 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 29 Jun 2024 09:09:23 -0700 Subject: [PATCH] fix(vertex_httpx.py): fix streaming for cloudflare proxy calls --- litellm/llms/vertex_httpx.py | 66 ++++++++++++++----- litellm/proxy/_super_secret_config.yaml | 6 +- .../tests/test_amazing_vertex_completion.py | 42 ++++++++++++ 3 files changed, 92 insertions(+), 22 deletions(-) diff --git a/litellm/llms/vertex_httpx.py b/litellm/llms/vertex_httpx.py index a0cfc98e766..a6dcd3daa2b 100644 --- a/litellm/llms/vertex_httpx.py +++ b/litellm/llms/vertex_httpx.py @@ -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}") diff --git a/litellm/proxy/_super_secret_config.yaml b/litellm/proxy/_super_secret_config.yaml index c28bb490112..ede853094ed 100644 --- a/litellm/proxy/_super_secret_config.yaml +++ b/litellm/proxy/_super_secret_config.yaml @@ -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: '' diff --git a/litellm/tests/test_amazing_vertex_completion.py b/litellm/tests/test_amazing_vertex_completion.py index 901d68ef3d0..6de3e11b84c 100644 --- a/litellm/tests/test_amazing_vertex_completion.py +++ b/litellm/tests/test_amazing_vertex_completion.py @@ -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"])