From 2da18a70d30917313fa57c7be1c414a131511007 Mon Sep 17 00:00:00 2001 From: Deniz Dalkilic Date: Fri, 7 Aug 2026 14:20:04 +0200 Subject: [PATCH] Add OpenAI gateway API key header support --- .env.example | 4 ++ assert_ai/core/model_client.py | 30 +++++++++++++ assert_ai/init/_llm.py | 2 + assert_ai/integrations/acs/language_model.py | 6 ++- tests/test_acs_language_model.py | 18 ++++++++ tests/test_init_llm.py | 29 ++++++++++++ tests/test_model_client.py | 47 ++++++++++++++++++++ 7 files changed, 135 insertions(+), 1 deletion(-) diff --git a/.env.example b/.env.example index 2b0f9baf..503878f6 100644 --- a/.env.example +++ b/.env.example @@ -26,6 +26,10 @@ AZURE_API_BASE=https://your-resource.openai.azure.com/ # OpenAI (for model strings like openai/gpt-4o) # OPENAI_API_KEY= # OPENAI_MODEL=gpt-4o +# Optional for OpenAI-compatible API gateways that require the key in a +# provider-specific header in addition to standard Bearer authentication. +# The header value is read from OPENAI_API_KEY; do not duplicate the secret. +# ASSERT_OPENAI_API_KEY_HEADER=api-key # Anthropic (for model strings like anthropic/claude-3.5-sonnet) # ANTHROPIC_API_KEY= diff --git a/assert_ai/core/model_client.py b/assert_ai/core/model_client.py index 62476c24..e071887f 100644 --- a/assert_ai/core/model_client.py +++ b/assert_ai/core/model_client.py @@ -442,6 +442,34 @@ def _maybe_inject_azure_aad_token(model: str, payload: dict[str, Any]) -> None: payload["azure_ad_token_provider"] = provider +def _maybe_inject_openai_api_key_header(model: str, payload: dict[str, Any]) -> None: + """Send ``OPENAI_API_KEY`` under an opt-in compatibility header. + + OpenAI-compatible gateways sometimes require an API-management header in + addition to the standard Bearer token. ``ASSERT_OPENAI_API_KEY_HEADER`` + names that header without duplicating the secret into another environment + variable. Explicit per-call headers take precedence. + """ + if _model_family(model) != "openai": + return + header_name = os.environ.get("ASSERT_OPENAI_API_KEY_HEADER", "").strip() + api_key = os.environ.get("OPENAI_API_KEY", "").strip() + if not header_name or not api_key: + return + + configured_headers = payload.get("extra_headers") + if configured_headers is None: + headers: dict[str, Any] = {} + elif isinstance(configured_headers, Mapping): + headers = dict(configured_headers) + else: + raise ValueError("extra_headers must be a mapping") + + if not any(str(name).lower() == header_name.lower() for name in headers): + headers[header_name] = api_key + payload["extra_headers"] = headers + + def _supports_web_search_preview(model: str) -> bool: """Whether this model can use the Responses API web_search_preview tool. @@ -647,6 +675,7 @@ def _build_chat_payload( payload["reasoning_effort"] = resolved_options.reasoning_effort _maybe_inject_azure_aad_token(model, payload) payload.update(resolved_options.extra_kwargs) + _maybe_inject_openai_api_key_header(model, payload) return payload @@ -674,6 +703,7 @@ def _build_responses_payload( payload["reasoning_effort"] = resolved_options.reasoning_effort _maybe_inject_azure_aad_token(model, payload) payload.update(resolved_options.extra_kwargs) + _maybe_inject_openai_api_key_header(model, payload) return payload diff --git a/assert_ai/init/_llm.py b/assert_ai/init/_llm.py index 912b7d17..31c10afe 100644 --- a/assert_ai/init/_llm.py +++ b/assert_ai/init/_llm.py @@ -33,6 +33,7 @@ def chat_completion( _classify_llm_error, _force_chat_completions, _maybe_inject_azure_aad_token, + _maybe_inject_openai_api_key_header, ) kwargs: dict[str, Any] = { @@ -50,6 +51,7 @@ def chat_completion( # to whatever key/cred LiteLLM finds in the environment, defeating # the documented ``ASSERT_AZURE_USE_AAD=1`` opt-in. _maybe_inject_azure_aad_token(model, kwargs) + _maybe_inject_openai_api_key_header(model, kwargs) try: response = litellm.completion(**kwargs) diff --git a/assert_ai/integrations/acs/language_model.py b/assert_ai/integrations/acs/language_model.py index 0c4fc85c..e50839ef 100644 --- a/assert_ai/integrations/acs/language_model.py +++ b/assert_ai/integrations/acs/language_model.py @@ -34,7 +34,10 @@ def __init__( def complete(self, system: str, user: str) -> str: """Return the raw assistant text for the ACS generator's JSON plan prompt.""" - from assert_ai.core.model_client import _maybe_inject_azure_aad_token + from assert_ai.core.model_client import ( + _maybe_inject_azure_aad_token, + _maybe_inject_openai_api_key_header, + ) litellm = _assert_litellm_module() payload: dict[str, Any] = { @@ -54,6 +57,7 @@ def complete(self, system: str, user: str) -> str: # models and for the ``key`` auth mode, so existing API-key # users are unaffected. _maybe_inject_azure_aad_token(self.model, payload) + _maybe_inject_openai_api_key_header(self.model, payload) try: response = litellm.completion(**payload) diff --git a/tests/test_acs_language_model.py b/tests/test_acs_language_model.py index 6323a414..a8249b20 100644 --- a/tests/test_acs_language_model.py +++ b/tests/test_acs_language_model.py @@ -104,6 +104,24 @@ def fake_completion(**kwargs): assert "response_format" not in calls[1] +def test_assert_language_model_injects_openai_compatibility_header(monkeypatch) -> None: + import litellm + + captured: dict = {} + + def fake_completion(**kwargs): + captured.update(kwargs) + return _ok_response() + + monkeypatch.setenv("ASSERT_OPENAI_API_KEY_HEADER", "api-key") + monkeypatch.setenv("OPENAI_API_KEY", "gateway-secret") + monkeypatch.setattr(litellm, "completion", fake_completion) + + AssertLanguageModel("openai/gpt-5.4").complete("sys", "usr") + + assert captured["extra_headers"] == {"api-key": "gateway-secret"} + + # ── Azure AD token provider injection (PR #237 follow-up) ────────────── # # The ACS LiteLLM call site is the third place in the codebase that diff --git a/tests/test_init_llm.py b/tests/test_init_llm.py index c66aa604..adbd7f4d 100644 --- a/tests/test_init_llm.py +++ b/tests/test_init_llm.py @@ -147,5 +147,34 @@ def fake_completion(**kwargs: Any) -> Any: self.assertNotIn("azure_ad_token_provider", captured) +class InitChatCompletionOpenAICompatibilityHeaderTest(unittest.TestCase): + def test_openai_model_gets_configured_api_key_header(self) -> None: + from assert_ai.init import _llm + + captured: dict[str, Any] = {} + + def fake_completion(**kwargs: Any) -> Any: + captured.update(kwargs) + return _fake_response("ok") + + with ( + patch.dict( + os.environ, + { + "ASSERT_OPENAI_API_KEY_HEADER": "api-key", + "OPENAI_API_KEY": "gateway-secret", + }, + ), + patch("litellm.completion", side_effect=fake_completion), + ): + result = _llm.chat_completion( + model="openai/gpt-5.4", + messages=[{"role": "user", "content": "hi"}], + ) + + self.assertEqual(result, "ok") + self.assertEqual(captured["extra_headers"], {"api-key": "gateway-secret"}) + + if __name__ == "__main__": unittest.main() diff --git a/tests/test_model_client.py b/tests/test_model_client.py index fc5aa7a1..9f83cb21 100644 --- a/tests/test_model_client.py +++ b/tests/test_model_client.py @@ -58,6 +58,53 @@ async def fake_acompletion(**kwargs): self.assertEqual(response.request_payload["model"], "openai/gpt-5-mini") self.assertEqual(response.request_payload["messages"], [{"role": "user", "content": "say hi"}]) + async def test_generate_injects_configured_openai_api_key_header(self) -> None: + captured: dict[str, object] = {} + + async def fake_acompletion(**kwargs): + captured.update(kwargs) + return { + "choices": [ + { + "finish_reason": "stop", + "message": {"role": "assistant", "content": "ok"}, + } + ] + } + + fake_litellm = SimpleNamespace(acompletion=fake_acompletion) + with ( + patch.dict( + os.environ, + { + "ASSERT_OPENAI_API_KEY_HEADER": "api-key", + "OPENAI_API_KEY": "gateway-secret", + }, + ), + patch.object(model_client, "_get_litellm_module", return_value=fake_litellm), + ): + await model_client.generate("openai/gpt-5-mini", "say hi") + + self.assertEqual(captured["extra_headers"], {"api-key": "gateway-secret"}) + + def test_explicit_openai_api_key_header_takes_precedence(self) -> None: + with patch.dict( + os.environ, + { + "ASSERT_OPENAI_API_KEY_HEADER": "api-key", + "OPENAI_API_KEY": "environment-secret", + }, + ): + payload = model_client._build_chat_payload( + "openai/gpt-5-mini", + "say hi", + model_client.GenerateOptions( + extra_kwargs={"extra_headers": {"Api-Key": "explicit-secret"}} + ), + ) + + self.assertEqual(payload["extra_headers"], {"Api-Key": "explicit-secret"}) + async def test_generate_structured_adds_json_schema_response_format(self) -> None: captured: dict[str, object] = {}