Fix: Send Gemini API key via x-goog-api-key header with custom api_base (#16085)

* Add gemini api key in the custom api url

* Update tests

* Use api key n the header

* Use api key n the header

* fix mypy error

* fix mypy error

* fix test gemini auth
This commit is contained in:
Sameer Kankute 2025-11-05 20:42:13 +05:30 • committed by GitHub
parent 781f9df883
commit c45fad3855
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 74 additions and 10 deletions

View file

@ -1743,13 +1743,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
messages: List[AllMessageValues],
optional_params: Dict,
litellm_params: Dict,
api_key: Optional[str] = None,
api_key: Optional[Union[str, Dict]] = None,
api_base: Optional[str] = None,
) -> Dict:
default_headers = {
"Content-Type": "application/json",
}
if api_key is not None:
if isinstance(api_key, dict):
default_headers.update(api_key)
elif api_key is not None:
default_headers["Authorization"] = f"Bearer {api_key}"
if headers is not None:
default_headers.update(headers)

View file

@ -308,9 +308,8 @@ class VertexBase:
raise ValueError(
"Missing gemini_api_key, please set `GEMINI_API_KEY`"
)
auth_header = (
gemini_api_key # cloudflare expects api key as bearer token
)
if gemini_api_key is not None:
auth_header = {"x-goog-api-key": gemini_api_key} # type: ignore[assignment]
else:
url = "{}:{}".format(api_base, endpoint)

View file

@ -378,8 +378,8 @@ async def test_gemini_custom_api_base_proxy_integration():
expected_url = f"{custom_api_base}/models/{model}:{endpoint}"
assert result_url == expected_url, f"Expected {expected_url}, got {result_url}"
# Verify the auth header is set to the API key
assert auth_header == "test-api-key", f"Expected 'test-api-key', got {auth_header}"
# Verify the auth header is set to the API key as a dictionary
assert auth_header == {"x-goog-api-key": "test-api-key"}, f"Expected {{'x-goog-api-key': 'test-api-key'}}, got {auth_header}"
print(f"✅ Custom API base URL construction test passed: {result_url}")
@ -399,6 +399,9 @@ async def test_gemini_custom_api_base_proxy_integration():
expected_streaming_url = f"{custom_api_base}/models/{model}:{endpoint}?alt=sse"
assert result_url_streaming == expected_streaming_url, f"Expected {expected_streaming_url}, got {result_url_streaming}"
# Verify the auth header is also set correctly for streaming
assert auth_header_streaming == {"x-goog-api-key": "test-api-key"}, f"Expected {{'x-goog-api-key': 'test-api-key'}}, got {auth_header_streaming}"
print(f"✅ Custom API base streaming URL test passed: {result_url_streaming}")
# Test case 3: Error handling - missing API key
@ -462,7 +465,8 @@ async def test_gemini_proxy_config_with_custom_api_base():
expected_url = f"{model_config['litellm_params']['api_base']}/models/{model}:generateContent"
assert result_url == expected_url, f"Expected {expected_url}, got {result_url} for model {model}"
assert auth_header == model_config["litellm_params"]["api_key"], f"Expected API key, got {auth_header} for model {model}"
expected_auth_header = {"x-goog-api-key": model_config["litellm_params"]["api_key"]}
assert auth_header == expected_auth_header, f"Expected {expected_auth_header}, got {auth_header} for model {model}"
print(f"✅ Model {model} configuration test passed: {result_url}")

View file

@ -718,7 +718,7 @@ class TestVertexBase:
None,
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent",
"gemini-2.5-flash-lite",
"test-api-key",
{"x-goog-api-key": "test-api-key"},
"https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent"
),
# Test case 2: Gemini with custom API base and streaming
@ -731,7 +731,7 @@ class TestVertexBase:
None,
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent",
"gemini-2.5-flash-lite",
"test-api-key",
{"x-goog-api-key": "test-api-key"},
"https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse"
),
# Test case 3: Non-Gemini provider with custom API base
@ -878,3 +878,62 @@ class TestVertexBase:
expected_no_streaming_url = "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent"
assert result_url_no_streaming == expected_no_streaming_url, f"Expected {expected_no_streaming_url}, got {result_url_no_streaming}"
@pytest.mark.parametrize(
"api_base, custom_llm_provider, gemini_api_key, endpoint, stream, auth_header, url, model, expected_auth_header, expected_url",
[
# Test case 1: Gemini with custom API base
(
"https://proxy.example.com/generativelanguage.googleapis.com/v1beta",
"gemini",
"test-api-key",
"generateContent",
False,
None,
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent",
"gemini-2.5-flash-lite",
{"x-goog-api-key": "test-api-key"},
"https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent"
),
# Test case 2: Gemini with custom API base and streaming
(
"https://proxy.example.com/generativelanguage.googleapis.com/v1beta",
"gemini",
"test-api-key",
"generateContent",
True,
None,
"https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent",
"gemini-2.5-flash-lite",
{"x-goog-api-key": "test-api-key"},
"https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse"
),
],
)
def test_check_custom_proxy_minimal_gemini_key_param(
self,
api_base,
custom_llm_provider,
gemini_api_key,
endpoint,
stream,
auth_header,
url,
model,
expected_auth_header,
expected_url,
):
"""Single focused test to ensure ?key is appended (and &alt=sse for streaming)."""
vertex_base = VertexBase()
result_auth_header, result_url = vertex_base._check_custom_proxy(
api_base=api_base,
custom_llm_provider=custom_llm_provider,
gemini_api_key=gemini_api_key,
endpoint=endpoint,
stream=stream,
auth_header=auth_header,
url=url,
model=model,
)
assert result_auth_header == expected_auth_header
assert result_url == expected_url