fix(cloudflare): route native Workers AI provider through OpenAI-compatible endpoint

This commit is contained in:
mateo-berri 2026-06-22 19:36:00 -07:00
parent 6f6aec2930
commit 1815cc9bee
4 changed files with 261 additions and 286 deletions

View file

@ -1,26 +1,11 @@
import json from typing import List, Optional, Union
import time
from typing import AsyncIterator, Iterator, List, Optional, Union
import httpx import httpx
import litellm from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.litellm_core_utils.url_utils import encode_url_path_segments from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import (
BaseConfig,
BaseLLMException,
LiteLLMLoggingObj,
)
from litellm.secret_managers.main import get_secret_str from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import (
ChatCompletionToolCallChunk,
ChatCompletionUsageBlock,
GenericStreamingChunk,
ModelResponse,
Usage,
)
class CloudflareError(BaseLLMException): class CloudflareError(BaseLLMException):
@ -34,26 +19,32 @@ class CloudflareError(BaseLLMException):
message=message, message=message,
request=self.request, request=self.request,
response=self.response, response=self.response,
) # Call the base class constructor with the parameters it needs )
class CloudflareChatConfig(BaseConfig): class CloudflareChatConfig(OpenAIGPTConfig):
max_tokens: Optional[int] = None def get_complete_url(
stream: Optional[bool] = None
def __init__(
self, self,
max_tokens: Optional[int] = None, api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None, stream: Optional[bool] = None,
) -> None: ) -> str:
locals_ = locals().copy() if api_base is None:
for key, value in locals_.items(): account_id = get_secret_str("CLOUDFLARE_ACCOUNT_ID")
if key != "self" and value is not None: api_base = (
setattr(self.__class__, key, value) f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/v1"
)
@classmethod return super().get_complete_url(
def get_config(cls): api_base=api_base,
return super().get_config() api_key=api_key,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
stream=stream,
)
def validate_environment( def validate_environment(
self, self,
@ -67,107 +58,18 @@ class CloudflareChatConfig(BaseConfig):
) -> dict: ) -> dict:
if api_key is None: if api_key is None:
raise ValueError( raise ValueError(
"Missing CloudflareError API Key - A call is being made to cloudflare but no key is set either in the environment variables or via params" "Missing Cloudflare API Key - A call is being made to cloudflare but no key is set either in the environment variables or via params"
) )
headers = { return super().validate_environment(
"accept": "application/json", headers=headers,
"content-type": "apbplication/json", model=model,
"Authorization": "Bearer " + api_key, messages=messages,
} optional_params=optional_params,
return headers litellm_params=litellm_params,
api_key=api_key,
def get_complete_url( api_base=api_base,
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
if api_base is None:
account_id = get_secret_str("CLOUDFLARE_ACCOUNT_ID")
api_base = (
f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run/"
)
encoded_model = encode_url_path_segments(model, field_name="model")
return api_base + encoded_model
def get_supported_openai_params(self, model: str) -> List[str]:
return [
"stream",
"max_tokens",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
supported_openai_params = self.get_supported_openai_params(model=model)
for param, value in non_default_params.items():
if param == "max_completion_tokens":
optional_params["max_tokens"] = value
elif param in supported_openai_params:
optional_params[param] = value
return optional_params
def transform_request(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
config = litellm.CloudflareChatConfig.get_config()
for k, v in config.items():
if k not in optional_params:
optional_params[k] = v
data = {
"messages": messages,
**optional_params,
}
return data
def transform_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: str,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> 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
prompt_tokens = litellm.utils.get_token_count(messages=messages, model=model)
completion_tokens = len(
encoding.encode(model_response["choices"][0]["message"].get("content", ""))
) )
model_response.created = int(time.time())
model_response.model = "cloudflare/" + model
usage = Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
)
setattr(model_response, "usage", usage)
return model_response
def get_error_class( def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException: ) -> BaseLLMException:
@ -175,48 +77,3 @@ class CloudflareChatConfig(BaseConfig):
status_code=status_code, status_code=status_code,
message=error_message, message=error_message,
) )
def get_model_response_iterator(
self,
streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse],
sync_stream: bool,
json_mode: Optional[bool] = False,
):
return CloudflareChatResponseIterator(
streaming_response=streaming_response,
sync_stream=sync_stream,
json_mode=json_mode,
)
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
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"]
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,
)
return returned_chunk
except json.JSONDecodeError:
raise ValueError(f"Failed to decode JSON from chunk: {chunk}")

