Files
2026-07-13 13:22:34 +08:00

172 lines
5.8 KiB
Python

from unittest import mock
import pytest
from aiohttp import ClientTimeout
from fastapi.encoders import jsonable_encoder
from mlflow.environment_variables import MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS
from mlflow.gateway.config import EndpointConfig
from mlflow.gateway.exceptions import AIGatewayException
from mlflow.gateway.providers.huggingface import HFTextGenerationInferenceServerProvider
from mlflow.gateway.schemas import chat, completions, embeddings
from tests.gateway.tools import MockAsyncResponse
from tests.helper_functions import skip_if_hf_hub_unhealthy
pytestmark = skip_if_hf_hub_unhealthy()
def completions_config():
return {
"name": "completions",
"endpoint_type": "llm/v1/completions",
"model": {
"provider": "huggingface-text-generation-inference",
"name": "hf-tgi",
"config": {"hf_server_url": "https://testserverurl.com"},
},
}
def embedding_config():
return {
"name": "embeddings",
"endpoint_type": "llm/v1/embeddings",
"model": {
"provider": "huggingface-text-generation-inference",
"name": "hf-tgi",
"config": {"hf_server_url": "https://testserverurl.com"},
},
}
def chat_config():
return {
"name": "chat",
"endpoint_type": "llm/v1/chat",
"model": {
"provider": "huggingface-text-generation-inference",
"name": "hf-tgi",
"config": {"hf_server_url": "https://testserverurl.com"},
},
}
def completions_response():
return {
"generated_text": "this is a test response",
"details": {
"finish_reason": "length",
"generated_tokens": 5,
"seed": 0,
"prefill": [{"text": "This"}, {"text": "is"}, {"text": "a"}, {"text": "test"}],
},
}
def test_get_provider_name():
config = completions_config()
provider = HFTextGenerationInferenceServerProvider(EndpointConfig(**config))
assert provider.DISPLAY_NAME == "Hugging Face Text Generation Inference"
assert provider.get_provider_name() == "huggingface"
@pytest.mark.asyncio
async def test_completions():
resp = completions_response()
config = completions_config()
with (
mock.patch("time.time", return_value=1677858242),
mock.patch("aiohttp.ClientSession.post", return_value=MockAsyncResponse(resp)) as mock_post,
):
provider = HFTextGenerationInferenceServerProvider(EndpointConfig(**config))
payload = {
"prompt": "This is a test",
"n": 1,
"max_tokens": 1000,
}
response = await provider.completions(completions.RequestPayload(**payload))
assert jsonable_encoder(response) == {
"id": None,
"object": "text_completion",
"created": 1677858242,
"model": "hf-tgi",
"choices": [
{
"text": "this is a test response",
"index": 0,
"finish_reason": "length",
}
],
"usage": {"prompt_tokens": 4, "completion_tokens": 5, "total_tokens": 9},
}
mock_post.assert_called_once_with(
"https://testserverurl.com/generate",
json={
"inputs": "This is a test",
"parameters": {
"max_new_tokens": 1000,
"details": True,
"decoder_input_details": True,
},
},
timeout=ClientTimeout(total=MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS.get()),
)
@pytest.mark.asyncio
async def test_completions_temperature_is_scaled_correctly():
resp = completions_response()
config = completions_config()
with mock.patch(
"aiohttp.ClientSession.post", return_value=MockAsyncResponse(resp)
) as mock_post:
provider = HFTextGenerationInferenceServerProvider(EndpointConfig(**config))
payload = {
"prompt": "This is a test",
"temperature": 0.5,
}
await provider.completions(completions.RequestPayload(**payload))
assert mock_post.call_args[1]["json"]["parameters"]["temperature"] == 0.5 * 50
@pytest.mark.asyncio
async def test_completion_fails_with_multiple_candidates():
config = chat_config()
provider = HFTextGenerationInferenceServerProvider(EndpointConfig(**config))
payload = {"prompt": "This is a test", "n": 2}
with pytest.raises(AIGatewayException, match=r".*") as e:
await provider.completions(completions.RequestPayload(**payload))
assert "'n' must be '1' for the Text Generation Inference provider." in e.value.detail
assert e.value.status_code == 422
@pytest.mark.asyncio
async def test_chat_is_not_supported_for_tgi():
config = chat_config()
provider = HFTextGenerationInferenceServerProvider(EndpointConfig(**config))
payload = {"messages": [{"role": "user", "content": "TGI, can you chat with me? I'm lonely."}]}
with pytest.raises(AIGatewayException, match=r".*") as e:
await provider.chat(chat.RequestPayload(**payload))
assert (
"The chat route is not implemented for Hugging Face Text Generation Inference models."
in e.value.detail
)
assert e.value.status_code == 501
@pytest.mark.asyncio
async def test_embeddings_are_not_supported_for_tgi():
config = embedding_config()
provider = HFTextGenerationInferenceServerProvider(EndpointConfig(**config))
payload = {"input": "give me that sweet, sweet vector, please."}
with pytest.raises(AIGatewayException, match=r".*") as e:
await provider.embeddings(embeddings.RequestPayload(**payload))
assert (
"The embeddings route is not implemented for Hugging Face Text Generation Inference models."
in e.value.detail
)
assert e.value.status_code == 501