fix(image_generation): propagate extra_headers to OpenAI image generation

Add headers parameter to image_generation() and aimage_generation() methods
in OpenAI provider, and pass headers from images/main.py to ensure custom
headers like cf-aig-authorization are properly forwarded to the OpenAI API.
Aligns behavior with completion() method and Azure provider implementation.
This commit is contained in:
Zero Clover 2026-02-25 01:29:06 +08:00
parent 4652c73259
commit a1c939b2ef
Failed to extract signature
2 changed files with 8 additions and 1 deletions

View file

@ -483,6 +483,7 @@ def image_generation( # noqa: PLR0915
organization=organization,
aimg_generation=aimg_generation,
client=client,
headers=headers,
)
elif custom_llm_provider == "bedrock":
if model is None:

View file

@ -1401,6 +1401,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
client=None,
max_retries=None,
organization: Optional[str] = None,
headers: Optional[dict] = None,
):
response = None
try:
@ -1414,6 +1415,8 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
client=client,
)
if headers:
data["extra_headers"] = headers
response = await openai_aclient.images.generate(**data, timeout=timeout) # type: ignore
stringified_response = response.model_dump()
## LOGGING
@ -1446,6 +1449,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
client=None,
aimg_generation=None,
organization: Optional[str] = None,
headers: Optional[dict] = None,
) -> ImageResponse:
data = {}
try:
@ -1455,7 +1459,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
raise OpenAIError(status_code=422, message="max retries must be an int")
if aimg_generation is True:
return self.aimage_generation(data=data, prompt=prompt, logging_obj=logging_obj, model_response=model_response, api_base=api_base, api_key=api_key, timeout=timeout, client=client, max_retries=max_retries, organization=organization) # type: ignore
return self.aimage_generation(data=data, prompt=prompt, logging_obj=logging_obj, model_response=model_response, api_base=api_base, api_key=api_key, timeout=timeout, client=client, max_retries=max_retries, organization=organization, headers=headers) # type: ignore
openai_client: OpenAI = self._get_openai_client( # type: ignore
is_async=False,
@ -1480,6 +1484,8 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
)
## COMPLETION CALL
if headers:
data["extra_headers"] = headers
_response = openai_client.images.generate(**data, timeout=timeout) # type: ignore
response = _response.model_dump()