mlflow--mlflow
1498 行
52 KiB
Python
1498 行
52 KiB
Python
from unittest import mock
|
|
|
|
import pytest
|
|
from aiohttp import ClientTimeout
|
|
from fastapi.encoders import jsonable_encoder
|
|
from pydantic import ValidationError
|
|
|
|
from mlflow.environment_variables import MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS
|
|
from mlflow.exceptions import MlflowException
|
|
from mlflow.gateway.config import EndpointConfig, OpenAIConfig
|
|
from mlflow.gateway.exceptions import AIGatewayException
|
|
from mlflow.gateway.providers.base import PassthroughAction
|
|
from mlflow.gateway.providers.openai import OpenAIAdapter, OpenAIProvider
|
|
from mlflow.gateway.schemas import chat, completions, embeddings
|
|
|
|
from tests.gateway.tools import (
|
|
MockAsyncResponse,
|
|
MockAsyncStreamingResponse,
|
|
mock_http_client,
|
|
)
|
|
|
|
|
|
def chat_config():
|
|
return {
|
|
"name": "chat",
|
|
"endpoint_type": "llm/v1/chat",
|
|
"model": {
|
|
"provider": "openai",
|
|
"name": "gpt-4o-mini",
|
|
"config": {
|
|
"openai_api_base": "https://api.openai.com/v1",
|
|
"openai_api_key": "key",
|
|
},
|
|
},
|
|
}
|
|
|
|
|
|
def chat_response():
|
|
return {
|
|
"id": "chatcmpl-abc123",
|
|
"object": "chat.completion",
|
|
"created": 1677858242,
|
|
"model": "gpt-4o-mini",
|
|
"usage": {
|
|
"prompt_tokens": 13,
|
|
"completion_tokens": 7,
|
|
"total_tokens": 20,
|
|
},
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": "\n\nThis is a test!",
|
|
},
|
|
"finish_reason": "stop",
|
|
"index": 0,
|
|
}
|
|
],
|
|
"headers": {"Content-Type": "application/json"},
|
|
}
|
|
|
|
|
|
def completions_response():
|
|
return {
|
|
"id": "chatcmpl-abc123",
|
|
"object": "text.completion",
|
|
"created": 1677858242,
|
|
"model": "gpt-4o-mini",
|
|
"usage": {
|
|
"prompt_tokens": 13,
|
|
"completion_tokens": 7,
|
|
"total_tokens": 20,
|
|
},
|
|
"choices": [
|
|
{
|
|
"text": "\n\nThis is a test!",
|
|
"index": 0,
|
|
"finish_reason": "stop",
|
|
}
|
|
],
|
|
"headers": {"Content-Type": "application/json"},
|
|
}
|
|
|
|
|
|
async def _run_test_chat(provider):
|
|
resp = chat_response()
|
|
mock_client = mock_http_client(MockAsyncResponse(resp))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client) as mock_build_client:
|
|
payload = {
|
|
"messages": [{"role": "user", "content": "Tell me a joke"}],
|
|
"temperature": 0.5,
|
|
"top_p": 0.9,
|
|
"presence_penalty": 0.1,
|
|
"frequency_penalty": 0.2,
|
|
}
|
|
response = await provider.chat(chat.RequestPayload(**payload))
|
|
assert jsonable_encoder(response) == {
|
|
"id": "chatcmpl-abc123",
|
|
"object": "chat.completion",
|
|
"created": 1677858242,
|
|
"model": "gpt-4o-mini",
|
|
"provider": "openai",
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": "\n\nThis is a test!",
|
|
"tool_calls": None,
|
|
"refusal": None,
|
|
},
|
|
"finish_reason": "stop",
|
|
"index": 0,
|
|
}
|
|
],
|
|
"usage": {
|
|
"prompt_tokens": 13,
|
|
"completion_tokens": 7,
|
|
"total_tokens": 20,
|
|
},
|
|
}
|
|
mock_build_client.assert_called_once()
|
|
call_headers = mock_build_client.call_args.kwargs["headers"]
|
|
assert call_headers.get("authorization") == "Bearer key"
|
|
mock_client.post.assert_called_once_with(
|
|
"https://api.openai.com/v1/chat/completions",
|
|
json={
|
|
"model": "gpt-4o-mini",
|
|
"n": 1,
|
|
**payload,
|
|
},
|
|
timeout=ClientTimeout(total=MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS.get()),
|
|
)
|
|
|
|
|
|
def test_get_headers_uses_server_key_by_default():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
merged = provider._get_headers(
|
|
headers={"authorization": "Bearer client-key", "X-Custom": "value"}
|
|
)
|
|
assert merged["authorization"] == "Bearer key"
|
|
assert merged["X-Custom"] == "value"
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"user_agent",
|
|
[
|
|
"claude-cli/2.0.37 (external, cli)",
|
|
"Codex-Desktop/26.422.2437.0",
|
|
"GeminiCLI/0.39.0/gemini-2.0-pro (darwin; x64)",
|
|
],
|
|
)
|
|
def test_get_headers_preserves_client_key_for_credential_agents(user_agent):
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
merged = provider._get_headers(
|
|
headers={"authorization": "Bearer client-key", "user-agent": user_agent}
|
|
)
|
|
assert merged["authorization"] == "Bearer client-key"
|
|
|
|
|
|
def test_get_headers_preserves_azure_api_key_for_credential_agents():
|
|
provider = OpenAIProvider(EndpointConfig(**azure_config(api_type="azure")))
|
|
merged = provider._get_headers(
|
|
headers={"api-key": "client-azure-key", "user-agent": "claude-cli/2.0.37 (external, cli)"}
|
|
)
|
|
assert merged["api-key"] == "client-azure-key"
|
|
assert "authorization" not in merged
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_chat():
|
|
config = chat_config()
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
await _run_test_chat(provider)
|
|
|
|
|
|
def chat_stream_response():
|
|
return [
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":null,"delta":{"role":"assistant"}}]}\n',
|
|
b"\n",
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":null,"delta":{"content":"test"}}]}\n',
|
|
b"\n",
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":"stop","delta":{}}]}\n',
|
|
b"\n",
|
|
b"data: [DONE]\n",
|
|
]
|
|
|
|
|
|
def chat_stream_response_incomplete():
|
|
return [
|
|
# contains first half of a chunk
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test","choi'
|
|
# contains second half of first chunk and first half of second chunk
|
|
b'ces":[{"index":0,"finish_reason":null,"delta":{"role":"assistant"}}]}\n\n'
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"te',
|
|
# contains second half of second chunk
|
|
b'st","choices":[{"index":0,"finish_reason":null,"delta":{"content":"test"}}]}\n',
|
|
b"\n",
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":"stop","delta":{}}]}\n',
|
|
b"\n",
|
|
b"data: [DONE]\n",
|
|
]
|
|
|
|
|
|
async def _run_test_chat_stream(resp, provider):
|
|
mock_client = mock_http_client(MockAsyncStreamingResponse(resp))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client) as mock_build_client:
|
|
payload = {"messages": [{"role": "user", "content": "Tell me a joke"}]}
|
|
response = provider.chat_stream(chat.RequestPayload(**payload))
|
|
|
|
chunks = [jsonable_encoder(chunk) async for chunk in response]
|
|
assert chunks == [
|
|
{
|
|
"choices": [
|
|
{
|
|
"delta": {
|
|
"content": None,
|
|
"role": "assistant",
|
|
"tool_calls": None,
|
|
},
|
|
"finish_reason": None,
|
|
"index": 0,
|
|
}
|
|
],
|
|
"created": 1,
|
|
"id": "test-id",
|
|
"model": "test",
|
|
"provider": "openai",
|
|
"object": "chat.completion.chunk",
|
|
"usage": None,
|
|
},
|
|
{
|
|
"choices": [
|
|
{
|
|
"delta": {
|
|
"content": "test",
|
|
"role": None,
|
|
"tool_calls": None,
|
|
},
|
|
"finish_reason": None,
|
|
"index": 0,
|
|
}
|
|
],
|
|
"created": 1,
|
|
"id": "test-id",
|
|
"model": "test",
|
|
"provider": "openai",
|
|
"object": "chat.completion.chunk",
|
|
"usage": None,
|
|
},
|
|
{
|
|
"choices": [
|
|
{
|
|
"delta": {
|
|
"content": None,
|
|
"role": None,
|
|
"tool_calls": None,
|
|
},
|
|
"finish_reason": "stop",
|
|
"index": 0,
|
|
}
|
|
],
|
|
"created": 1,
|
|
"id": "test-id",
|
|
"model": "test",
|
|
"provider": "openai",
|
|
"object": "chat.completion.chunk",
|
|
"usage": None,
|
|
},
|
|
]
|
|
|
|
mock_build_client.assert_called_once()
|
|
call_headers = mock_build_client.call_args.kwargs["headers"]
|
|
assert call_headers.get("authorization") == "Bearer key"
|
|
mock_client.post.assert_called_once_with(
|
|
"https://api.openai.com/v1/chat/completions",
|
|
json={
|
|
"model": "gpt-4o-mini",
|
|
"n": 1,
|
|
"stream_options": {"include_usage": True},
|
|
**payload,
|
|
},
|
|
timeout=ClientTimeout(total=MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS.get()),
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize("resp", [chat_stream_response(), chat_stream_response_incomplete()])
|
|
@pytest.mark.asyncio
|
|
async def test_chat_stream(resp):
|
|
config = chat_config()
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
await _run_test_chat_stream(resp, provider)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_chat_stream_with_function_calling():
|
|
config = chat_config()
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
|
|
resp = [
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":null,"delta":{"role":"assistant",'
|
|
b'"tool_calls":[{"index":0,"id":"call_001","function":{"name":"get_weather"},'
|
|
b'"type":"function"}]}}]}\n',
|
|
b"\n",
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":null,"delta":{'
|
|
b'"tool_calls":[{"index":0,"function":{"arguments":"{\\"location\\":"'
|
|
b"}}]}}]}\n",
|
|
b"\n",
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":"stop","delta":{'
|
|
b'"tool_calls":[{"index":0,"function":{"arguments":"\\"Singapore\\"}"'
|
|
b"}}]}}]}\n",
|
|
b"\n",
|
|
b"data: [DONE]\n",
|
|
]
|
|
mock_client = mock_http_client(MockAsyncStreamingResponse(resp))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client):
|
|
payload = {"messages": [{"role": "user", "content": "What's the weather in Singapore?"}]}
|
|
response = provider.chat_stream(chat.RequestPayload(**payload))
|
|
|
|
chunks = [jsonable_encoder(chunk) async for chunk in response]
|
|
assert chunks == [
|
|
{
|
|
"id": "test-id",
|
|
"object": "chat.completion.chunk",
|
|
"created": 1,
|
|
"model": "test",
|
|
"provider": "openai",
|
|
"choices": [
|
|
{
|
|
"index": 0,
|
|
"finish_reason": None,
|
|
"delta": {
|
|
"role": "assistant",
|
|
"content": None,
|
|
"tool_calls": [
|
|
{
|
|
"index": 0,
|
|
"id": "call_001",
|
|
"type": "function",
|
|
"function": {"name": "get_weather", "arguments": None},
|
|
}
|
|
],
|
|
},
|
|
}
|
|
],
|
|
"usage": None,
|
|
},
|
|
{
|
|
"id": "test-id",
|
|
"object": "chat.completion.chunk",
|
|
"created": 1,
|
|
"model": "test",
|
|
"provider": "openai",
|
|
"choices": [
|
|
{
|
|
"index": 0,
|
|
"finish_reason": None,
|
|
"delta": {
|
|
"role": None,
|
|
"content": None,
|
|
"tool_calls": [
|
|
{
|
|
"index": 0,
|
|
"id": None,
|
|
"type": None,
|
|
"function": {"name": None, "arguments": '{"location":'},
|
|
}
|
|
],
|
|
},
|
|
}
|
|
],
|
|
"usage": None,
|
|
},
|
|
{
|
|
"id": "test-id",
|
|
"object": "chat.completion.chunk",
|
|
"created": 1,
|
|
"model": "test",
|
|
"provider": "openai",
|
|
"choices": [
|
|
{
|
|
"index": 0,
|
|
"finish_reason": "stop",
|
|
"delta": {
|
|
"role": None,
|
|
"content": None,
|
|
"tool_calls": [
|
|
{
|
|
"index": 0,
|
|
"id": None,
|
|
"type": None,
|
|
"function": {"name": None, "arguments": '"Singapore"}'},
|
|
}
|
|
],
|
|
},
|
|
}
|
|
],
|
|
"usage": None,
|
|
},
|
|
]
|
|
|
|
|
|
def completions_config():
|
|
return {
|
|
"name": "completions",
|
|
"endpoint_type": "llm/v1/completions",
|
|
"model": {
|
|
"provider": "openai",
|
|
"name": "gpt-4-32k",
|
|
"config": {
|
|
"openai_api_key": "key",
|
|
"openai_organization": "test-organization",
|
|
},
|
|
},
|
|
}
|
|
|
|
|
|
async def _run_test_completions(resp, provider):
|
|
mock_client = mock_http_client(MockAsyncResponse(resp))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client) as mock_build_client:
|
|
payload = {
|
|
"prompt": "This is a test",
|
|
}
|
|
response = await provider.completions(completions.RequestPayload(**payload))
|
|
assert jsonable_encoder(response) == {
|
|
"id": "chatcmpl-abc123",
|
|
"object": "text_completion",
|
|
"created": 1677858242,
|
|
"model": "gpt-4o-mini",
|
|
"choices": [{"text": "\n\nThis is a test!", "index": 0, "finish_reason": "stop"}],
|
|
"usage": {"prompt_tokens": 13, "completion_tokens": 7, "total_tokens": 20},
|
|
}
|
|
mock_build_client.assert_called_once()
|
|
call_headers = mock_build_client.call_args.kwargs["headers"]
|
|
assert call_headers.get("authorization") == "Bearer key"
|
|
assert call_headers.get("OpenAI-Organization") == "test-organization"
|
|
mock_client.post.assert_called_once_with(
|
|
"https://api.openai.com/v1/completions",
|
|
json={
|
|
"model": "gpt-4-32k",
|
|
"n": 1,
|
|
"prompt": "This is a test",
|
|
},
|
|
timeout=ClientTimeout(total=MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS.get()),
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize("resp", [completions_response(), chat_response()])
|
|
@pytest.mark.asyncio
|
|
async def test_completions(resp):
|
|
config = completions_config()
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
await _run_test_completions(resp, provider)
|
|
|
|
|
|
@pytest.mark.parametrize("prompt", [{"set1", "set2"}, ["list1"], [1], ["list1", "list2"], [1, 2]])
|
|
@pytest.mark.asyncio
|
|
async def test_completions_throws_if_prompt_contains_non_string(prompt):
|
|
config = completions_config()
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
payload = {"prompt": prompt}
|
|
with pytest.raises(ValidationError, match=r"prompt"):
|
|
await provider.completions(completions.RequestPayload(**payload))
|
|
|
|
|
|
def completions_stream_response():
|
|
return [
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":null,'
|
|
b'"delta":{"role":"assistant", "content": ""}}]}\n',
|
|
b"\n",
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":null,'
|
|
b'"delta":{"content":"test"}}]}\n',
|
|
b"\n",
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":"length","delta":{}}]}\n',
|
|
b"\n",
|
|
b"data: [DONE]\n",
|
|
]
|
|
|
|
|
|
def completions_stream_response_incomplete():
|
|
return [
|
|
# contains first half of a chunk
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test","choi',
|
|
# contains second half of first chunk and first half of second chunk
|
|
b'ces":[{"index":0,"finish_reason":null,"delta":{"role":"assistant", '
|
|
b'"content": ""}}]}\n\ndata: {"id":"test-id","object":"chat.comp',
|
|
# contains second half of second chunk
|
|
b'letion.chunk","created":1,"model":"test","choices":[{"index":0,"finish_reason":null,'
|
|
b'"delta":{"content":"test"}}]}\n',
|
|
b"\n",
|
|
b'data: {"id":"test-id","object":"chat.completion.chunk","created":1,"model":"test",'
|
|
b'"choices":[{"index":0,"finish_reason":"length","delta":{}}]}\n',
|
|
b"\n",
|
|
b"data: [DONE]\n",
|
|
]
|
|
|
|
|
|
async def _run_test_completions_stream(resp, provider):
|
|
mock_client = mock_http_client(MockAsyncStreamingResponse(resp))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client) as mock_build_client:
|
|
payload = {"prompt": "This is a test"}
|
|
response = provider.completions_stream(completions.RequestPayload(**payload))
|
|
|
|
chunks = [jsonable_encoder(chunk) async for chunk in response]
|
|
assert chunks == [
|
|
{
|
|
"choices": [
|
|
{
|
|
"text": "",
|
|
"finish_reason": None,
|
|
"index": 0,
|
|
}
|
|
],
|
|
"created": 1,
|
|
"id": "test-id",
|
|
"model": "test",
|
|
"object": "text_completion_chunk",
|
|
"usage": None,
|
|
},
|
|
{
|
|
"choices": [
|
|
{
|
|
"text": "test",
|
|
"finish_reason": None,
|
|
"index": 0,
|
|
}
|
|
],
|
|
"created": 1,
|
|
"id": "test-id",
|
|
"model": "test",
|
|
"object": "text_completion_chunk",
|
|
"usage": None,
|
|
},
|
|
{
|
|
"choices": [
|
|
{
|
|
"text": None,
|
|
"finish_reason": "length",
|
|
"index": 0,
|
|
}
|
|
],
|
|
"created": 1,
|
|
"id": "test-id",
|
|
"model": "test",
|
|
"object": "text_completion_chunk",
|
|
"usage": None,
|
|
},
|
|
]
|
|
|
|
mock_build_client.assert_called_once()
|
|
call_headers = mock_build_client.call_args.kwargs["headers"]
|
|
assert call_headers.get("authorization") == "Bearer key"
|
|
assert call_headers.get("OpenAI-Organization") == "test-organization"
|
|
mock_client.post.assert_called_once_with(
|
|
"https://api.openai.com/v1/completions",
|
|
json={
|
|
"model": "gpt-4-32k",
|
|
"n": 1,
|
|
"prompt": "This is a test",
|
|
"stream_options": {"include_usage": True},
|
|
},
|
|
timeout=ClientTimeout(total=MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS.get()),
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"resp", [completions_stream_response(), completions_stream_response_incomplete()]
|
|
)
|
|
@pytest.mark.asyncio
|
|
async def test_completions_stream(resp):
|
|
config = completions_config()
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
await _run_test_completions_stream(resp, provider)
|
|
|
|
|
|
def embedding_config():
|
|
return {
|
|
"name": "embeddings",
|
|
"endpoint_type": "llm/v1/embeddings",
|
|
"model": {
|
|
"provider": "openai",
|
|
"name": "text-embedding-ada-002",
|
|
"config": {
|
|
"openai_api_base": "https://api.openai.com/v1",
|
|
"openai_api_key": "key",
|
|
},
|
|
},
|
|
}
|
|
|
|
|
|
async def _run_test_embeddings(provider):
|
|
resp = {
|
|
"object": "list",
|
|
"data": [
|
|
{
|
|
"object": "embedding",
|
|
"embedding": [
|
|
0.0023064255,
|
|
-0.009327292,
|
|
-0.0028842222,
|
|
],
|
|
"index": 0,
|
|
}
|
|
],
|
|
"model": "text-embedding-ada-002",
|
|
"usage": {"prompt_tokens": 8, "total_tokens": 8},
|
|
"headers": {"Content-Type": "application/json"},
|
|
}
|
|
mock_client = mock_http_client(MockAsyncResponse(resp))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client) as mock_build_client:
|
|
payload = {"input": "This is a test"}
|
|
response = await provider.embeddings(embeddings.RequestPayload(**payload))
|
|
assert jsonable_encoder(response) == {
|
|
"object": "list",
|
|
"data": [
|
|
{
|
|
"object": "embedding",
|
|
"embedding": [
|
|
0.0023064255,
|
|
-0.009327292,
|
|
-0.0028842222,
|
|
],
|
|
"index": 0,
|
|
}
|
|
],
|
|
"model": "text-embedding-ada-002",
|
|
"usage": {"prompt_tokens": 8, "total_tokens": 8},
|
|
}
|
|
mock_build_client.assert_called_once()
|
|
call_headers = mock_build_client.call_args.kwargs["headers"]
|
|
assert call_headers.get("authorization") == "Bearer key"
|
|
mock_client.post.assert_called_once_with(
|
|
"https://api.openai.com/v1/embeddings",
|
|
json={"model": "text-embedding-ada-002", "input": "This is a test"},
|
|
timeout=ClientTimeout(total=MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS.get()),
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_embeddings():
|
|
config = embedding_config()
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
await _run_test_embeddings(provider)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_embeddings_batch_input():
|
|
resp = {
|
|
"object": "list",
|
|
"data": [
|
|
{
|
|
"object": "embedding",
|
|
"embedding": [
|
|
0.1,
|
|
0.2,
|
|
0.3,
|
|
],
|
|
"index": 0,
|
|
},
|
|
{
|
|
"object": "embedding",
|
|
"embedding": [
|
|
0.4,
|
|
0.5,
|
|
0.6,
|
|
],
|
|
"index": 1,
|
|
},
|
|
],
|
|
"model": "text-embedding-ada-002",
|
|
"usage": {"prompt_tokens": 8, "total_tokens": 8},
|
|
"headers": {"Content-Type": "application/json"},
|
|
}
|
|
config = embedding_config()
|
|
mock_client = mock_http_client(MockAsyncResponse(resp))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client) as mock_build_client:
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
payload = {"input": ["1", "2"]}
|
|
response = await provider.embeddings(embeddings.RequestPayload(**payload))
|
|
assert jsonable_encoder(response) == {
|
|
"object": "list",
|
|
"data": [
|
|
{
|
|
"object": "embedding",
|
|
"embedding": [
|
|
0.1,
|
|
0.2,
|
|
0.3,
|
|
],
|
|
"index": 0,
|
|
},
|
|
{
|
|
"object": "embedding",
|
|
"embedding": [
|
|
0.4,
|
|
0.5,
|
|
0.6,
|
|
],
|
|
"index": 1,
|
|
},
|
|
],
|
|
"model": "text-embedding-ada-002",
|
|
"usage": {"prompt_tokens": 8, "total_tokens": 8},
|
|
}
|
|
mock_build_client.assert_called_once()
|
|
call_headers = mock_build_client.call_args.kwargs["headers"]
|
|
assert call_headers.get("authorization") == "Bearer key"
|
|
mock_client.post.assert_called_once_with(
|
|
"https://api.openai.com/v1/embeddings",
|
|
json={
|
|
"model": "text-embedding-ada-002",
|
|
"input": ["1", "2"],
|
|
},
|
|
timeout=ClientTimeout(total=MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS.get()),
|
|
)
|
|
|
|
|
|
def azure_config(api_type: str):
|
|
return {
|
|
"name": "completions",
|
|
"endpoint_type": "llm/v1/completions",
|
|
"model": {
|
|
"provider": "openai",
|
|
"name": "gpt-4o-mini",
|
|
"config": {
|
|
"openai_api_type": api_type,
|
|
"openai_api_key": "key",
|
|
"openai_api_base": "https://test-azureopenai.openai.azure.com/",
|
|
"openai_deployment_name": "test-gpt35",
|
|
"openai_api_version": "2023-05-15",
|
|
},
|
|
},
|
|
}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_azure_openai():
|
|
resp = chat_response()
|
|
config = azure_config(api_type="azure")
|
|
mock_client = mock_http_client(MockAsyncResponse(resp))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client) as mock_build_client:
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
payload = {
|
|
"prompt": "This is a test",
|
|
}
|
|
response = await provider.completions(completions.RequestPayload(**payload))
|
|
assert jsonable_encoder(response) == {
|
|
"id": "chatcmpl-abc123",
|
|
"object": "text_completion",
|
|
"created": 1677858242,
|
|
"model": "gpt-4o-mini",
|
|
"choices": [{"text": "\n\nThis is a test!", "index": 0, "finish_reason": "stop"}],
|
|
"usage": {"prompt_tokens": 13, "completion_tokens": 7, "total_tokens": 20},
|
|
}
|
|
mock_build_client.assert_called_once()
|
|
call_headers = mock_build_client.call_args.kwargs["headers"]
|
|
assert call_headers.get("api-key") == "key"
|
|
mock_client.post.assert_called_once_with(
|
|
(
|
|
"https://test-azureopenai.openai.azure.com/openai/deployments/test-gpt35"
|
|
"/completions?api-version=2023-05-15"
|
|
),
|
|
json={
|
|
"n": 1,
|
|
"prompt": "This is a test",
|
|
},
|
|
timeout=ClientTimeout(total=MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS.get()),
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_azuread_openai():
|
|
resp = chat_response()
|
|
config = azure_config(api_type="azuread")
|
|
mock_client = mock_http_client(MockAsyncResponse(resp))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client) as mock_build_client:
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
payload = {
|
|
"prompt": "This is a test",
|
|
}
|
|
response = await provider.completions(completions.RequestPayload(**payload))
|
|
assert jsonable_encoder(response) == {
|
|
"id": "chatcmpl-abc123",
|
|
"object": "text_completion",
|
|
"created": 1677858242,
|
|
"model": "gpt-4o-mini",
|
|
"choices": [{"text": "\n\nThis is a test!", "index": 0, "finish_reason": "stop"}],
|
|
"usage": {"prompt_tokens": 13, "completion_tokens": 7, "total_tokens": 20},
|
|
}
|
|
mock_build_client.assert_called_once()
|
|
call_headers = mock_build_client.call_args.kwargs["headers"]
|
|
assert call_headers.get("authorization") == "Bearer key"
|
|
mock_client.post.assert_called_once_with(
|
|
(
|
|
"https://test-azureopenai.openai.azure.com/openai/deployments/test-gpt35"
|
|
"/completions?api-version=2023-05-15"
|
|
),
|
|
json={
|
|
"n": 1,
|
|
"prompt": "This is a test",
|
|
},
|
|
timeout=ClientTimeout(total=MLFLOW_GATEWAY_ROUTE_TIMEOUT_SECONDS.get()),
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("api_type", "api_base", "deployment_name", "api_version", "organization"),
|
|
[
|
|
# OpenAI API type
|
|
("openai", None, None, None, None),
|
|
("openai", "https://api.openai.com/v1", None, None, None),
|
|
("openai", "https://api.openai.com/v1", None, "2023-05-15", None),
|
|
("openAI", "https://api.openai.com/v1", None, "2023-05-15", None),
|
|
("openAI", "https://api.openai.com/v1", None, "2023-05-15", "test-organization"),
|
|
# Azure API type
|
|
("azure", "https://test.openai.azure.com", "mock-dep", "2023-05-15", None),
|
|
("AZURe", "https://test.openai.azure.com", "mock-dep", "2023-04-12", None),
|
|
# AzureAD API type
|
|
("azuread", "https://test.openai.azure.com", "mock-dep", "2023-05-15", None),
|
|
("azureAD", "https://test.openai.azure.com", "mock-dep", "2023-04-12", None),
|
|
],
|
|
)
|
|
def test_openai_provider_can_be_constructed_with_valid_configs(
|
|
api_type,
|
|
api_base,
|
|
deployment_name,
|
|
api_version,
|
|
organization,
|
|
):
|
|
openai_config = OpenAIConfig(
|
|
openai_api_key="mock-api-key",
|
|
openai_api_type=api_type,
|
|
openai_api_base=api_base,
|
|
openai_deployment_name=deployment_name,
|
|
openai_api_version=api_version,
|
|
openai_organization=organization,
|
|
)
|
|
route_config = EndpointConfig(
|
|
name="completions",
|
|
endpoint_type="llm/v1/completions",
|
|
model={
|
|
"provider": "openai",
|
|
"name": "text-davinci-003",
|
|
"config": dict(openai_config),
|
|
},
|
|
)
|
|
provider = OpenAIProvider(route_config)
|
|
assert provider.openai_config == openai_config
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("api_type", "api_base", "deployment_name", "api_version", "organization"),
|
|
[
|
|
# Invalid API Type
|
|
("invalidtype", None, None, None, None),
|
|
# Deployment name is specified when API type is not 'azure' or 'azuread'
|
|
("openai", None, "mock-deployment", None, None),
|
|
# Missing required API base, deployment name, and / or api version fields
|
|
("azure", None, None, None, None),
|
|
("azure", "https://test.openai.azure.com", "mock-dep", None, None),
|
|
# Organization is specified when API type is not 'openai'
|
|
("azure", "https://test.openai.azure.com", "mock-dep", "2023", "test-org"),
|
|
# Missing required API base, deployment name, and / or api version fields
|
|
("azuread", None, None, None, None),
|
|
("azuread", "https://test.openai.azure.com", "mock-dep", None, None),
|
|
# Organization is specified when API type is not 'openai'
|
|
("azuread", "https://test.openai.azure.com", "mock", "2023", "test-org"),
|
|
],
|
|
)
|
|
def test_invalid_openai_configs_throw_on_construction(
|
|
api_type,
|
|
api_base,
|
|
deployment_name,
|
|
api_version,
|
|
organization,
|
|
):
|
|
with pytest.raises(MlflowException, match="OpenAI"):
|
|
OpenAIConfig(
|
|
openai_api_key="mock-api-key",
|
|
openai_api_type=api_type,
|
|
openai_api_base=api_base,
|
|
openai_deployment_name=deployment_name,
|
|
openai_api_version=api_version,
|
|
openai_organization=organization,
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_param_model_is_not_permitted():
|
|
config = azure_config("azuread")
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
payload = {
|
|
"prompt": "This should fail",
|
|
"max_tokens": 5000,
|
|
"model": "something-else",
|
|
}
|
|
with pytest.raises(AIGatewayException, match=r".*") as e:
|
|
await provider.completions(completions.RequestPayload(**payload))
|
|
assert "The parameter 'model' is not permitted" in e.value.detail
|
|
assert e.value.status_code == 422
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_openai_passthrough_chat():
|
|
config = chat_config()
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
|
|
# Mock OpenAI API response
|
|
mock_response = {
|
|
"id": "chatcmpl-123",
|
|
"object": "chat.completion",
|
|
"created": 1677858242,
|
|
"model": "gpt-4o-mini",
|
|
"choices": [
|
|
{
|
|
"index": 0,
|
|
"message": {"role": "assistant", "content": "Hello from passthrough!"},
|
|
"finish_reason": "stop",
|
|
}
|
|
],
|
|
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
|
}
|
|
|
|
with mock.patch(
|
|
"mlflow.gateway.providers.openai.send_request", return_value=mock_response
|
|
) as mock_send:
|
|
payload = {"messages": [{"role": "user", "content": "Hello"}]}
|
|
custom_headers = {
|
|
"X-Custom-Header": "custom-value",
|
|
"X-Request-ID": "req-123",
|
|
"host": "example.com",
|
|
"content-length": "100",
|
|
"authorization": "Bearer key",
|
|
}
|
|
response = await provider.passthrough(
|
|
PassthroughAction.OPENAI_CHAT, payload, headers=custom_headers
|
|
)
|
|
|
|
# Verify send_request was called with correct parameters
|
|
assert mock_send.called
|
|
call_kwargs = mock_send.call_args[1]
|
|
assert call_kwargs["path"] == "chat/completions"
|
|
assert call_kwargs["payload"]["model"] == "gpt-4o-mini"
|
|
assert call_kwargs["payload"]["messages"] == [{"role": "user", "content": "Hello"}]
|
|
|
|
# Verify provider headers are propagated correctly
|
|
assert call_kwargs["headers"]["authorization"] == "Bearer key"
|
|
|
|
# Verify custom headers are propagated correctly
|
|
assert call_kwargs["headers"]["X-Custom-Header"] == "custom-value"
|
|
assert call_kwargs["headers"]["X-Request-ID"] == "req-123"
|
|
|
|
# Verify gateway specific headers are not propagated
|
|
assert "host" not in call_kwargs["headers"]
|
|
assert "content-length" not in call_kwargs["headers"]
|
|
|
|
# Verify response is raw OpenAI format
|
|
assert response == mock_response
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_openai_passthrough_embeddings():
|
|
embeddings_config = {
|
|
"name": "embeddings",
|
|
"endpoint_type": "llm/v1/embeddings",
|
|
"model": {
|
|
"provider": "openai",
|
|
"name": "text-embedding-3-small",
|
|
"config": {
|
|
"openai_api_base": "https://api.openai.com/v1",
|
|
"openai_api_key": "key",
|
|
},
|
|
},
|
|
}
|
|
provider = OpenAIProvider(EndpointConfig(**embeddings_config))
|
|
|
|
# Mock OpenAI API response
|
|
mock_response = {
|
|
"object": "list",
|
|
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
|
|
"model": "text-embedding-3-small",
|
|
"usage": {"prompt_tokens": 5, "total_tokens": 5},
|
|
}
|
|
|
|
with mock.patch(
|
|
"mlflow.gateway.providers.openai.send_request", return_value=mock_response
|
|
) as mock_send:
|
|
payload = {"input": "Test input"}
|
|
custom_headers = {"X-Custom-Header": "custom-value"}
|
|
response = await provider.passthrough(
|
|
PassthroughAction.OPENAI_EMBEDDINGS, payload, headers=custom_headers
|
|
)
|
|
|
|
# Verify send_request was called with correct parameters
|
|
assert mock_send.called
|
|
call_kwargs = mock_send.call_args[1]
|
|
assert call_kwargs["path"] == "embeddings"
|
|
assert call_kwargs["payload"]["model"] == "text-embedding-3-small"
|
|
assert call_kwargs["payload"]["input"] == "Test input"
|
|
|
|
# Verify provider headers are propagated correctly
|
|
assert call_kwargs["headers"]["authorization"] == "Bearer key"
|
|
|
|
# Verify custom headers are propagated correctly
|
|
assert call_kwargs["headers"]["X-Custom-Header"] == "custom-value"
|
|
|
|
# Verify response is raw OpenAI format
|
|
assert response == mock_response
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_openai_passthrough_responses():
|
|
config = chat_config()
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
|
|
# Mock OpenAI Responses API response (using correct Responses API schema)
|
|
mock_response = {
|
|
"id": "resp-123",
|
|
"object": "response",
|
|
"created": 1677858242,
|
|
"model": "gpt-4o-mini",
|
|
"status": "completed",
|
|
"output": [{"type": "text", "text": "Response from Responses API"}],
|
|
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
|
}
|
|
|
|
with mock.patch(
|
|
"mlflow.gateway.providers.openai.send_request", return_value=mock_response
|
|
) as mock_send:
|
|
# Responses API uses 'input' and 'instructions' instead of 'messages'
|
|
payload = {
|
|
"input": [{"type": "text", "text": "Hello"}],
|
|
"instructions": "You are a helpful assistant",
|
|
"response_format": {"type": "text"},
|
|
}
|
|
custom_headers = {"X-Trace-ID": "trace-456"}
|
|
response = await provider.passthrough(
|
|
PassthroughAction.OPENAI_RESPONSES, payload, headers=custom_headers
|
|
)
|
|
|
|
# Verify send_request was called with correct parameters
|
|
assert mock_send.called
|
|
call_kwargs = mock_send.call_args[1]
|
|
assert call_kwargs["path"] == "responses"
|
|
assert call_kwargs["payload"]["model"] == "gpt-4o-mini"
|
|
assert call_kwargs["payload"]["input"] == [{"type": "text", "text": "Hello"}]
|
|
assert call_kwargs["payload"]["instructions"] == "You are a helpful assistant"
|
|
|
|
# Verify provider headers are propagated correctly
|
|
assert call_kwargs["headers"]["authorization"] == "Bearer key"
|
|
|
|
# Verify custom headers are propagated correctly
|
|
assert call_kwargs["headers"]["X-Trace-ID"] == "trace-456"
|
|
|
|
# Verify response is raw OpenAI Responses API format
|
|
assert response == mock_response
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_azure_openai_passthrough_chat_removes_model():
|
|
azure_chat_config = {
|
|
"name": "chat",
|
|
"endpoint_type": "llm/v1/chat",
|
|
"model": {
|
|
"provider": "openai",
|
|
"name": "gpt-4o-mini",
|
|
"config": {
|
|
"openai_api_type": "azure",
|
|
"openai_api_base": "https://my-org.openai.azure.com/",
|
|
"openai_deployment_name": "my-deployment",
|
|
"openai_api_version": "2023-05-15",
|
|
"openai_api_key": "key",
|
|
},
|
|
},
|
|
}
|
|
provider = OpenAIProvider(EndpointConfig(**azure_chat_config))
|
|
|
|
# Mock OpenAI API response
|
|
mock_response = {
|
|
"id": "chatcmpl-123",
|
|
"object": "chat.completion",
|
|
"created": 1677858242,
|
|
"model": "gpt-4o-mini",
|
|
"choices": [
|
|
{
|
|
"index": 0,
|
|
"message": {"role": "assistant", "content": "Hello from Azure!"},
|
|
"finish_reason": "stop",
|
|
}
|
|
],
|
|
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
|
|
}
|
|
|
|
with mock.patch(
|
|
"mlflow.gateway.providers.openai.send_request", return_value=mock_response
|
|
) as mock_send:
|
|
payload = {"messages": [{"role": "user", "content": "Hello"}]}
|
|
custom_headers = {"X-Azure-Custom": "azure-header"}
|
|
response = await provider.passthrough(
|
|
PassthroughAction.OPENAI_CHAT, payload, headers=custom_headers
|
|
)
|
|
|
|
# Verify send_request was called
|
|
assert mock_send.called
|
|
call_kwargs = mock_send.call_args[1]
|
|
assert call_kwargs["path"] == "chat/completions"
|
|
# Azure OpenAI should NOT have model in payload
|
|
assert "model" not in call_kwargs["payload"]
|
|
assert call_kwargs["payload"]["messages"] == [{"role": "user", "content": "Hello"}]
|
|
|
|
# Verify provider headers are propagated correctly (Azure uses api-key header)
|
|
assert call_kwargs["headers"]["api-key"] == "key"
|
|
|
|
# Verify custom headers are propagated correctly
|
|
assert call_kwargs["headers"]["X-Azure-Custom"] == "azure-header"
|
|
|
|
# Verify response is raw OpenAI format
|
|
assert response == mock_response
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_validate_passthrough_action_error_shows_correct_endpoint():
|
|
config = chat_config()
|
|
provider = OpenAIProvider(EndpointConfig(**config))
|
|
|
|
with pytest.raises(
|
|
AIGatewayException,
|
|
match=r"Unsupported passthrough endpoint "
|
|
r"'/gemini/v1beta/models/\{endpoint_name\}:generateContent' for OpenAI provider",
|
|
):
|
|
provider._validate_passthrough_action(PassthroughAction.GEMINI_GENERATE_CONTENT)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_chat_with_structured_output():
|
|
config = EndpointConfig(**chat_config())
|
|
provider = OpenAIProvider(config)
|
|
|
|
json_schema = {
|
|
"name": "math_response",
|
|
"strict": True,
|
|
"schema": {
|
|
"type": "object",
|
|
"properties": {
|
|
"steps": {"type": "array", "items": {"type": "string"}},
|
|
"final_answer": {"type": "string"},
|
|
},
|
|
"required": ["steps", "final_answer"],
|
|
"additionalProperties": False,
|
|
},
|
|
}
|
|
|
|
resp = {
|
|
"id": "chatcmpl-abc123",
|
|
"object": "chat.completion",
|
|
"created": 1677858242,
|
|
"model": "gpt-4o-mini",
|
|
"usage": {
|
|
"prompt_tokens": 13,
|
|
"completion_tokens": 50,
|
|
"total_tokens": 63,
|
|
},
|
|
"choices": [
|
|
{
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": '{"steps": ["1 + 1 = 2"], "final_answer": "2"}',
|
|
},
|
|
"finish_reason": "stop",
|
|
"index": 0,
|
|
}
|
|
],
|
|
}
|
|
|
|
mock_client = mock_http_client(MockAsyncResponse(resp))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client):
|
|
payload = {
|
|
"messages": [{"role": "user", "content": "What is 1+1?"}],
|
|
"temperature": 0.0,
|
|
"response_format": {"type": "json_schema", "json_schema": json_schema},
|
|
}
|
|
response = await provider.chat(chat.RequestPayload(**payload))
|
|
|
|
# Verify the response_format was passed correctly
|
|
assert (
|
|
response.choices[0].message.content == '{"steps": ["1 + 1 = 2"], "final_answer": "2"}'
|
|
)
|
|
assert response.choices[0].finish_reason == "stop"
|
|
|
|
|
|
# Tests for passthrough token extraction
|
|
def test_extract_passthrough_token_usage():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
result = {
|
|
"id": "chatcmpl-123",
|
|
"object": "chat.completion",
|
|
"usage": {
|
|
"prompt_tokens": 10,
|
|
"completion_tokens": 20,
|
|
"total_tokens": 30,
|
|
},
|
|
}
|
|
token_usage = provider._extract_passthrough_token_usage(PassthroughAction.OPENAI_CHAT, result)
|
|
assert token_usage == {
|
|
"input_tokens": 10,
|
|
"output_tokens": 20,
|
|
"total_tokens": 30,
|
|
}
|
|
|
|
|
|
def test_extract_passthrough_token_usage_no_usage():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
result = {
|
|
"id": "chatcmpl-123",
|
|
"object": "chat.completion",
|
|
}
|
|
token_usage = provider._extract_passthrough_token_usage(PassthroughAction.OPENAI_CHAT, result)
|
|
assert token_usage is None
|
|
|
|
|
|
def test_extract_passthrough_token_usage_partial():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
result = {
|
|
"usage": {
|
|
"prompt_tokens": 10,
|
|
},
|
|
}
|
|
token_usage = provider._extract_passthrough_token_usage(PassthroughAction.OPENAI_CHAT, result)
|
|
assert token_usage == {"input_tokens": 10}
|
|
|
|
|
|
def test_extract_streaming_token_usage():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
chunk = (
|
|
b'data: {"id":"chatcmpl-123","usage":'
|
|
b'{"prompt_tokens":10,"completion_tokens":20,"total_tokens":30}}\n\n'
|
|
)
|
|
result = provider._extract_streaming_token_usage(chunk)
|
|
assert result == {
|
|
"input_tokens": 10,
|
|
"output_tokens": 20,
|
|
"total_tokens": 30,
|
|
}
|
|
|
|
|
|
def test_extract_streaming_token_usage_no_usage_in_chunk():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
chunk = b'data: {"id":"chatcmpl-123","choices":[{"delta":{"content":"Hello"}}]}\n\n'
|
|
result = provider._extract_streaming_token_usage(chunk)
|
|
assert result == {}
|
|
|
|
|
|
def test_extract_streaming_token_usage_done_chunk():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
chunk = b"data: [DONE]\n\n"
|
|
result = provider._extract_streaming_token_usage(chunk)
|
|
assert result == {}
|
|
|
|
|
|
def test_extract_streaming_token_usage_invalid_json():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
chunk = b"data: {invalid json}\n\n"
|
|
result = provider._extract_streaming_token_usage(chunk)
|
|
assert result == {}
|
|
|
|
|
|
def test_extract_streaming_token_usage_non_data_line():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
chunk = b"event: message\n\n"
|
|
result = provider._extract_streaming_token_usage(chunk)
|
|
assert result == {}
|
|
|
|
|
|
def test_extract_streaming_token_usage_responses_api():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
# Responses API returns usage in data.response.usage with input_tokens/output_tokens
|
|
chunk = (
|
|
b'data: {"type":"response.completed","response":{"id":"resp_123",'
|
|
b'"usage":{"input_tokens":9,"output_tokens":65,"total_tokens":74}}}\n'
|
|
)
|
|
result = provider._extract_streaming_token_usage(chunk)
|
|
assert result == {
|
|
"input_tokens": 9,
|
|
"output_tokens": 65,
|
|
"total_tokens": 74,
|
|
}
|
|
|
|
|
|
def test_extract_passthrough_token_usage_with_cached_tokens():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
result = {
|
|
"id": "chatcmpl-123",
|
|
"usage": {
|
|
"prompt_tokens": 50,
|
|
"completion_tokens": 20,
|
|
"total_tokens": 70,
|
|
"prompt_tokens_details": {"cached_tokens": 30},
|
|
},
|
|
}
|
|
token_usage = provider._extract_passthrough_token_usage(PassthroughAction.OPENAI_CHAT, result)
|
|
assert token_usage == {
|
|
"input_tokens": 50,
|
|
"output_tokens": 20,
|
|
"total_tokens": 70,
|
|
"cache_read_input_tokens": 30,
|
|
}
|
|
|
|
|
|
def test_extract_passthrough_token_usage_responses_api_with_cached_tokens():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
result = {
|
|
"id": "resp_123",
|
|
"usage": {
|
|
"input_tokens": 100,
|
|
"output_tokens": 50,
|
|
"total_tokens": 150,
|
|
"input_tokens_details": {"cached_tokens": 40},
|
|
},
|
|
}
|
|
token_usage = provider._extract_passthrough_token_usage(
|
|
PassthroughAction.OPENAI_RESPONSES, result
|
|
)
|
|
assert token_usage == {
|
|
"input_tokens": 100,
|
|
"output_tokens": 50,
|
|
"total_tokens": 150,
|
|
"cache_read_input_tokens": 40,
|
|
}
|
|
|
|
|
|
def test_extract_streaming_token_usage_with_cached_tokens():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
chunk = (
|
|
b'data: {"id":"chatcmpl-123","usage":'
|
|
b'{"prompt_tokens":50,"completion_tokens":20,"total_tokens":70,'
|
|
b'"prompt_tokens_details":{"cached_tokens":30}}}\n\n'
|
|
)
|
|
result = provider._extract_streaming_token_usage(chunk)
|
|
assert result == {
|
|
"input_tokens": 50,
|
|
"output_tokens": 20,
|
|
"total_tokens": 70,
|
|
"cache_read_input_tokens": 30,
|
|
}
|
|
|
|
|
|
def test_extract_streaming_token_usage_responses_api_with_cached_tokens():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
chunk = (
|
|
b'data: {"type":"response.completed","response":{"id":"resp_123",'
|
|
b'"usage":{"input_tokens":100,"output_tokens":50,"total_tokens":150,'
|
|
b'"input_tokens_details":{"cached_tokens":40}}}}\n'
|
|
)
|
|
result = provider._extract_streaming_token_usage(chunk)
|
|
assert result == {
|
|
"input_tokens": 100,
|
|
"output_tokens": 50,
|
|
"total_tokens": 150,
|
|
"cache_read_input_tokens": 40,
|
|
}
|
|
|
|
|
|
def test_openai_adapter_build_chat_usage_with_cached_tokens():
|
|
usage_data = {
|
|
"prompt_tokens": 50,
|
|
"completion_tokens": 20,
|
|
"total_tokens": 70,
|
|
"prompt_tokens_details": {"cached_tokens": 30},
|
|
}
|
|
usage = OpenAIAdapter._build_chat_usage(usage_data)
|
|
assert usage.prompt_tokens == 50
|
|
assert usage.completion_tokens == 20
|
|
assert usage.total_tokens == 70
|
|
assert usage.prompt_tokens_details is not None
|
|
assert usage.prompt_tokens_details.cached_tokens == 30
|
|
|
|
|
|
def test_openai_adapter_build_chat_usage_without_cached_tokens():
|
|
usage_data = {
|
|
"prompt_tokens": 50,
|
|
"completion_tokens": 20,
|
|
"total_tokens": 70,
|
|
}
|
|
usage = OpenAIAdapter._build_chat_usage(usage_data)
|
|
assert usage.prompt_tokens == 50
|
|
assert usage.completion_tokens == 20
|
|
assert usage.total_tokens == 70
|
|
assert usage.prompt_tokens_details is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_proxy_non_streaming():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
mock_client = mock_http_client(MockAsyncResponse(chat_response()))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client):
|
|
result = await provider.proxy(
|
|
path="v1/chat/completions",
|
|
payload={"messages": [{"role": "user", "content": "Hello"}]},
|
|
)
|
|
|
|
assert result["id"] == "chatcmpl-abc123"
|
|
mock_client.post.assert_called_once_with(
|
|
"https://api.openai.com/v1/chat/completions",
|
|
json={"messages": [{"role": "user", "content": "Hello"}]},
|
|
timeout=mock.ANY,
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_proxy_strips_mlflow_auth_header_but_preserves_client_key():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
mock_client = mock_http_client(MockAsyncResponse(chat_response()))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client) as mock_session:
|
|
await provider.proxy(
|
|
path="v1/responses",
|
|
payload={"messages": [{"role": "user", "content": "Hello"}]},
|
|
headers={
|
|
"authorization": "Bearer client-key",
|
|
"x-mlflow-authorization": "Basic dXNlcjpwYXNz",
|
|
"user-agent": "codex_cli_rs/1.0",
|
|
},
|
|
)
|
|
|
|
sent_headers = mock_session.call_args.kwargs["headers"]
|
|
assert sent_headers["authorization"] == "Bearer client-key"
|
|
assert "x-mlflow-authorization" not in {k.lower() for k in sent_headers}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"header_name",
|
|
["X-MLflow-Authorization", "X-MLFLOW-AUTHORIZATION", "x-Mlflow-authorization"],
|
|
)
|
|
async def test_proxy_strips_mlflow_auth_header_mixed_case(header_name):
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
mock_client = mock_http_client(MockAsyncResponse(chat_response()))
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client) as mock_session:
|
|
await provider.proxy(
|
|
path="v1/responses",
|
|
payload={"messages": [{"role": "user", "content": "Hello"}]},
|
|
headers={
|
|
"authorization": "Bearer client-key",
|
|
header_name: "Basic dXNlcjpwYXNz",
|
|
"user-agent": "codex_cli_rs/1.0",
|
|
},
|
|
)
|
|
|
|
sent_headers = mock_session.call_args.kwargs["headers"]
|
|
assert sent_headers["authorization"] == "Bearer client-key"
|
|
assert "x-mlflow-authorization" not in {k.lower() for k in sent_headers}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_proxy_streaming():
|
|
provider = OpenAIProvider(EndpointConfig(**chat_config()))
|
|
chunk_data = (
|
|
b'data: {"id":"chatcmpl-1","object":"chat.completion.chunk","created":1,'
|
|
b'"model":"gpt-4o-mini","choices":[{"index":0,"delta":{"content":"Hi"},'
|
|
b'"finish_reason":null}]}\n\n'
|
|
)
|
|
chunks = [chunk_data, b"data: [DONE]\n\n"]
|
|
mock_client = mock_http_client(
|
|
MockAsyncStreamingResponse(chunks, headers={"Content-Type": "text/event-stream"})
|
|
)
|
|
|
|
with mock.patch("aiohttp.ClientSession", return_value=mock_client):
|
|
result = await provider.proxy(
|
|
path="v1/chat/completions",
|
|
payload={"messages": [{"role": "user", "content": "Hello"}], "stream": True},
|
|
)
|
|
collected = [chunk async for chunk in result]
|
|
|
|
assert len(collected) == 2
|
|
assert b"chatcmpl-1" in collected[0]
|
|
assert b"[DONE]" in collected[1]
|