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]