-
Notifications
You must be signed in to change notification settings - Fork 3.1k
Expand file tree
/
Copy pathtest_openai_image_generator.py
More file actions
357 lines (307 loc) · 15.3 KB
/
Copy pathtest_openai_image_generator.py
File metadata and controls
357 lines (307 loc) · 15.3 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
# SPDX-FileCopyrightText: 2022-present deepset GmbH <info@deepset.ai>
#
# SPDX-License-Identifier: Apache-2.0
import base64
import os
from unittest.mock import AsyncMock, MagicMock, Mock, patch
import pytest
from openai import AsyncOpenAI
from openai.types import ImagesResponse
from openai.types.image import Image
import haystack.components.generators.openai_image_generator as openai_image_generator_module
from haystack.components.generators.openai_image_generator import OpenAIImageGenerator
from haystack.utils import Secret
@pytest.fixture
def mock_image_response():
with patch("openai.resources.images.Images.generate") as mock_image_generate:
image_response = ImagesResponse(
created=1630000000, data=[Image(b64_json="test-b64-json", revised_prompt="test-prompt")]
)
mock_image_generate.return_value = image_response
yield mock_image_generate
class TestOpenAIImageGenerator:
def test_init_default(self, monkeypatch: pytest.MonkeyPatch) -> None:
component = OpenAIImageGenerator()
assert component.model == "gpt-image-2"
assert component.quality == "auto"
assert component.size == "1024x1024"
assert component.api_key == Secret.from_env_var("OPENAI_API_KEY")
assert component.api_base_url is None
assert component.organization is None
assert component.timeout is None
assert component.max_retries is None
assert component.http_client_kwargs is None
assert component.client is None
assert component.async_client is None
def test_init_with_params(self, monkeypatch: pytest.MonkeyPatch) -> None:
component = OpenAIImageGenerator(
model="gpt-image-1",
quality="high",
size="1024x1536",
api_key=Secret.from_env_var("EXAMPLE_API_KEY"),
api_base_url="https://api.openai.com",
organization="test-org",
timeout=60,
max_retries=10,
)
assert component.model == "gpt-image-1"
assert component.quality == "high"
assert component.size == "1024x1536"
assert component.api_key == Secret.from_env_var("EXAMPLE_API_KEY")
assert component.api_base_url == "https://api.openai.com"
assert component.organization == "test-org"
assert pytest.approx(component.timeout) == 60.0
assert component.max_retries == 10
assert component.client is None
assert component.async_client is None
def test_init_max_retries_0(self, monkeypatch: pytest.MonkeyPatch) -> None:
component = OpenAIImageGenerator(max_retries=0)
assert component.max_retries == 0
def test_init_invalid_quality_falls_back_to_auto(self, caplog: pytest.LogCaptureFixture) -> None:
component = OpenAIImageGenerator(quality="hd") # type: ignore[arg-type]
assert component.quality == "auto"
assert "Invalid quality" in caplog.text
def test_init_non_default_response_format_warns(self, caplog: pytest.LogCaptureFixture) -> None:
OpenAIImageGenerator(response_format="url") # type: ignore[arg-type]
assert "response_format is ignored" in caplog.text
def test_to_dict(self) -> None:
generator = OpenAIImageGenerator()
data = generator.to_dict()
assert data == {
"type": "haystack.components.generators.openai_image_generator.OpenAIImageGenerator",
"init_parameters": {
"model": "gpt-image-2",
"quality": "auto",
"size": "1024x1024",
"api_key": {"type": "env_var", "env_vars": ["OPENAI_API_KEY"], "strict": True},
"api_base_url": None,
"organization": None,
"timeout": None,
"max_retries": None,
"http_client_kwargs": None,
},
}
def test_to_dict_with_params(self) -> None:
generator = OpenAIImageGenerator(
model="gpt-image-1",
quality="high",
size="1024x1536",
api_key=Secret.from_env_var("EXAMPLE_API_KEY"),
api_base_url="https://api.openai.com",
organization="test-org",
timeout=60,
max_retries=10,
http_client_kwargs={"proxy": "http://localhost:8080"},
)
data = generator.to_dict()
assert data == {
"type": "haystack.components.generators.openai_image_generator.OpenAIImageGenerator",
"init_parameters": {
"model": "gpt-image-1",
"quality": "high",
"size": "1024x1536",
"api_key": {"type": "env_var", "env_vars": ["EXAMPLE_API_KEY"], "strict": True},
"api_base_url": "https://api.openai.com",
"organization": "test-org",
"timeout": 60,
"max_retries": 10,
"http_client_kwargs": {"proxy": "http://localhost:8080"},
},
}
def test_from_dict(self) -> None:
data = {
"type": "haystack.components.generators.openai_image_generator.OpenAIImageGenerator",
"init_parameters": {
"model": "gpt-image-2",
"quality": "auto",
"size": "1024x1024",
"api_key": {"type": "env_var", "env_vars": ["OPENAI_API_KEY"], "strict": True},
"api_base_url": None,
"organization": None,
"http_client_kwargs": None,
},
}
generator = OpenAIImageGenerator.from_dict(data)
assert generator.model == "gpt-image-2"
assert generator.quality == "auto"
assert generator.size == "1024x1024"
assert generator.api_key.to_dict() == {"type": "env_var", "env_vars": ["OPENAI_API_KEY"], "strict": True}
assert generator.http_client_kwargs is None
def test_to_dict_from_dict_roundtrip_preserves_client_settings(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""`timeout` and `max_retries` decide the client the component builds, so a pipeline
that survives a save/load round trip must keep them. Without them in `to_dict` they
silently revert to the `OPENAI_TIMEOUT`/`OPENAI_MAX_RETRIES` fallbacks (30s, 5)."""
monkeypatch.setenv("OPENAI_API_KEY", "test-api-key")
generator = OpenAIImageGenerator(timeout=120.0, max_retries=10)
restored = OpenAIImageGenerator.from_dict(generator.to_dict())
assert restored.timeout == 120.0
assert restored.max_retries == 10
assert restored._client_kwargs()["timeout"] == 120.0
assert restored._client_kwargs()["max_retries"] == 10
def test_from_dict_default_params(self) -> None:
data = {
"type": "haystack.components.generators.openai_image_generator.OpenAIImageGenerator",
"init_parameters": {},
}
generator = OpenAIImageGenerator.from_dict(data)
assert generator.model == "gpt-image-2"
assert generator.quality == "auto"
assert generator.size == "1024x1024"
assert generator.api_key.to_dict() == {"type": "env_var", "env_vars": ["OPENAI_API_KEY"], "strict": True}
assert generator.api_base_url is None
assert generator.organization is None
assert generator.timeout is None
assert generator.max_retries is None
assert generator.http_client_kwargs is None
def test_run(self, mock_image_response: MagicMock) -> None:
generator = OpenAIImageGenerator(api_key=Secret.from_token("test-api-key"))
response = generator.run("Show me a picture of a black cat.")
assert generator.client is not None
assert isinstance(response, dict)
assert "images" in response and "revised_prompt" in response
assert response["images"] == ["test-b64-json"]
assert response["revised_prompt"] == "test-prompt"
@pytest.mark.skipif(
not os.environ.get("OPENAI_API_KEY", None),
reason="Export an env var called OPENAI_API_KEY containing the OpenAI API key to run this test.",
)
@pytest.mark.integration
@pytest.mark.slow
def test_live_run(self) -> None:
generator = OpenAIImageGenerator(model="gpt-image-1-mini", size="1024x1024", quality="low")
response = generator.run("A nice cat")
assert isinstance(response, dict)
assert isinstance(response["revised_prompt"], str)
image_str = response["images"][0]
assert isinstance(image_str, str) and image_str
decoded = base64.b64decode(image_str, validate=True)
assert decoded.startswith(b"\x89PNG\r\n\x1a\n")
class TestOpenAIImageGeneratorAsync:
def test_async_client_none_before_warm_up(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("OPENAI_API_KEY", "test-api-key")
component = OpenAIImageGenerator()
assert component.async_client is None
@pytest.mark.asyncio
async def test_async_client_after_warm_up_async(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("OPENAI_API_KEY", "test-api-key")
component = OpenAIImageGenerator()
await component.warm_up_async()
assert isinstance(component.async_client, AsyncOpenAI)
assert component.async_client.api_key == "test-api-key"
@pytest.mark.asyncio
async def test_run_async(self) -> None:
generator = OpenAIImageGenerator(api_key=Secret.from_token("test-api-key"))
image_response = ImagesResponse(
created=1630000000, data=[Image(b64_json="test-b64-json", revised_prompt="test-prompt")]
)
mock_async_client = Mock()
mock_async_client.images.generate = AsyncMock(return_value=image_response)
generator.async_client = mock_async_client
response = await generator.run_async("Show me a picture of a black cat.")
assert isinstance(response, dict)
assert "images" in response and "revised_prompt" in response
assert response["images"] == ["test-b64-json"]
assert response["revised_prompt"] == "test-prompt"
mock_async_client.images.generate.assert_awaited_once()
@pytest.mark.asyncio
async def test_run_async_triggers_warm_up(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("OPENAI_API_KEY", "test-api-key")
generator = OpenAIImageGenerator()
assert generator.async_client is None
image_response = ImagesResponse(
created=1630000000, data=[Image(b64_json="test-b64-json", revised_prompt="test-prompt")]
)
with patch("openai.resources.images.AsyncImages.generate", new=AsyncMock(return_value=image_response)):
response = await generator.run_async("Show me a picture of a black cat.")
assert isinstance(generator.async_client, AsyncOpenAI)
assert response["images"] == ["test-b64-json"]
assert response["revised_prompt"] == "test-prompt"
@pytest.mark.asyncio
@pytest.mark.skipif(
not os.environ.get("OPENAI_API_KEY", None),
reason="Export an env var called OPENAI_API_KEY containing the OpenAI API key to run this test.",
)
@pytest.mark.integration
@pytest.mark.slow
async def test_live_run_async(self) -> None:
generator = OpenAIImageGenerator(model="gpt-image-1-mini", size="1024x1024", quality="low")
response = await generator.run_async("A nice cat")
assert isinstance(response, dict)
assert isinstance(response["revised_prompt"], str)
image_str = response["images"][0]
assert isinstance(image_str, str) and image_str
decoded = base64.b64decode(image_str, validate=True)
assert decoded.startswith(b"\x89PNG\r\n\x1a\n")
@pytest.fixture
def mock_openai_clients(monkeypatch):
monkeypatch.setenv("OPENAI_API_KEY", "fake")
sync_cls = MagicMock(name="OpenAI")
async_cls = MagicMock(name="AsyncOpenAI")
async_cls.return_value.close = AsyncMock()
monkeypatch.setattr(openai_image_generator_module, "OpenAI", sync_cls)
monkeypatch.setattr(openai_image_generator_module, "AsyncOpenAI", async_cls)
return sync_cls, async_cls
class TestComponentLifecycle:
def test_warm_up_uses_default_timeout_and_max_retries(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("OPENAI_API_KEY", "fake-api-key")
generator = OpenAIImageGenerator()
generator.warm_up()
assert generator.client is not None
assert generator.client.max_retries == 5
assert generator.client.timeout == 30.0
def test_warm_up_uses_timeout_and_max_retries_from_parameters(self) -> None:
generator = OpenAIImageGenerator(api_key=Secret.from_token("fake-api-key"), timeout=40.0, max_retries=1)
generator.warm_up()
assert generator.client is not None
assert generator.client.max_retries == 1
assert generator.client.timeout == 40.0
def test_warm_up_uses_timeout_and_max_retries_from_env_vars(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("OPENAI_TIMEOUT", "100")
monkeypatch.setenv("OPENAI_MAX_RETRIES", "10")
generator = OpenAIImageGenerator(api_key=Secret.from_token("fake-api-key"))
generator.warm_up()
assert generator.client is not None
assert generator.client.max_retries == 10
assert generator.client.timeout == 100.0
def test_key_resolved_at_warm_up_not_init(self, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
generator = OpenAIImageGenerator()
with pytest.raises(ValueError, match="None of the .* environment variables are set"):
generator.warm_up()
def test_sync_lifecycle(self, mock_openai_clients: tuple[MagicMock, MagicMock]) -> None:
sync_cls, _ = mock_openai_clients
generator = OpenAIImageGenerator()
assert generator.client is None
assert generator.async_client is None
generator.warm_up()
assert generator.client is sync_cls.return_value
assert generator.async_client is None
generator.close()
sync_cls.return_value.close.assert_called_once() # type: ignore[attr-defined]
assert generator.client is None
async def test_async_lifecycle(self, mock_openai_clients: tuple[MagicMock, MagicMock]) -> None:
_, async_cls = mock_openai_clients
generator = OpenAIImageGenerator()
await generator.warm_up_async()
assert generator.async_client is async_cls.return_value
assert generator.client is None
await generator.close_async()
async_cls.return_value.close.assert_awaited_once() # type: ignore[union-attr]
assert generator.async_client is None
async def test_close_is_safe_without_warm_up(self, mock_openai_clients: tuple[MagicMock, MagicMock]) -> None:
generator = OpenAIImageGenerator()
generator.close()
await generator.close_async()
assert generator.client is None
assert generator.async_client is None
async def test_close_and_close_async_are_independent(
self, mock_openai_clients: tuple[MagicMock, MagicMock]
) -> None:
generator = OpenAIImageGenerator()
generator.warm_up()
await generator.warm_up_async()
generator.close()
assert generator.client is None
assert generator.async_client is not None
await generator.close_async()
assert generator.async_client is None