fix(cloudflare): parse OpenAI-compatible Workers AI responses

This commit is contained in:
Piyush Aryan 2026-05-31 03:54:50 +05:30 • committed by Piyush Aryan
parent 4c3efe9c7c
commit b9f9a8ed55
2 changed files with 229 additions and 35 deletions

View file

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

View file

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