XFE Git
XFE Studio Git
Git 首页 全局搜索
XFE 主站 文档 NuGet
公开
关注 0 Fork 0 Star 1
返回提交历史

XFEstudio/gpt4free

feat(g4f/client/async_client.py, g4f/client/async_client.py): enhance async and sync handling in client

0d868f64
kqlio67 <kqlio67@users.noreply.github.com>
提交于

代码差异

3 个文件 +335 -472
Modified g4f/client/__init__.py +0 -1
@@ -1,3 +1,2 @@
1 1 from .stubs import ChatCompletion, ChatCompletionChunk, ImagesResponse
2 2 from .client import Client
3 from .async_client import AsyncClient
Deleted g4f/client/async_client.py +0 -339
@@ -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")
Modified g4f/client/client.py +335 -132
@@ -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