mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +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
|
||||
import time
|
||||
from typing import AsyncIterator, Iterator, List, Optional, Union
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segments
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.base_llm.chat.transformation import (
|
||||
BaseConfig,
|
||||
BaseLLMException,
|
||||
LiteLLMLoggingObj,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionUsageBlock,
|
||||
GenericStreamingChunk,
|
||||
ModelResponse,
|
||||
Usage,
|
||||
)
|
||||
|
||||
|
||||
class CloudflareError(BaseLLMException):
|
||||
|
|
@ -34,26 +19,32 @@ class CloudflareError(BaseLLMException):
|
|||
message=message,
|
||||
request=self.request,
|
||||
response=self.response,
|
||||
) # Call the base class constructor with the parameters it needs
|
||||
)
|
||||
|
||||
|
||||
class CloudflareChatConfig(BaseConfig):
|
||||
max_tokens: Optional[int] = None
|
||||
stream: Optional[bool] = None
|
||||
|
||||
def __init__(
|
||||
class CloudflareChatConfig(OpenAIGPTConfig):
|
||||
def get_complete_url(
|
||||
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,
|
||||
) -> None:
|
||||
locals_ = locals().copy()
|
||||
for key, value in locals_.items():
|
||||
if key != "self" and value is not None:
|
||||
setattr(self.__class__, key, value)
|
||||
|
||||
@classmethod
|
||||
def get_config(cls):
|
||||
return super().get_config()
|
||||
) -> 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/v1"
|
||||
)
|
||||
return super().get_complete_url(
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
stream=stream,
|
||||
)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
|
|
@ -67,107 +58,18 @@ class CloudflareChatConfig(BaseConfig):
|
|||
) -> dict:
|
||||
if api_key is None:
|
||||
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 = {
|
||||
"accept": "application/json",
|
||||
"content-type": "apbplication/json",
|
||||
"Authorization": "Bearer " + api_key,
|
||||
}
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
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", ""))
|
||||
return super().validate_environment(
|
||||
headers=headers,
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
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(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
|
|
@ -175,48 +77,3 @@ class CloudflareChatConfig(BaseConfig):
|
|||
status_code=status_code,
|
||||
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
|
||||
or litellm.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
|
||||
|
|
|
|||
|
|
@ -9,9 +9,7 @@ import pytest
|
|||
from litellm import acompletion, completion
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
|
||||
FAKE_API_BASE = (
|
||||
"https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/run/"
|
||||
)
|
||||
FAKE_API_BASE = "https://fake-cloudflare.example.com/client/v4/accounts/fake-acct/ai/v1"
|
||||
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]:
|
||||
return {
|
||||
"result": {
|
||||
"response": "I am a large language model created to assist you.",
|
||||
},
|
||||
"success": True,
|
||||
"errors": [],
|
||||
"messages": [],
|
||||
"id": "chatcmpl-cf",
|
||||
"object": "chat.completion",
|
||||
"created": 1234567890,
|
||||
"model": "@cf/meta/llama-2-7b-chat-int8",
|
||||
"choices": [
|
||||
{
|
||||
"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]:
|
||||
base = {
|
||||
"id": "chatcmpl-cf",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1234567890,
|
||||
"model": "@cf/meta/llama-2-7b-chat-int8",
|
||||
}
|
||||
return [
|
||||
json.dumps({"response": "I am"}),
|
||||
json.dumps({"response": " a language"}),
|
||||
json.dumps({"response": " model."}),
|
||||
]
|
||||
|
||||
|
||||
def _streaming_chunks_response_text() -> list[str]:
|
||||
return [
|
||||
json.dumps({"response_text": "I am"}),
|
||||
json.dumps({"response_text": " a language"}),
|
||||
json.dumps({"response_text": " model."}),
|
||||
json.dumps({**base, "choices": [{"index": 0, "delta": {"content": "I am"}}]}),
|
||||
json.dumps(
|
||||
{**base, "choices": [{"index": 0, "delta": {"content": " a language"}}]}
|
||||
),
|
||||
json.dumps(
|
||||
{
|
||||
**base,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"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 "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])
|
||||
def test_completion_cloudflare_stream(sync_mode):
|
||||
|
|
@ -153,76 +243,3 @@ def test_completion_cloudflare_stream(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_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
|
||||
|
||||
|
||||
def test_get_complete_url_encodes_model_path_segment():
|
||||
def test_supported_params_include_tools_and_tool_choice():
|
||||
config = CloudflareChatConfig()
|
||||
|
||||
assert (
|
||||
config.get_complete_url(
|
||||
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/",
|
||||
api_key="cf-key",
|
||||
model="@cf/meta/llama?x=1#frag",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
== "https://api.cloudflare.com/client/v4/accounts/acct/ai/run/%40cf/meta/llama%3Fx%3D1%23frag"
|
||||
params = config.get_supported_openai_params(model="@cf/meta/llama-2-7b-chat-int8")
|
||||
|
||||
assert "tools" in params
|
||||
assert "tool_choice" in params
|
||||
assert "stream" in params
|
||||
assert "max_tokens" in params
|
||||
|
||||
|
||||
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"):
|
||||
config.get_complete_url(
|
||||
api_base="https://api.cloudflare.com/client/v4/accounts/acct/ai/run/",
|
||||
api_key="cf-key",
|
||||
model="../../accounts/other",
|
||||
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_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={},
|
||||
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