View file

@ -4241,7 +4241,7 @@ def completion( # type: ignore
api_base api_base
or litellm.api_base or litellm.api_base
or get_secret("CLOUDFLARE_API_BASE") or get_secret("CLOUDFLARE_API_BASE")
or f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run/" or f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/v1"
) )
custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict custom_prompt_dict = custom_prompt_dict or litellm.custom_prompt_dict

View file

@ -9,9 +9,7 @@ import pytest
from litellm import acompletion, completion from litellm import acompletion, completion
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
FAKE_API_BASE = ( FAKE_API_BASE = "https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/v1"
"https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/run/"
)
FAKE_API_KEY = "fake-cf-api-key" FAKE_API_KEY = "fake-cf-api-key"
@ -26,28 +24,78 @@ def _make_mock_response(json_data: Dict[str, Any]) -> MagicMock:
def _chat_response() -> Dict[str, Any]: def _chat_response() -> Dict[str, Any]:
return { return {
"result": { "id": "chatcmpl-cf",
"response": "I am a large language model created to assist you.", "object": "chat.completion",
}, "created": 1234567890,
"success": True, "model": "@cf/meta/llama-2-7b-chat-int8",
"errors": [], "choices": [
"messages": [], {
"index": 0,
"message": {
"role": "assistant",
"content": "I am a large language model created to assist you.",
},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 8, "completion_tokens": 11, "total_tokens": 19},
}
def _tool_call_response() -> Dict[str, Any]:
return {
"id": "chatcmpl-cf-tools",
"object": "chat.completion",
"created": 1234567890,
"model": "@cf/meta/llama-2-7b-chat-int8",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"city": "New York"}',
},
}
],
},
"finish_reason": "tool_calls",
}
],
"usage": {"prompt_tokens": 20, "completion_tokens": 9, "total_tokens": 29},
} }
def _streaming_chunks() -> list[str]: def _streaming_chunks() -> list[str]:
base = {
"id": "chatcmpl-cf",
"object": "chat.completion.chunk",
"created": 1234567890,
"model": "@cf/meta/llama-2-7b-chat-int8",
}
return [ return [
json.dumps({"response": "I am"}), json.dumps({**base, "choices": [{"index": 0, "delta": {"content": "I am"}}]}),
json.dumps({"response": " a language"}), json.dumps(
json.dumps({"response": " model."}), {**base, "choices": [{"index": 0, "delta": {"content": " a language"}}]}
] ),
json.dumps(
{
def _streaming_chunks_response_text() -> list[str]: **base,
return [ "choices": [
json.dumps({"response_text": "I am"}), {
json.dumps({"response_text": " a language"}), "index": 0,
json.dumps({"response_text": " model."}), "delta": {"content": " model."},
"finish_reason": "stop",
}
],
}
),
] ]
@ -85,6 +133,48 @@ def test_completion_cloudflare(sync_mode):
assert response.choices[0].message.content is not None assert response.choices[0].message.content is not None
assert "language model" in response.choices[0].message.content.lower() assert "language model" in response.choices[0].message.content.lower()
called_url = mock_post.call_args.kwargs.get("url") or mock_post.call_args.args[0]
assert called_url.endswith("/ai/v1/chat/completions")
assert "/ai/run/" not in called_url
def test_completion_cloudflare_tool_calls_sent_to_openai_endpoint():
messages = [{"role": "user", "content": "weather in New York?"}]
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
},
}
]
mock_resp = _make_mock_response(_tool_call_response())
with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post:
response = completion(
model="cloudflare/@cf/meta/llama-2-7b-chat-int8",
messages=messages,
tools=tools,
tool_choice="auto",
api_base=FAKE_API_BASE,
api_key=FAKE_API_KEY,
)
mock_post.assert_called_once()
sent_body = json.loads(mock_post.call_args.kwargs["data"])
assert sent_body["tools"] == tools
assert sent_body["tool_choice"] == "auto"
assert response.choices[0].finish_reason == "tool_calls"
tool_calls = response.choices[0].message.tool_calls
assert tool_calls is not None and len(tool_calls) == 1
assert tool_calls[0].function.name == "get_weather"
@pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.parametrize("sync_mode", [True, False])
def test_completion_cloudflare_stream(sync_mode): def test_completion_cloudflare_stream(sync_mode):
@ -153,76 +243,3 @@ def test_completion_cloudflare_stream(sync_mode):
if c.choices[0].delta.content if c.choices[0].delta.content
) )
assert "language" in content.lower() assert "language" in content.lower()
@pytest.mark.parametrize("sync_mode", [True, False])
def test_completion_cloudflare_stream_response_text(sync_mode):
"""Newer Cloudflare Workers AI models (e.g. Nemotron) emit `response_text`
instead of `response` in streamed chunks. The iterator must surface that
text so streaming output is not silently empty.
"""
messages = [{"role": "user", "content": "what llm are you"}]
raw_chunks = _streaming_chunks_response_text()
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/nvidia/nemotron-mini-4b-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/nvidia/nemotron-mini-4b-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()

