mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(cloudflare): route native Workers AI provider through OpenAI-compatible endpoint
This commit is contained in:
parent
6f6aec2930
commit
1815cc9bee
4 changed files with 261 additions and 286 deletions
|
|
@ -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}")
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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()
|
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue