diff --git a/litellm/llms/cloudflare/chat/transformation.py b/litellm/llms/cloudflare/chat/transformation.py index 66e253f304d..042b4b8d639 100644 --- a/litellm/llms/cloudflare/chat/transformation.py +++ b/litellm/llms/cloudflare/chat/transformation.py @@ -1,6 +1,5 @@ -import json import time -from typing import AsyncIterator, Iterator, List, Optional, Union +from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, Union import httpx @@ -37,6 +36,43 @@ class CloudflareError(BaseLLMException): ) # Call the base class constructor with the parameters it needs +def _extract_cloudflare_chat_content(result: Dict[str, Any]) -> str: + # Legacy Cloudflare fields take precedence, and "" is a valid response. + if "response" in result and result["response"] is not None: + if isinstance(result["response"], str): + return result["response"] + raise CloudflareError( + status_code=500, + message=f"Unable to parse Cloudflare chat response. Invalid response field: {result}", + ) + + if "response_text" in result and result["response_text"] is not None: + if isinstance(result["response_text"], str): + return result["response_text"] + raise CloudflareError( + status_code=500, + message=f"Unable to parse Cloudflare chat response. Invalid response_text field: {result}", + ) + + if "choices" in result: + choices = result["choices"] + if ( + isinstance(choices, list) + and len(choices) > 0 + and isinstance(choices[0], dict) + ): + message = choices[0].get("message") + if isinstance(message, dict) and message.get("content") is not None: + content = message["content"] + if isinstance(content, str): + return content + + raise CloudflareError( + status_code=500, + message=f"Unable to parse Cloudflare chat response. Response result: {result}", + ) + + class CloudflareChatConfig(BaseConfig): max_tokens: Optional[int] = None stream: Optional[bool] = None @@ -149,9 +185,10 @@ class CloudflareChatConfig(BaseConfig): ) -> ModelResponse: completion_response = raw_response.json() - # Support both "response" and "response_text" keys (newer models like Nemotron use "response_text") result = completion_response["result"] - model_response.choices[0].message.content = result.get("response") if result.get("response") is not None else result.get("response_text", "") # type: ignore + model_response.choices[0].message.content = _extract_cloudflare_chat_content( # type: ignore + result=result + ) prompt_tokens = litellm.utils.get_token_count(messages=messages, model=model) completion_tokens = len( @@ -191,32 +228,43 @@ class CloudflareChatConfig(BaseConfig): class CloudflareChatResponseIterator(BaseModelResponseIterator): def chunk_parser(self, chunk: dict) -> GenericStreamingChunk: - try: - text = "" - tool_use: Optional[ChatCompletionToolCallChunk] = None - is_finished = False - finish_reason = "" - usage: Optional[ChatCompletionUsageBlock] = None - provider_specific_fields = None + text = "" + tool_use: Optional[ChatCompletionToolCallChunk] = None + is_finished = False + finish_reason = "" + usage: Optional[ChatCompletionUsageBlock] = None + provider_specific_fields = None - index = int(chunk.get("index", 0)) + index = int(chunk.get("index", 0)) - if "response" in chunk and chunk["response"] is not None: - text = chunk["response"] - elif "response_text" in chunk and chunk["response_text"] is not None: - text = chunk["response_text"] + if "response" in chunk and chunk["response"] is not None: + text = chunk["response"] + elif "response_text" in chunk and chunk["response_text"] is not None: + text = chunk["response_text"] + elif "choices" in chunk and isinstance(chunk["choices"], list): + choices = chunk["choices"] + if len(choices) > 0 and isinstance(choices[0], dict): + choice = choices[0] + index = int(choice.get("index", index)) + delta = choice.get("delta") + if isinstance(delta, dict) and delta.get("content") is not None: + text = delta["content"] + if choice.get("finish_reason") is not None: + is_finished = True + finish_reason = choice["finish_reason"] - returned_chunk = GenericStreamingChunk( - text=text, - tool_use=tool_use, - is_finished=is_finished, - finish_reason=finish_reason, - usage=usage, - index=index, - provider_specific_fields=provider_specific_fields, - ) + if not is_finished and chunk.get("finish_reason") is not None: + is_finished = True + finish_reason = chunk["finish_reason"] - return returned_chunk + returned_chunk = GenericStreamingChunk( + text=text, + tool_use=tool_use, + is_finished=is_finished, + finish_reason=finish_reason, + usage=usage, + index=index, + provider_specific_fields=provider_specific_fields, + ) - except json.JSONDecodeError: - raise ValueError(f"Failed to decode JSON from chunk: {chunk}") + return returned_chunk diff --git a/tests/llm_translation/test_cloudflare.py b/tests/llm_translation/test_cloudflare.py index 5a6a0008398..533eef00007 100644 --- a/tests/llm_translation/test_cloudflare.py +++ b/tests/llm_translation/test_cloudflare.py @@ -7,6 +7,7 @@ import httpx import pytest from litellm import acompletion, completion +from litellm.llms.cloudflare.chat.transformation import CloudflareChatResponseIterator from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler FAKE_API_BASE = ( @@ -24,11 +25,9 @@ def _make_mock_response(json_data: Dict[str, Any]) -> MagicMock: return mock -def _chat_response() -> Dict[str, Any]: +def _chat_response(result: Dict[str, Any]) -> Dict[str, Any]: return { - "result": { - "response": "I am a large language model created to assist you.", - }, + "result": result, "success": True, "errors": [], "messages": [], @@ -51,10 +50,33 @@ def _streaming_chunks_response_text() -> list[str]: ] +def _streaming_chunks_openai_compatible() -> list[str]: + return [ + json.dumps({"choices": [{"delta": {"content": "I am"}}]}), + json.dumps({"choices": [{"delta": {"content": " a language"}}]}), + json.dumps({"choices": [{"delta": {"content": " model."}}]}), + ] + + +@pytest.mark.parametrize( + ("result", "expected_content"), + [ + ( + {"response": "I am a large language model created to assist you."}, + "I am a large language model created to assist you.", + ), + ({"response_text": "I am a language model."}, "I am a language model."), + ( + {"choices": [{"message": {"role": "assistant", "content": "Hello"}}]}, + "Hello", + ), + ({"response": ""}, ""), + ], +) @pytest.mark.parametrize("sync_mode", [True, False]) -def test_completion_cloudflare(sync_mode): +def test_completion_cloudflare(sync_mode, result, expected_content): messages = [{"role": "user", "content": "what llm are you"}] - mock_resp = _make_mock_response(_chat_response()) + mock_resp = _make_mock_response(_chat_response(result=result)) if sync_mode: with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post: @@ -83,7 +105,7 @@ def test_completion_cloudflare(sync_mode): assert response is not None assert response.choices[0].message.content is not None - assert "language model" in response.choices[0].message.content.lower() + assert response.choices[0].message.content == expected_content @pytest.mark.parametrize("sync_mode", [True, False]) @@ -226,3 +248,127 @@ def test_completion_cloudflare_stream_response_text(sync_mode): if c.choices[0].delta.content ) assert "language" in content.lower() + + +@pytest.mark.parametrize("sync_mode", [True, False]) +def test_completion_cloudflare_stream_openai_compatible(sync_mode): + messages = [{"role": "user", "content": "what llm are you"}] + raw_chunks = _streaming_chunks_openai_compatible() + + if sync_mode: + + def _iter_lines(): + for chunk in raw_chunks: + yield f"data: {chunk}" + yield "data: [DONE]" + + mock_resp = MagicMock() + mock_resp.iter_lines.return_value = _iter_lines() + mock_resp.status_code = 200 + mock_resp.headers = {"content-type": "text/event-stream"} + + with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post: + response = completion( + model="cloudflare/@cf/meta/llama-3.1-8b-instruct", + messages=messages, + max_tokens=15, + stream=True, + api_base=FAKE_API_BASE, + api_key=FAKE_API_KEY, + ) + chunks_received = list(response) + mock_post.assert_called_once() + else: + + async def _aiter_lines(): + for chunk in raw_chunks: + yield f"data: {chunk}" + yield "data: [DONE]" + + mock_resp = MagicMock() + mock_resp.aiter_lines.return_value = _aiter_lines() + mock_resp.status_code = 200 + mock_resp.headers = {"content-type": "text/event-stream"} + + async def _run(): + with patch.object( + AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp + ) as mock_post: + resp = await acompletion( + model="cloudflare/@cf/meta/llama-3.1-8b-instruct", + messages=messages, + max_tokens=15, + stream=True, + api_base=FAKE_API_BASE, + api_key=FAKE_API_KEY, + ) + received = [] + async for chunk in resp: + received.append(chunk) + mock_post.assert_called_once() + return received + + chunks_received = asyncio.run(_run()) + + assert len(chunks_received) > 0 + content = "".join( + c.choices[0].delta.content + for c in chunks_received + if c.choices[0].delta.content + ) + assert "language" in content.lower() + + +@pytest.mark.parametrize( + ("raw_chunk", "expected_text"), + [ + ({"response": "hello"}, "hello"), + ({"response_text": "hello"}, "hello"), + ({"choices": [{"delta": {"content": "hello"}}]}, "hello"), + ], +) +def test_cloudflare_streaming_chunk_parser_content(raw_chunk, expected_text): + iterator = CloudflareChatResponseIterator( + streaming_response=[], + sync_stream=True, + ) + + parsed_chunk = iterator.chunk_parser(raw_chunk) + + assert parsed_chunk["text"] == expected_text + assert parsed_chunk["is_finished"] is False + + +def test_cloudflare_streaming_chunk_parser_finish_reason(): + iterator = CloudflareChatResponseIterator( + streaming_response=[], + sync_stream=True, + ) + + parsed_chunk = iterator.chunk_parser( + {"choices": [{"delta": {}, "finish_reason": "stop"}]} + ) + + assert parsed_chunk["text"] == "" + assert parsed_chunk["is_finished"] is True + assert parsed_chunk["finish_reason"] == "stop" + + +@pytest.mark.parametrize( + "raw_chunk", + [ + {"response": "last token", "finish_reason": "stop"}, + {"response_text": "last token", "finish_reason": "stop"}, + ], +) +def test_cloudflare_legacy_streaming_chunk_parser_finish_reason(raw_chunk): + iterator = CloudflareChatResponseIterator( + streaming_response=[], + sync_stream=True, + ) + + parsed_chunk = iterator.chunk_parser(raw_chunk) + + assert parsed_chunk["text"] == "last token" + assert parsed_chunk["is_finished"] is True + assert parsed_chunk["finish_reason"] == "stop"