View file

@ -3,25 +3,126 @@ import pytest
from litellm.llms.cloudflare.chat.transformation import CloudflareChatConfig from litellm.llms.cloudflare.chat.transformation import CloudflareChatConfig
def test_get_complete_url_encodes_model_path_segment(): def test_supported_params_include_tools_and_tool_choice():
config = CloudflareChatConfig() config = CloudflareChatConfig()
assert ( params = config.get_supported_openai_params(model="@cf/meta/llama-2-7b-chat-int8")
config.get_complete_url(
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/", assert "tools" in params
api_key="cf-key", assert "tool_choice" in params
model="@cf/meta/llama?x=1#frag", assert "stream" in params
optional_params={}, assert "max_tokens" in params
litellm_params={},
)
== "https://api.cloudflare.com/client/v4/accounts/acct/ai/run/%40cf/meta/llama%3Fx%3D1%23frag" def test_get_complete_url_defaults_to_openai_compatible_endpoint(monkeypatch):
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
config = CloudflareChatConfig()
url = config.get_complete_url(
api_base=None,
api_key="cf-key",
model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
) )
with pytest.raises(ValueError, match="dot path segment"): assert (
config.get_complete_url( url
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/", == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
api_key="cf-key", )
model="../../accounts/other", assert "/ai/run/" not in url
def test_get_complete_url_appends_chat_completions_to_explicit_base():
config = CloudflareChatConfig()
url = config.get_complete_url(
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/v1",
api_key="cf-key",
model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
)
assert (
url
== "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
)
assert "/ai/run/" not in url
def test_get_complete_url_is_idempotent_for_full_base():
config = CloudflareChatConfig()
url = config.get_complete_url(
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions",
api_key="cf-key",
model="@cf/meta/llama-2-7b-chat-int8",
optional_params={},
litellm_params={},
)
assert (
url
== "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/chat/completions"
)
def test_transform_request_passes_tools_through_in_openai_format():
config = CloudflareChatConfig()
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
},
},
}
]
messages = [{"role": "user", "content": "weather in nyc?"}]
body = config.transform_request(
model="@cf/meta/llama-2-7b-chat-int8",
messages=messages,
optional_params={"tools": tools, "tool_choice": "auto"},
litellm_params={},
headers={},
)
assert body["messages"] == messages
assert body["model"] == "@cf/meta/llama-2-7b-chat-int8"
assert body["tools"] == tools
assert body["tool_choice"] == "auto"
def test_validate_environment_requires_api_key():
config = CloudflareChatConfig()
with pytest.raises(ValueError, match="Missing Cloudflare API Key"):
config.validate_environment(
headers={},
model="@cf/meta/llama-2-7b-chat-int8",
messages=[],
optional_params={}, optional_params={},
litellm_params={}, litellm_params={},
api_key=None,
) )
def test_validate_environment_sets_bearer_and_content_type():
config = CloudflareChatConfig()
headers = config.validate_environment(
headers={},
model="@cf/meta/llama-2-7b-chat-int8",
messages=[],
optional_params={},
litellm_params={},
api_key="cf-key",
)
assert headers["Authorization"] == "Bearer cf-key"
assert headers["Content-Type"] == "application/json"