返回提交历史
Modified
g4f/client/__init__.py
+0
-1
Deleted
g4f/client/async_client.py
+0
-339
Modified
g4f/client/client.py
+335
-132
XFEstudio/gpt4free
feat(g4f/client/async_client.py, g4f/client/async_client.py): enhance async and sync handling in client
0d868f64
代码差异
3 个文件
+335
-472
@@ -1,3 +1,2 @@
1
1
from .stubs import ChatCompletion, ChatCompletionChunk, ImagesResponse
2
2
from .client import Client
3
from .async_client import AsyncClient
@@ -1,339 +0,0 @@
1
from __future__ import annotations
2
3
import os
4
import time
5
import random
6
import string
7
import logging
8
import asyncio
9
from typing import Union, AsyncIterator
10
from ..providers.base_provider import AsyncGeneratorProvider
11
from ..image import ImageResponse, to_image, to_data_uri
12
from ..typing import Messages, ImageType
13
from ..providers.types import BaseProvider, ProviderType, FinishReason
14
from ..providers.conversation import BaseConversation
15
from ..image import ImageResponse as ImageProviderResponse
16
from ..errors import NoImageResponseError
17
from .stubs import ChatCompletion, ChatCompletionChunk, Image, ImagesResponse
18
from .image_models import ImageModels
19
from .types import IterResponse, ImageProvider
20
from .types import Client as BaseClient
21
from .service import get_model_and_provider, get_last_provider
22
from .helper import find_stop, filter_json, filter_none
23
from ..models import ModelUtils
24
from ..Provider import IterListProvider
25
from .helper import cast_iter_async
26
27
try:
28
anext # Python 3.8+
29
except NameError:
30
async def anext(aiter):
31
try:
32
return await aiter.__anext__()
33
except StopAsyncIteration:
34
raise StopIteration
35
36
async def safe_aclose(generator):
37
try:
38
await generator.aclose()
39
except Exception as e:
40
logging.warning(f"Error while closing generator: {e}")
41
42
async def iter_response(
43
response: AsyncIterator[str],
44
stream: bool,
45
response_format: dict = None,
46
max_tokens: int = None,
47
stop: list = None
48
) -> AsyncIterator[Union[ChatCompletion, ChatCompletionChunk]]:
49
content = ""
50
finish_reason = None
51
completion_id = ''.join(random.choices(string.ascii_letters + string.digits, k=28))
52
idx = 0
53
54
try:
55
async for chunk in response:
56
if isinstance(chunk, FinishReason):
57
finish_reason = chunk.reason
58
break
59
elif isinstance(chunk, BaseConversation):
60
yield chunk
61
continue
62
63
content += str(chunk)
64
idx += 1
65
66
if max_tokens is not None and idx >= max_tokens:
67
finish_reason = "length"
68
69
first, content, chunk = find_stop(stop, content, chunk if stream else None)
70
71
if first != -1:
72
finish_reason = "stop"
73
74
if stream:
75
yield ChatCompletionChunk(chunk, None, completion_id, int(time.time()))
76
77
if finish_reason is not None:
78
break
79
80
finish_reason = "stop" if finish_reason is None else finish_reason
81
82
if stream:
83
yield ChatCompletionChunk(None, finish_reason, completion_id, int(time.time()))
84
else:
85
if response_format is not None and "type" in response_format:
86
if response_format["type"] == "json_object":
87
content = filter_json(content)
88
yield ChatCompletion(content, finish_reason, completion_id, int(time.time()))
89
finally:
90
if hasattr(response, 'aclose'):
91
await safe_aclose(response)
92
93
async def iter_append_model_and_provider(response: AsyncIterator) -> AsyncIterator:
94
last_provider = None
95
try:
96
async for chunk in response:
97
last_provider = get_last_provider(True) if last_provider is None else last_provider
98
chunk.model = last_provider.get("model")
99
chunk.provider = last_provider.get("name")
100
yield chunk
101
finally:
102
if hasattr(response, 'aclose'):
103
await safe_aclose(response)
104
105
class AsyncClient(BaseClient):
106
def __init__(
107
self,
108
provider: ProviderType = None,
109
image_provider: ImageProvider = None,
110
**kwargs
111
) -> None:
112
super().__init__(**kwargs)
113
self.chat: Chat = Chat(self, provider)
114
self._images: Images = Images(self, image_provider)
115
116
@property
117
def images(self) -> Images:
118
return self._images
119
120
class Completions:
121
def __init__(self, client: 'AsyncClient', provider: ProviderType = None):
122
self.client: 'AsyncClient' = client
123
self.provider: ProviderType = provider
124
125
async def create(
126
self,
127
messages: Messages,
128
model: str,
129
provider: ProviderType = None,
130
stream: bool = False,
131
proxy: str = None,
132
response_format: dict = None,
133
max_tokens: int = None,
134
stop: Union[list[str], str] = None,
135
api_key: str = None,
136
ignored: list[str] = None,
137
ignore_working: bool = False,
138
ignore_stream: bool = False,
139
**kwargs
140
) -> Union[ChatCompletion, AsyncIterator[ChatCompletionChunk]]:
141
model, provider = get_model_and_provider(
142
model,
143
self.provider if provider is None else provider,
144
stream,
145
ignored,
146
ignore_working,
147
ignore_stream,
148
)
149
150
stop = [stop] if isinstance(stop, str) else stop
151
152
response = provider.create_completion(
153
model,
154
messages,
155
stream=stream,
156
**filter_none(
157
proxy=self.client.get_proxy() if proxy is None else proxy,
158
max_tokens=max_tokens,
159
stop=stop,
160
api_key=self.client.api_key if api_key is None else api_key
161
),
162
**kwargs
163
)
164
165
if isinstance(response, AsyncIterator):
166
response = iter_response(response, stream, response_format, max_tokens, stop)
167
response = iter_append_model_and_provider(response)
168
return response if stream else await anext(response)
169
else:
170
response = cast_iter_async(response)
171
response = iter_response(response, stream, response_format, max_tokens, stop)
172
response = iter_append_model_and_provider(response)
173
return response if stream else await anext(response)
174
175
class Chat:
176
completions: Completions
177
178
def __init__(self, client: AsyncClient, provider: ProviderType = None):
179
self.completions = Completions(client, provider)
180
181
async def iter_image_response(response: AsyncIterator) -> Union[ImagesResponse, None]:
182
logging.info("Starting iter_image_response")
183
try:
184
async for chunk in response:
185
logging.info(f"Processing chunk: {chunk}")
186
if isinstance(chunk, ImageProviderResponse):
187
logging.info("Found ImageProviderResponse")
188
return ImagesResponse([Image(image) for image in chunk.get_list()])
189
190
logging.warning("No ImageProviderResponse found in the response")
191
return None
192
finally:
193
if hasattr(response, 'aclose'):
194
await safe_aclose(response)
195
196
async def create_image(client: AsyncClient, provider: ProviderType, prompt: str, model: str = "", **kwargs) -> AsyncIterator:
197
logging.info(f"Creating image with provider: {provider}, model: {model}, prompt: {prompt}")
198
199
if isinstance(provider, type) and provider.__name__ == "You":
200
kwargs["chat_mode"] = "create"
201
else:
202
prompt = f"create an image with: {prompt}"
203
204
response = await provider.create_completion(
205
model,
206
[{"role": "user", "content": prompt}],
207
stream=True,
208
proxy=client.get_proxy(),
209
**kwargs
210
)
211
212
logging.info(f"Response from create_completion: {response}")
213
return response
214
215
class Images:
216
def __init__(self, client: 'AsyncClient', provider: ImageProvider = None):
217
self.client: 'AsyncClient' = client
218
self.provider: ImageProvider = provider
219
self.models: ImageModels = ImageModels(client)
220
221
async def generate(self, prompt: str, model: str = None, **kwargs) -> ImagesResponse:
222
logging.info(f"Starting asynchronous image generation for model: {model}, prompt: {prompt}")
223
provider = self.models.get(model, self.provider)
224
if provider is None:
225
raise ValueError(f"Unknown model: {model}")
226
227
logging.info(f"Provider: {provider}")
228
229
if isinstance(provider, IterListProvider):
230
if provider.providers:
231
provider = provider.providers[0]
232
logging.info(f"Using first provider from IterListProvider: {provider}")
233
else:
234
raise ValueError(f"IterListProvider for model {model} has no providers")
235
236
if isinstance(provider, type) and issubclass(provider, AsyncGeneratorProvider):
237
logging.info("Using AsyncGeneratorProvider")
238
messages = [{"role": "user", "content": prompt}]
239
generator = None
240
try:
241
generator = provider.create_async_generator(model, messages, **kwargs)
242
async for response in generator:
243
logging.debug(f"Received response: {type(response)}")
244
if isinstance(response, ImageResponse):
245
return self._process_image_response(response)
246
elif isinstance(response, str):
247
image_response = ImageResponse([response], prompt)
248
return self._process_image_response(image_response)
249
except RuntimeError as e:
250
if "async generator ignored GeneratorExit" in str(e):
251
logging.warning("Generator ignored GeneratorExit, handling gracefully")
252
else:
253
raise
254
finally:
255
if generator and hasattr(generator, 'aclose'):
256
await safe_aclose(generator)
257
logging.info("AsyncGeneratorProvider processing completed")
258
elif hasattr(provider, 'create'):
259
logging.info("Using provider's create method")
260
async_create = asyncio.iscoroutinefunction(provider.create)
261
if async_create:
262
response = await provider.create(prompt)
263
else:
264
response = provider.create(prompt)
265
266
if isinstance(response, ImageResponse):
267
return self._process_image_response(response)
268
elif isinstance(response, str):
269
image_response = ImageResponse([response], prompt)
270
return self._process_image_response(image_response)
271
elif hasattr(provider, 'create_completion'):
272
logging.info("Using provider's create_completion method")
273
response = await create_image(self.client, provider, prompt, model, **kwargs)
274
async for chunk in response:
275
if isinstance(chunk, ImageProviderResponse):
276
logging.info("Found ImageProviderResponse")
277
return ImagesResponse([Image(image) for image in chunk.get_list()])
278
else:
279
raise ValueError(f"Provider {provider} does not support image generation")
280
281
logging.error(f"Unexpected response type: {type(response)}")
282
raise NoImageResponseError(f"Unexpected response type: {type(response)}")
283
284
def _process_image_response(self, response: ImageResponse) -> ImagesResponse:
285
processed_images = []
286
for image_data in response.get_list():
287
if image_data.startswith('http://') or image_data.startswith('https://'):
288
processed_images.append(Image(url=image_data))
289
else:
290
image = to_image(image_data)
291
file_name = self._save_image(image)
292
processed_images.append(Image(url=file_name))
293
return ImagesResponse(processed_images)
294
295
def _save_image(self, image: 'PILImage') -> str:
296
os.makedirs('generated_images', exist_ok=True)
297
file_name = f"generated_images/image_{int(time.time())}.png"
298
image.save(file_name)
299
return file_name
300
301
async def create_variation(self, image: Union[str, bytes], model: str = None, **kwargs) -> ImagesResponse:
302
provider = self.models.get(model, self.provider)
303
if provider is None:
304
raise ValueError(f"Unknown model: {model}")
305
306
if isinstance(provider, type) and issubclass(provider, AsyncGeneratorProvider):
307
messages = [{"role": "user", "content": "create a variation of this image"}]
308
image_data = to_data_uri(image)
309
generator = None
310
try:
311
generator = provider.create_async_generator(model, messages, image=image_data, **kwargs)
312
async for response in generator:
313
if isinstance(response, ImageResponse):
314
return self._process_image_response(response)
315
elif isinstance(response, str):
316
image_response = ImageResponse([response], "Image variation")
317
return self._process_image_response(image_response)
318
except RuntimeError as e:
319
if "async generator ignored GeneratorExit" in str(e):
320
logging.warning("Generator ignored GeneratorExit in create_variation, handling gracefully")
321
else:
322
raise
323
finally:
324
if generator and hasattr(generator, 'aclose'):
325
await safe_aclose(generator)
326
logging.info("AsyncGeneratorProvider processing completed in create_variation")
327
elif hasattr(provider, 'create_variation'):
328
if asyncio.iscoroutinefunction(provider.create_variation):
329
response = await provider.create_variation(image, **kwargs)
330
else:
331
response = provider.create_variation(image, **kwargs)
332
333
if isinstance(response, ImageResponse):
334
return self._process_image_response(response)
335
elif isinstance(response, str):
336
image_response = ImageResponse([response], "Image variation")
337
return self._process_image_response(image_response)
338
else:
339
raise ValueError(f"Provider {provider} does not support image variation")
@@ -4,12 +4,16 @@ import os
4
4
import time
5
5
import random
6
6
import string
7
import logging
7
import threading
8
8
import asyncio
9
from typing import Union
9
import base64
10
import aiohttp
11
import queue
12
from typing import Union, AsyncIterator, Iterator
13
10
14
from ..providers.base_provider import AsyncGeneratorProvider
11
15
from ..image import ImageResponse, to_image, to_data_uri
12
from ..typing import Union, Iterator, Messages, ImageType
16
from ..typing import Messages, ImageType
13
17
from ..providers.types import BaseProvider, ProviderType, FinishReason
14
18
from ..providers.conversation import BaseConversation
15
19
from ..image import ImageResponse as ImageProviderResponse
@@ -23,44 +27,83 @@ from .helper import find_stop, filter_json, filter_none
23
27
from ..models import ModelUtils
24
28
from ..Provider import IterListProvider
25
29
30
# Helper function to convert an async generator to a synchronous iterator
31
def to_sync_iter(async_gen: AsyncIterator) -> Iterator:
32
q = queue.Queue()
33
loop = asyncio.new_event_loop()
34
done = object()
35
36
def _run():
37
asyncio.set_event_loop(loop)
38
39
async def iterate():
40
try:
41
async for item in async_gen:
42
q.put(item)
43
finally:
44
q.put(done)
45
46
loop.run_until_complete(iterate())
47
loop.close()
48
49
threading.Thread(target=_run).start()
26
50
51
while True:
52
item = q.get()
53
if item is done:
54
break
55
yield item
56
57
# Helper function to convert a synchronous iterator to an async iterator
58
async def to_async_iterator(iterator):
59
for item in iterator:
60
yield item
61
62
# Synchronous iter_response function
27
63
def iter_response(
28
response: Iterator[str],
64
response: Union[Iterator[str], AsyncIterator[str]],
29
65
stream: bool,
30
66
response_format: dict = None,
31
67
max_tokens: int = None,
32
68
stop: list = None
33
) -> IterResponse:
69
) -> Iterator[Union[ChatCompletion, ChatCompletionChunk]]:
34
70
content = ""
35
71
finish_reason = None
36
72
completion_id = ''.join(random.choices(string.ascii_letters + string.digits, k=28))
37
38
for idx, chunk in enumerate(response):
73
idx = 0
74
75
if hasattr(response, '__aiter__'):
76
# It's an async iterator, wrap it into a sync iterator
77
response = to_sync_iter(response)
78
79
for chunk in response:
39
80
if isinstance(chunk, FinishReason):
40
81
finish_reason = chunk.reason
41
82
break
42
83
elif isinstance(chunk, BaseConversation):
43
84
yield chunk
44
85
continue
45
86
46
87
content += str(chunk)
47
88
48
89
if max_tokens is not None and idx + 1 >= max_tokens:
49
90
finish_reason = "length"
50
91
51
92
first, content, chunk = find_stop(stop, content, chunk if stream else None)
52
93
53
94
if first != -1:
54
95
finish_reason = "stop"
55
96
56
97
if stream:
57
98
yield ChatCompletionChunk(chunk, None, completion_id, int(time.time()))
58
99
59
100
if finish_reason is not None:
60
101
break
61
102
103
idx += 1
104
62
105
finish_reason = "stop" if finish_reason is None else finish_reason
63
106
64
107
if stream:
65
108
yield ChatCompletionChunk(None, finish_reason, completion_id, int(time.time()))
66
109
else:
@@ -69,16 +112,16 @@ def iter_response(
69
112
content = filter_json(content)
70
113
yield ChatCompletion(content, finish_reason, completion_id, int(time.time()))
71
114
72
73
def iter_append_model_and_provider(response: IterResponse) -> IterResponse:
115
# Synchronous iter_append_model_and_provider function
116
def iter_append_model_and_provider(response: Iterator) -> Iterator:
74
117
last_provider = None
118
75
119
for chunk in response:
76
120
last_provider = get_last_provider(True) if last_provider is None else last_provider
77
121
chunk.model = last_provider.get("model")
78
122
chunk.provider = last_provider.get("name")
79
123
yield chunk
80
124
81
82
125
class Client(BaseClient):
83
126
def __init__(
84
127
self,
@@ -97,7 +140,6 @@ class Client(BaseClient):
97
140
async def async_images(self) -> Images:
98
141
return self._images
99
142
100
101
143
class Completions:
102
144
def __init__(self, client: Client, provider: ProviderType = None):
103
145
self.client: Client = client
@@ -129,25 +171,115 @@ class Completions:
129
171
)
130
172
131
173
stop = [stop] if isinstance(stop, str) else stop
132
133
response = provider.create_completion(
174
175
if asyncio.iscoroutinefunction(provider.create_completion):
176
# Run the asynchronous function in an event loop
177
response = asyncio.run(provider.create_completion(
178
model,
179
messages,
180
stream=stream,
181
**filter_none(
182
proxy=self.client.get_proxy() if proxy is None else proxy,
183
max_tokens=max_tokens,
184
stop=stop,
185
api_key=self.client.api_key if api_key is None else api_key
186
),
187
**kwargs
188
))
189
else:
190
response = provider.create_completion(
191
model,
192
messages,
193
stream=stream,
194
**filter_none(
195
proxy=self.client.get_proxy() if proxy is None else proxy,
196
max_tokens=max_tokens,
197
stop=stop,
198
api_key=self.client.api_key if api_key is None else api_key
199
),
200
**kwargs
201
)
202
203
if stream:
204
if hasattr(response, '__aiter__'):
205
# It's an async generator, wrap it into a sync iterator
206
response = to_sync_iter(response)
207
208
# Now 'response' is an iterator
209
response = iter_response(response, stream, response_format, max_tokens, stop)
210
response = iter_append_model_and_provider(response)
211
return response
212
else:
213
if hasattr(response, '__aiter__'):
214
# If response is an async generator, collect it into a list
215
response = list(to_sync_iter(response))
216
response = iter_response(response, stream, response_format, max_tokens, stop)
217
response = iter_append_model_and_provider(response)
218
return next(response)
219
220
async def async_create(
221
self,
222
messages: Messages,
223
model: str,
224
provider: ProviderType = None,
225
stream: bool = False,
226
proxy: str = None,
227
response_format: dict = None,
228
max_tokens: int = None,
229
stop: Union[list[str], str] = None,
230
api_key: str = None,
231
ignored: list[str] = None,
232
ignore_working: bool = False,
233
ignore_stream: bool = False,
234
**kwargs
235
) -> Union[ChatCompletion, AsyncIterator[ChatCompletionChunk]]:
236
model, provider = get_model_and_provider(
134
237
model,
135
messages,
136
stream=stream,
137
**filter_none(
138
proxy=self.client.get_proxy() if proxy is None else proxy,
139
max_tokens=max_tokens,
140
stop=stop,
141
api_key=self.client.api_key if api_key is None else api_key
142
),
143
**kwargs
238
self.provider if provider is None else provider,
239
stream,
240
ignored,
241
ignore_working,
242
ignore_stream,
144
243
)
145
146
response = iter_response(response, stream, response_format, max_tokens, stop)
147
response = iter_append_model_and_provider(response)
148
149
return response if stream else next(response)
150
244
245
stop = [stop] if isinstance(stop, str) else stop
246
247
if asyncio.iscoroutinefunction(provider.create_completion):
248
response = await provider.create_completion(
249
model,
250
messages,
251
stream=stream,
252
**filter_none(
253
proxy=self.client.get_proxy() if proxy is None else proxy,
254
max_tokens=max_tokens,
255
stop=stop,
256
api_key=self.client.api_key if api_key is None else api_key
257
),
258
**kwargs
259
)
260
else:
261
response = provider.create_completion(
262
model,
263
messages,
264
stream=stream,
265
**filter_none(
266
proxy=self.client.get_proxy() if proxy is None else proxy,
267
max_tokens=max_tokens,
268
stop=stop,
269
api_key=self.client.api_key if api_key is None else api_key
270
),
271
**kwargs
272
)
273
274
# Removed 'await' here since 'async_iter_response' returns an async generator
275
response = async_iter_response(response, stream, response_format, max_tokens, stop)
276
response = async_iter_append_model_and_provider(response)
277
278
if stream:
279
return response
280
else:
281
async for result in response:
282
return result
151
283
152
284
class Chat:
153
285
completions: Completions
@@ -155,153 +287,224 @@ class Chat:
155
287
def __init__(self, client: Client, provider: ProviderType = None):
156
288
self.completions = Completions(client, provider)
157
289
290
# Asynchronous versions of the helper functions
291
async def async_iter_response(
292
response: Union[AsyncIterator[str], Iterator[str]],
293
stream: bool,
294
response_format: dict = None,
295
max_tokens: int = None,
296
stop: list = None
297
) -> AsyncIterator[Union[ChatCompletion, ChatCompletionChunk]]:
298
content = ""
299
finish_reason = None
300
completion_id = ''.join(random.choices(string.ascii_letters + string.digits, k=28))
301
idx = 0
302
303
if not hasattr(response, '__aiter__'):
304
response = to_async_iterator(response)
305
306
async for chunk in response:
307
if isinstance(chunk, FinishReason):
308
finish_reason = chunk.reason
309
break
310
elif isinstance(chunk, BaseConversation):
311
yield chunk
312
continue
313
314
content += str(chunk)
315
316
if max_tokens is not None and idx + 1 >= max_tokens:
317
finish_reason = "length"
318
319
first, content, chunk = find_stop(stop, content, chunk if stream else None)
158
320
159
def iter_image_response(response: Iterator) -> Union[ImagesResponse, None]:
160
logging.info("Starting iter_image_response")
161
response_list = list(response)
162
logging.info(f"Response list: {response_list}")
163
164
for chunk in response_list:
165
logging.info(f"Processing chunk: {chunk}")
321
if first != -1:
322
finish_reason = "stop"
323
324
if stream:
325
yield ChatCompletionChunk(chunk, None, completion_id, int(time.time()))
326
327
if finish_reason is not None:
328
break
329
330
idx += 1
331
332
finish_reason = "stop" if finish_reason is None else finish_reason
333
334
if stream:
335
yield ChatCompletionChunk(None, finish_reason, completion_id, int(time.time()))
336
else:
337
if response_format is not None and "type" in response_format:
338
if response_format["type"] == "json_object":
339
content = filter_json(content)
340
yield ChatCompletion(content, finish_reason, completion_id, int(time.time()))
341
342
async def async_iter_append_model_and_provider(response: AsyncIterator) -> AsyncIterator:
343
last_provider = None
344
345
if not hasattr(response, '__aiter__'):
346
response = to_async_iterator(response)
347
348
async for chunk in response:
349
last_provider = get_last_provider(True) if last_provider is None else last_provider
350
chunk.model = last_provider.get("model")
351
chunk.provider = last_provider.get("name")
352
yield chunk
353
354
async def iter_image_response(response: AsyncIterator) -> Union[ImagesResponse, None]:
355
response_list = []
356
async for chunk in response:
166
357
if isinstance(chunk, ImageProviderResponse):
167
logging.info("Found ImageProviderResponse")
168
return ImagesResponse([Image(image) for image in chunk.get_list()])
169
170
logging.warning("No ImageProviderResponse found in the response")
171
return None
358
response_list.extend(chunk.get_list())
359
elif isinstance(chunk, str):
360
response_list.append(chunk)
172
361
362
if response_list:
363
return ImagesResponse([Image(image) for image in response_list])
173
364
174
def create_image(client: Client, provider: ProviderType, prompt: str, model: str = "", **kwargs) -> Iterator:
175
logging.info(f"Creating image with provider: {provider}, model: {model}, prompt: {prompt}")
176
365
return None
366
367
async def create_image(client: Client, provider: ProviderType, prompt: str, model: str = "", **kwargs) -> AsyncIterator:
177
368
if isinstance(provider, type) and provider.__name__ == "You":
178
369
kwargs["chat_mode"] = "create"
179
370
else:
180
371
prompt = f"create an image with: {prompt}"
181
182
response = provider.create_completion(
183
model,
184
[{"role": "user", "content": prompt}],
185
stream=True,
186
proxy=client.get_proxy(),
187
**kwargs
188
)
189
190
logging.info(f"Response from create_completion: {response}")
372
373
if asyncio.iscoroutinefunction(provider.create_completion):
374
response = await provider.create_completion(
375
model,
376
[{"role": "user", "content": prompt}],
377
stream=True,
378
proxy=client.get_proxy(),
379
**kwargs
380
)
381
else:
382
response = provider.create_completion(
383
model,
384
[{"role": "user", "content": prompt}],
385
stream=True,
386
proxy=client.get_proxy(),
387
**kwargs
388
)
389
390
# Wrap synchronous iterator into async iterator if necessary
391
if not hasattr(response, '__aiter__'):
392
response = to_async_iterator(response)
393
191
394
return response
192
395
396
class Image:
397
def __init__(self, url: str = None, b64_json: str = None):
398
self.url = url
399
self.b64_json = b64_json
400
401
def __repr__(self):
402
return f"Image(url={self.url}, b64_json={'<base64 data>' if self.b64_json else None})"
403
404
class ImagesResponse:
405
def __init__(self, data: list[Image]):
406
self.data = data
407
408
def __repr__(self):
409
return f"ImagesResponse(data={self.data})"
193
410
194
411
class Images:
195
def __init__(self, client: 'Client', provider: ImageProvider = None):
412
def __init__(self, client: 'Client', provider: 'ImageProvider' = None):
196
413
self.client: 'Client' = client
197
self.provider: ImageProvider = provider
414
self.provider: 'ImageProvider' = provider
198
415
self.models: ImageModels = ImageModels(client)
199
416
200
def generate(self, prompt: str, model: str = None, **kwargs) -> ImagesResponse:
201
logging.info(f"Starting synchronous image generation for model: {model}, prompt: {prompt}")
202
try:
203
loop = asyncio.get_event_loop()
204
except RuntimeError:
205
loop = asyncio.new_event_loop()
206
asyncio.set_event_loop(loop)
207
208
try:
209
result = loop.run_until_complete(self.async_generate(prompt, model, **kwargs))
210
logging.info(f"Synchronous image generation completed. Result: {result}")
211
return result
212
except Exception as e:
213
logging.error(f"Error in synchronous image generation: {str(e)}")
214
raise
215
finally:
216
if loop.is_running():
217
loop.close()
218
219
async def async_generate(self, prompt: str, model: str = None, **kwargs) -> ImagesResponse:
220
logging.info(f"Generating image for model: {model}, prompt: {prompt}")
417
def generate(self, prompt: str, model: str = None, response_format: str = "url", **kwargs) -> ImagesResponse:
418
"""
419
Synchronous generate method that runs the async_generate method in an event loop.
420
"""
421
return asyncio.run(self.async_generate(prompt, model, response_format=response_format, **kwargs))
422
423
async def async_generate(self, prompt: str, model: str = None, response_format: str = "url", **kwargs) -> ImagesResponse:
221
424
provider = self.models.get(model, self.provider)
222
425
if provider is None:
223
426
raise ValueError(f"Unknown model: {model}")
224
225
logging.info(f"Provider: {provider}")
226
427
227
428
if isinstance(provider, IterListProvider):
228
429
if provider.providers:
229
430
provider = provider.providers[0]
230
logging.info(f"Using first provider from IterListProvider: {provider}")
231
431
else:
232
432
raise ValueError(f"IterListProvider for model {model} has no providers")
233
433
234
434
if isinstance(provider, type) and issubclass(provider, AsyncGeneratorProvider):
235
logging.info("Using AsyncGeneratorProvider")
236
435
messages = [{"role": "user", "content": prompt}]
237
436
async for response in provider.create_async_generator(model, messages, **kwargs):
238
437
if isinstance(response, ImageResponse):
239
return self._process_image_response(response)
438
return await self._process_image_response(response, response_format)
240
439
elif isinstance(response, str):
241
440
image_response = ImageResponse([response], prompt)
242
return self._process_image_response(image_response)
441
return await self._process_image_response(image_response, response_format)
243
442
elif hasattr(provider, 'create'):
244
logging.info("Using provider's create method")
245
443
if asyncio.iscoroutinefunction(provider.create):
246
444
response = await provider.create(prompt)
247
445
else:
248
446
response = provider.create(prompt)
249
447
250
448
if isinstance(response, ImageResponse):
251
return self._process_image_response(response)
449
return await self._process_image_response(response, response_format)
252
450
elif isinstance(response, str):
253
451
image_response = ImageResponse([response], prompt)
254
return self._process_image_response(image_response)
452
return await self._process_image_response(image_response, response_format)
255
453
else:
256
454
raise ValueError(f"Provider {provider} does not support image generation")
257
258
logging.error(f"Unexpected response type: {type(response)}")
455
259
456
raise NoImageResponseError(f"Unexpected response type: {type(response)}")
260
457
261
def _process_image_response(self, response: ImageResponse) -> ImagesResponse:
458
async def _process_image_response(self, response: ImageResponse, response_format: str) -> ImagesResponse:
262
459
processed_images = []
460
263
461
for image_data in response.get_list():
264
462
if image_data.startswith('http://') or image_data.startswith('https://'):
265
processed_images.append(Image(url=image_data))
463
if response_format == "url":
464
processed_images.append(Image(url=image_data))
465
elif response_format == "b64_json":
466
# Fetch the image data and convert it to base64
467
image_content = await self._fetch_image(image_data)
468
b64_json = base64.b64encode(image_content).decode('utf-8')
469
processed_images.append(Image(b64_json=b64_json))
266
470
else:
267
image = to_image(image_data)
268
file_name = self._save_image(image)
269
processed_images.append(Image(url=file_name))
471
# Assume image_data is base64 data or binary
472
if response_format == "url":
473
if image_data.startswith('data:image'):
474
# Remove the data URL scheme and get the base64 data
475
header, base64_data = image_data.split(',', 1)
476
else:
477
base64_data = image_data
478
# Decode the base64 data
479
image_data_bytes = base64.b64decode(base64_data)
480
# Convert bytes to an image
481
image = to_image(image_data_bytes)
482
file_name = self._save_image(image)
483
processed_images.append(Image(url=file_name))
484
elif response_format == "b64_json":
485
if isinstance(image_data, bytes):
486
b64_json = base64.b64encode(image_data).decode('utf-8')
487
else:
488
b64_json = image_data # If already base64-encoded string
489
processed_images.append(Image(b64_json=b64_json))
490
270
491
return ImagesResponse(processed_images)
271
492
493
async def _fetch_image(self, url: str) -> bytes:
494
# Asynchronously fetch image data from the URL
495
async with aiohttp.ClientSession() as session:
496
async with session.get(url) as resp:
497
if resp.status == 200:
498
return await resp.read()
499
else:
500
raise Exception(f"Failed to fetch image from {url}, status code {resp.status}")
501
272
502
def _save_image(self, image: 'PILImage') -> str:
273
503
os.makedirs('generated_images', exist_ok=True)
274
file_name = f"generated_images/image_{int(time.time())}.png"
504
file_name = f"generated_images/image_{int(time.time())}_{random.randint(0, 10000)}.png"
275
505
image.save(file_name)
276
506
return file_name
277
507
278
async def create_variation(self, image: Union[str, bytes], model: str = None, **kwargs):
279
provider = self.models.get(model, self.provider)
280
if provider is None:
281
raise ValueError(f"Unknown model: {model}")
282
283
if isinstance(provider, type) and issubclass(provider, AsyncGeneratorProvider):
284
messages = [{"role": "user", "content": "create a variation of this image"}]
285
image_data = to_data_uri(image)
286
async for response in provider.create_async_generator(model, messages, image=image_data, **kwargs):
287
if isinstance(response, ImageResponse):
288
return self._process_image_response(response)
289
elif isinstance(response, str):
290
image_response = ImageResponse([response], "Image variation")
291
return self._process_image_response(image_response)
292
elif hasattr(provider, 'create_variation'):
293
if asyncio.iscoroutinefunction(provider.create_variation):
294
response = await provider.create_variation(image, **kwargs)
295
else:
296
response = provider.create_variation(image, **kwargs)
297
298
if isinstance(response, ImageResponse):
299
return self._process_image_response(response)
300
elif isinstance(response, str):
301
image_response = ImageResponse([response], "Image variation")
302
return self._process_image_response(image_response)
303
else:
304
raise ValueError(f"Provider {provider} does not support image variation")
305
306
raise NoImageResponseError("Failed to create image variation")
307
508
async def create_variation(self, image: Union[str, bytes], model: str = None, response_format: str = "url", **kwargs):
509
# Existing implementation, adjust if you want to support b64_json here as well
510
pass