mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
test(cloudflare): fake the HTTP boundary instead of patching litellm internals
This commit is contained in:
parent
30fdb6226b
commit
70c919aece
3 changed files with 111 additions and 79 deletions
|
|
@ -1026,27 +1026,47 @@ def test_hosted_vllm_embedding(monkeypatch):
|
|||
assert json_data["model"] == "jina-embeddings-v3"
|
||||
|
||||
|
||||
class _RecordingHTTPHandler(HTTPHandler):
|
||||
def __init__(self, response):
|
||||
super().__init__()
|
||||
self.response = response
|
||||
self.requests = []
|
||||
|
||||
def post(self, url: str, **kwargs):
|
||||
self.requests.append({"url": url, **kwargs})
|
||||
return self.response
|
||||
|
||||
|
||||
def test_cloudflare_embedding_dispatch(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "cloudflare_api_key", None)
|
||||
monkeypatch.setattr(litellm, "api_key", None)
|
||||
monkeypatch.setattr(litellm, "api_base", None)
|
||||
|
||||
with patch(
|
||||
"litellm.main.base_llm_http_handler.embedding",
|
||||
return_value=litellm.EmbeddingResponse(),
|
||||
) as mock_embedding:
|
||||
embedding(
|
||||
model="cloudflare/@cf/baai/bge-large-en-v1.5",
|
||||
input=["Hello world"],
|
||||
api_key="cf-key",
|
||||
api_base="https://example.com/ai/v1",
|
||||
caching=False,
|
||||
)
|
||||
response_json = {
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "embedding": [0.1, 0.2, 0.3], "index": 0}],
|
||||
"model": "@cf/baai/bge-large-en-v1.5",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
}
|
||||
raw_response = MagicMock()
|
||||
raw_response.status_code = 200
|
||||
raw_response.headers = {"content-type": "application/json"}
|
||||
raw_response.json.return_value = response_json
|
||||
raw_response.text = json.dumps(response_json)
|
||||
client = _RecordingHTTPHandler(raw_response)
|
||||
|
||||
dispatch = mock_embedding.call_args.kwargs
|
||||
assert dispatch["custom_llm_provider"] == "cloudflare"
|
||||
assert dispatch["api_key"] == "cf-key"
|
||||
assert dispatch["api_base"] == "https://example.com/ai/v1"
|
||||
embedding(
|
||||
model="cloudflare/@cf/baai/bge-large-en-v1.5",
|
||||
input=["Hello world"],
|
||||
api_key="cf-key",
|
||||
api_base="https://example.com/ai/v1",
|
||||
client=client,
|
||||
caching=False,
|
||||
)
|
||||
|
||||
request = client.requests[0]
|
||||
assert request["url"] == "https://example.com/ai/v1/embeddings"
|
||||
assert request["headers"]["Authorization"] == "Bearer cf-key"
|
||||
|
||||
|
||||
def test_cloudflare_embedding_dispatch_requires_api_key(monkeypatch):
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import json
|
||||
from unittest.mock import Mock, patch
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -10,6 +10,18 @@ from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
|||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
|
||||
class _RecordingHTTPHandler(HTTPHandler):
|
||||
def __init__(self, response):
|
||||
super().__init__()
|
||||
self.response = response
|
||||
self.requests = []
|
||||
|
||||
def post(self, url: str, **kwargs):
|
||||
self.requests.append({"url": url, **kwargs})
|
||||
return self.response
|
||||
|
||||
|
||||
|
||||
def test_provider_config_manager_returns_cloudflare_embedding_config():
|
||||
config = ProviderConfigManager.get_provider_embedding_config(
|
||||
model="@cf/baai/bge-large-en-v1.5",
|
||||
|
|
@ -90,7 +102,6 @@ def test_get_error_class():
|
|||
|
||||
def test_embedding_routes_to_cloudflare_openai_compatible_endpoint(monkeypatch):
|
||||
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
|
||||
client = HTTPHandler()
|
||||
response_json = {
|
||||
"object": "list",
|
||||
"data": [
|
||||
|
|
@ -108,17 +119,17 @@ def test_embedding_routes_to_cloudflare_openai_compatible_endpoint(monkeypatch):
|
|||
raw_response.headers = {"content-type": "application/json"}
|
||||
raw_response.json.return_value = response_json
|
||||
raw_response.text = json.dumps(response_json)
|
||||
client = _RecordingHTTPHandler(raw_response)
|
||||
|
||||
with patch.object(HTTPHandler, "post", return_value=raw_response) as mock_post:
|
||||
response = litellm.embedding(
|
||||
model="cloudflare/@cf/baai/bge-large-en-v1.5",
|
||||
input=["hello"],
|
||||
api_key="cf-key",
|
||||
client=client,
|
||||
caching=False,
|
||||
)
|
||||
response = litellm.embedding(
|
||||
model="cloudflare/@cf/baai/bge-large-en-v1.5",
|
||||
input=["hello"],
|
||||
api_key="cf-key",
|
||||
client=client,
|
||||
caching=False,
|
||||
)
|
||||
|
||||
request = mock_post.call_args.kwargs
|
||||
request = client.requests[0]
|
||||
body = json.loads(request["data"])
|
||||
assert request["url"] == "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1/embeddings"
|
||||
assert request["headers"]["Authorization"] == "Bearer cf-key"
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import json
|
||||
from unittest.mock import MagicMock, Mock, patch
|
||||
from unittest.mock import MagicMock, Mock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -12,6 +12,18 @@ from litellm.types.rerank import RerankResponse
|
|||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
|
||||
class _RecordingHTTPHandler(HTTPHandler):
|
||||
def __init__(self, response):
|
||||
super().__init__()
|
||||
self.response = response
|
||||
self.requests = []
|
||||
|
||||
def post(self, url: str, **kwargs):
|
||||
self.requests.append({"url": url, **kwargs})
|
||||
return self.response
|
||||
|
||||
|
||||
|
||||
def test_provider_config_manager_returns_cloudflare_rerank_config():
|
||||
config = ProviderConfigManager.get_provider_rerank_config(
|
||||
model="@cf/baai/bge-reranker-base",
|
||||
|
|
@ -23,17 +35,14 @@ def test_provider_config_manager_returns_cloudflare_rerank_config():
|
|||
assert isinstance(config, CloudflareRerankConfig)
|
||||
|
||||
|
||||
def test_get_complete_url_uses_native_workers_ai_endpoint():
|
||||
def test_get_complete_url_uses_native_workers_ai_endpoint(monkeypatch):
|
||||
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "account-id")
|
||||
config = CloudflareRerankConfig()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.cloudflare.rerank.transformation.get_secret_str",
|
||||
return_value="account-id",
|
||||
):
|
||||
url = config.get_complete_url(
|
||||
api_base=None,
|
||||
model="@cf/baai/bge-reranker-base",
|
||||
)
|
||||
url = config.get_complete_url(
|
||||
api_base=None,
|
||||
model="@cf/baai/bge-reranker-base",
|
||||
)
|
||||
|
||||
assert url == ("https://api.cloudflare.com/client/v4/accounts/account-id/ai/run/%40cf/baai/bge-reranker-base")
|
||||
|
||||
|
|
@ -72,31 +81,26 @@ def test_get_complete_url_handles_supported_base_shapes(api_base, expected):
|
|||
assert config.get_complete_url(api_base, "@cf/baai/bge-reranker-base") == expected
|
||||
|
||||
|
||||
def test_get_complete_url_requires_account_id():
|
||||
def test_get_complete_url_requires_account_id(monkeypatch):
|
||||
monkeypatch.delenv("CLOUDFLARE_ACCOUNT_ID", raising=False)
|
||||
config = CloudflareRerankConfig()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.llms.cloudflare.rerank.transformation.get_secret_str",
|
||||
return_value=None,
|
||||
),
|
||||
pytest.raises(ValueError, match="CLOUDFLARE_ACCOUNT_ID"),
|
||||
):
|
||||
with pytest.raises(ValueError, match="CLOUDFLARE_ACCOUNT_ID"):
|
||||
config.get_complete_url(None, "@cf/baai/bge-reranker-base")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
"model,error_match",
|
||||
(
|
||||
"../graphql",
|
||||
"@cf/baai/../graphql",
|
||||
"/@cf/baai/bge-reranker-base",
|
||||
("../graphql", "cannot be a dot path segment"),
|
||||
("@cf/baai/../graphql", "cannot be a dot path segment"),
|
||||
("/@cf/baai/bge-reranker-base", "model is required"),
|
||||
),
|
||||
)
|
||||
def test_get_complete_url_rejects_path_traversal(model):
|
||||
def test_get_complete_url_rejects_path_traversal(model, error_match):
|
||||
config = CloudflareRerankConfig()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
with pytest.raises(ValueError, match=error_match):
|
||||
config.get_complete_url(
|
||||
"https://api.cloudflare.com/client/v4/accounts/account-id/ai/run",
|
||||
model,
|
||||
|
|
@ -133,16 +137,11 @@ def test_validate_environment_and_supported_params():
|
|||
)
|
||||
|
||||
|
||||
def test_validate_environment_requires_api_key():
|
||||
def test_validate_environment_requires_api_key(monkeypatch):
|
||||
monkeypatch.delenv("CLOUDFLARE_API_KEY", raising=False)
|
||||
config = CloudflareRerankConfig()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.llms.cloudflare.rerank.transformation.get_secret_str",
|
||||
return_value=None,
|
||||
),
|
||||
pytest.raises(ValueError, match="Cloudflare API Key"),
|
||||
):
|
||||
with pytest.raises(ValueError, match="Cloudflare API Key"):
|
||||
config.validate_environment({}, "@cf/baai/bge-reranker-base")
|
||||
|
||||
|
||||
|
|
@ -214,18 +213,21 @@ def test_transform_rerank_request():
|
|||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"params",
|
||||
[
|
||||
{"documents": ("document",)},
|
||||
{"query": "query"},
|
||||
{"query": "query", "documents": "document"},
|
||||
{"query": "query", "documents": ()},
|
||||
],
|
||||
"params,error_match",
|
||||
(
|
||||
({"documents": ("document",)}, "query is required for Cloudflare rerank"),
|
||||
({"query": "query"}, "documents is required for Cloudflare rerank"),
|
||||
(
|
||||
{"query": "query", "documents": "document"},
|
||||
"documents is required for Cloudflare rerank",
|
||||
),
|
||||
({"query": "query", "documents": ()}, "documents is required for Cloudflare rerank"),
|
||||
),
|
||||
)
|
||||
def test_transform_rerank_request_validates_required_params(params):
|
||||
def test_transform_rerank_request_validates_required_params(params, error_match):
|
||||
config = CloudflareRerankConfig()
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
with pytest.raises(ValueError, match=error_match):
|
||||
config.transform_rerank_request(
|
||||
model="@cf/baai/bge-reranker-base",
|
||||
optional_rerank_params=params,
|
||||
|
|
@ -372,7 +374,6 @@ def test_get_error_class():
|
|||
|
||||
|
||||
def test_litellm_rerank_sends_cloudflare_request():
|
||||
client = HTTPHandler()
|
||||
response_json = {
|
||||
"result": {"response": [{"id": 0, "score": 0.98}]},
|
||||
"success": True,
|
||||
|
|
@ -381,19 +382,19 @@ def test_litellm_rerank_sends_cloudflare_request():
|
|||
raw_response.status_code = 200
|
||||
raw_response.json.return_value = response_json
|
||||
raw_response.text = json.dumps(response_json)
|
||||
client = _RecordingHTTPHandler(raw_response)
|
||||
|
||||
with patch.object(HTTPHandler, "post", return_value=raw_response) as mock_post:
|
||||
response = litellm.rerank(
|
||||
model="cloudflare/@cf/baai/bge-reranker-base",
|
||||
query="What is LiteLLM?",
|
||||
documents=["LiteLLM is an LLM gateway.", "A recipe for soup."],
|
||||
top_n=1,
|
||||
api_key="test-key",
|
||||
api_base="https://api.cloudflare.com/client/v4/accounts/account-id/ai/run",
|
||||
client=client,
|
||||
)
|
||||
response = litellm.rerank(
|
||||
model="cloudflare/@cf/baai/bge-reranker-base",
|
||||
query="What is LiteLLM?",
|
||||
documents=["LiteLLM is an LLM gateway.", "A recipe for soup."],
|
||||
top_n=1,
|
||||
api_key="test-key",
|
||||
api_base="https://api.cloudflare.com/client/v4/accounts/account-id/ai/run",
|
||||
client=client,
|
||||
)
|
||||
|
||||
request = mock_post.call_args.kwargs
|
||||
request = client.requests[0]
|
||||
assert request["url"].endswith("/ai/run/%40cf/baai/bge-reranker-base")
|
||||
assert request["headers"]["Authorization"] == "Bearer test-key"
|
||||
assert json.loads(request["data"]) == {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue