返回提交历史
Modified
g4f/client/async_client.py
+160
-162
XFEstudio/gpt4free
feat(g4f/client/async_client.py): enhance image and chat response handling
b3ddad4a
代码差异
1 个文件
+160
-162
@@ -1,32 +1,27 @@
1
1
from __future__ import annotations
2
2
3
import os
3
4
import time
4
5
import random
5
6
import string
7
import logging
6
8
import asyncio
7
import base64
8
from aiohttp import ClientSession, BaseConnector
9
10
from .types import Client as BaseClient
11
from .types import ProviderType, FinishReason
12
from .stubs import ChatCompletion, ChatCompletionChunk, ImagesResponse, Image
13
from .types import AsyncIterResponse, ImageProvider
14
from .image_models import ImageModels
15
from .helper import filter_json, find_stop, filter_none, cast_iter_async
16
from .service import get_last_provider, get_model_and_provider
17
from ..Provider import ProviderUtils
18
from ..typing import Union, Messages, AsyncIterator, ImageType
19
from ..errors import NoImageResponseError, ProviderNotFoundError
20
from ..requests.aiohttp import get_connector
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
21
14
from ..providers.conversation import BaseConversation
22
from ..image import ImageResponse as ImageProviderResponse, ImageDataResponse
23
24
try:
25
anext
26
except NameError:
27
async def anext(iter):
28
async for chunk in iter:
29
return chunk
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
30
25
31
26
async def iter_response(
32
27
response: AsyncIterator[str],
@@ -34,30 +29,37 @@ async def iter_response(
34
29
response_format: dict = None,
35
30
max_tokens: int = None,
36
31
stop: list = None
37
) -> AsyncIterResponse:
32
) -> AsyncIterator[ChatCompletion | ChatCompletionChunk]:
38
33
content = ""
39
34
finish_reason = None
40
35
completion_id = ''.join(random.choices(string.ascii_letters + string.digits, k=28))
41
count: int = 0
42
async for chunk in response:
36
37
async for idx, chunk in enumerate(response):
43
38
if isinstance(chunk, FinishReason):
44
39
finish_reason = chunk.reason
45
40
break
46
41
elif isinstance(chunk, BaseConversation):
47
42
yield chunk
48
43
continue
44
49
45
content += str(chunk)
50
count += 1
51
if max_tokens is not None and count >= max_tokens:
46
47
if max_tokens is not None and idx + 1 >= max_tokens:
52
48
finish_reason = "length"
53
first, content, chunk = find_stop(stop, content, chunk)
49
50
first, content, chunk = find_stop(stop, content, chunk if stream else None)
51
54
52
if first != -1:
55
53
finish_reason = "stop"
54
56
55
if stream:
57
56
yield ChatCompletionChunk(chunk, None, completion_id, int(time.time()))
57
58
58
if finish_reason is not None:
59
59
break
60
60
61
finish_reason = "stop" if finish_reason is None else finish_reason
62
61
63
if stream:
62
64
yield ChatCompletionChunk(None, finish_reason, completion_id, int(time.time()))
63
65
else:
@@ -66,12 +68,12 @@ async def iter_response(
66
68
content = filter_json(content)
67
69
yield ChatCompletion(content, finish_reason, completion_id, int(time.time()))
68
70
69
async def iter_append_model_and_provider(response: AsyncIterResponse) -> AsyncIterResponse:
71
async def iter_append_model_and_provider(response: AsyncIterator) -> AsyncIterator:
70
72
last_provider = None
71
73
async for chunk in response:
72
74
last_provider = get_last_provider(True) if last_provider is None else last_provider
73
75
chunk.model = last_provider.get("model")
74
chunk.provider = last_provider.get("name")
76
chunk.provider = last_provider.get("name")
75
77
yield chunk
76
78
77
79
class AsyncClient(BaseClient):
@@ -80,59 +82,32 @@ class AsyncClient(BaseClient):
80
82
provider: ProviderType = None,
81
83
image_provider: ImageProvider = None,
82
84
**kwargs
83
):
85
) -> None:
84
86
super().__init__(**kwargs)
85
87
self.chat: Chat = Chat(self, provider)
86
self.images: Images = Images(self, image_provider)
88
self._images: Images = Images(self, image_provider)
87
89
88
def create_response(
89
messages: Messages,
90
model: str,
91
provider: ProviderType = None,
92
stream: bool = False,
93
proxy: str = None,
94
max_tokens: int = None,
95
stop: list[str] = None,
96
api_key: str = None,
97
**kwargs
98
):
99
has_asnyc = hasattr(provider, "create_async_generator")
100
if has_asnyc:
101
create = provider.create_async_generator
102
else:
103
create = provider.create_completion
104
response = create(
105
model, messages,
106
stream=stream,
107
**filter_none(
108
proxy=proxy,
109
max_tokens=max_tokens,
110
stop=stop,
111
api_key=api_key
112
),
113
**kwargs
114
)
115
if not has_asnyc:
116
response = cast_iter_async(response)
117
return response
90
@property
91
def images(self) -> Images:
92
return self._images
118
93
119
class Completions():
94
class Completions:
120
95
def __init__(self, client: AsyncClient, provider: ProviderType = None):
121
96
self.client: AsyncClient = client
122
97
self.provider: ProviderType = provider
123
98
124
def create(
99
async def create(
125
100
self,
126
101
messages: Messages,
127
102
model: str,
128
103
provider: ProviderType = None,
129
104
stream: bool = False,
130
105
proxy: str = None,
106
response_format: dict = None,
131
107
max_tokens: int = None,
132
108
stop: Union[list[str], str] = None,
133
109
api_key: str = None,
134
response_format: dict = None,
135
ignored : list[str] = None,
110
ignored: list[str] = None,
136
111
ignore_working: bool = False,
137
112
ignore_stream: bool = False,
138
113
**kwargs
@@ -143,133 +118,156 @@ class Completions():
143
118
stream,
144
119
ignored,
145
120
ignore_working,
146
ignore_stream
121
ignore_stream,
147
122
)
123
148
124
stop = [stop] if isinstance(stop, str) else stop
149
response = create_response(
150
messages, model,
151
provider, stream,
152
proxy=self.client.get_proxy() if proxy is None else proxy,
153
max_tokens=max_tokens,
154
stop=stop,
155
api_key=self.client.api_key if api_key is None else api_key,
125
126
response = await provider.create_completion(
127
model,
128
messages,
129
stream=stream,
130
**filter_none(
131
proxy=self.client.get_proxy() if proxy is None else proxy,
132
max_tokens=max_tokens,
133
stop=stop,
134
api_key=self.client.api_key if api_key is None else api_key
135
),
156
136
**kwargs
157
137
)
138
158
139
response = iter_response(response, stream, response_format, max_tokens, stop)
159
140
response = iter_append_model_and_provider(response)
160
return response if stream else anext(response)
141
142
return response if stream else await anext(response)
161
143
162
class Chat():
144
class Chat:
163
145
completions: Completions
164
146
165
147
def __init__(self, client: AsyncClient, provider: ProviderType = None):
166
148
self.completions = Completions(client, provider)
167
149
168
async def iter_image_response(
169
response: AsyncIterator,
170
response_format: str = None,
171
connector: BaseConnector = None,
172
proxy: str = None
173
) -> Union[ImagesResponse, None]:
150
async def iter_image_response(response: AsyncIterator) -> Union[ImagesResponse, None]:
151
logging.info("Starting iter_image_response")
174
152
async for chunk in response:
153
logging.info(f"Processing chunk: {chunk}")
175
154
if isinstance(chunk, ImageProviderResponse):
176
if response_format == "b64_json":
177
async with ClientSession(
178
connector=get_connector(connector, proxy),
179
cookies=chunk.options.get("cookies")
180
) as session:
181
async def fetch_image(image):
182
async with session.get(image) as response:
183
return base64.b64encode(await response.content.read()).decode()
184
images = await asyncio.gather(*[fetch_image(image) for image in chunk.get_list()])
185
return ImagesResponse([Image(None, image, chunk.alt) for image in images], int(time.time()))
186
return ImagesResponse([Image(image, None, chunk.alt) for image in chunk.get_list()], int(time.time()))
187
elif isinstance(chunk, ImageDataResponse):
188
return ImagesResponse([Image(None, image, chunk.alt) for image in chunk.get_list()], int(time.time()))
155
logging.info("Found ImageProviderResponse")
156
return ImagesResponse([Image(image) for image in chunk.get_list()])
157
158
logging.warning("No ImageProviderResponse found in the response")
159
return None
160
161
async def create_image(client: AsyncClient, provider: ProviderType, prompt: str, model: str = "", **kwargs) -> AsyncIterator:
162
logging.info(f"Creating image with provider: {provider}, model: {model}, prompt: {prompt}")
189
163
190
def create_image(provider: ProviderType, prompt: str, model: str = "", **kwargs) -> AsyncIterator:
191
164
if isinstance(provider, type) and provider.__name__ == "You":
192
165
kwargs["chat_mode"] = "create"
193
166
else:
194
prompt = f"create a image with: {prompt}"
195
return provider.create_async_generator(
167
prompt = f"create an image with: {prompt}"
168
169
response = await provider.create_completion(
196
170
model,
197
171
[{"role": "user", "content": prompt}],
198
172
stream=True,
173
proxy=client.get_proxy(),
199
174
**kwargs
200
175
)
176
177
logging.info(f"Response from create_completion: {response}")
178
return response
201
179
202
class Images():
180
class Images:
203
181
def __init__(self, client: AsyncClient, provider: ImageProvider = None):
204
182
self.client: AsyncClient = client
205
183
self.provider: ImageProvider = provider
206
184
self.models: ImageModels = ImageModels(client)
207
185
208
def get_provider(self, model: str, provider: ProviderType = None):
209
if isinstance(provider, str):
210
if provider in ProviderUtils.convert:
211
provider = ProviderUtils.convert[provider]
186
async def generate(self, prompt: str, model: str = None, **kwargs) -> ImagesResponse:
187
logging.info(f"Starting asynchronous image generation for model: {model}, prompt: {prompt}")
188
provider = self.models.get(model, self.provider)
189
if provider is None:
190
raise ValueError(f"Unknown model: {model}")
191
192
logging.info(f"Provider: {provider}")
193
194
if isinstance(provider, IterListProvider):
195
if provider.providers:
196
provider = provider.providers[0]
197
logging.info(f"Using first provider from IterListProvider: {provider}")
212
198
else:
213
raise ProviderNotFoundError(f'Provider not found: {provider}')
199
raise ValueError(f"IterListProvider for model {model} has no providers")
200
201
if isinstance(provider, type) and issubclass(provider, AsyncGeneratorProvider):
202
logging.info("Using AsyncGeneratorProvider")
203
messages = [{"role": "user", "content": prompt}]
204
async for response in provider.create_async_generator(model, messages, **kwargs):
205
if isinstance(response, ImageResponse):
206
return self._process_image_response(response)
207
elif isinstance(response, str):
208
image_response = ImageResponse([response], prompt)
209
return self._process_image_response(image_response)
210
elif hasattr(provider, 'create'):
211
logging.info("Using provider's create method")
212
if asyncio.iscoroutinefunction(provider.create):
213
response = await provider.create(prompt)
214
else:
215
response = provider.create(prompt)
216
217
if isinstance(response, ImageResponse):
218
return self._process_image_response(response)
219
elif isinstance(response, str):
220
image_response = ImageResponse([response], prompt)
221
return self._process_image_response(image_response)
214
222
else:
215
provider = self.models.get(model, self.provider)
216
return provider
223
raise ValueError(f"Provider {provider} does not support image generation")
217
224
218
async def generate(
219
self,
220
prompt,
221
model: str = "",
222
provider: ProviderType = None,
223
response_format: str = None,
224
connector: BaseConnector = None,
225
proxy: str = None,
226
**kwargs
227
) -> ImagesResponse:
228
provider = self.get_provider(model, provider)
229
if hasattr(provider, "create_async_generator"):
230
response = create_image(
231
provider,
232
prompt,
233
**filter_none(
234
response_format=response_format,
235
connector=connector,
236
proxy=self.client.get_proxy() if proxy is None else proxy,
237
),
238
**kwargs
239
)
225
logging.error(f"Unexpected response type: {type(response)}")
226
raise NoImageResponseError(f"Unexpected response type: {type(response)}")
227
228
def _process_image_response(self, response: ImageResponse) -> ImagesResponse:
229
processed_images = []
230
for image_data in response.get_list():
231
if image_data.startswith('http://') or image_data.startswith('https://'):
232
processed_images.append(Image(url=image_data))
233
else:
234
image = to_image(image_data)
235
file_name = self._save_image(image)
236
processed_images.append(Image(url=file_name))
237
return ImagesResponse(processed_images)
238
239
def _save_image(self, image: 'PILImage') -> str:
240
os.makedirs('generated_images', exist_ok=True)
241
file_name = f"generated_images/image_{int(time.time())}.png"
242
image.save(file_name)
243
return file_name
244
245
async def create_variation(self, image: Union[str, bytes], model: str = None, **kwargs) -> ImagesResponse:
246
provider = self.models.get(model, self.provider)
247
if provider is None:
248
raise ValueError(f"Unknown model: {model}")
249
250
if isinstance(provider, type) and issubclass(provider, AsyncGeneratorProvider):
251
messages = [{"role": "user", "content": "create a variation of this image"}]
252
image_data = to_data_uri(image)
253
async for response in provider.create_async_generator(model, messages, image=image_data, **kwargs):
254
if isinstance(response, ImageResponse):
255
return self._process_image_response(response)
256
elif isinstance(response, str):
257
image_response = ImageResponse([response], "Image variation")
258
return self._process_image_response(image_response)
259
elif hasattr(provider, 'create_variation'):
260
if asyncio.iscoroutinefunction(provider.create_variation):
261
response = await provider.create_variation(image, **kwargs)
262
else:
263
response = provider.create_variation(image, **kwargs)
264
265
if isinstance(response, ImageResponse):
266
return self._process_image_response(response)
267
elif isinstance(response, str):
268
image_response = ImageResponse([response], "Image variation")
269
return self._process_image_response(image_response)
240
270
else:
241
response = await provider.create_async(prompt)
242
return ImagesResponse([Image(image) for image in response.get_list()])
243
image = await iter_image_response(response, response_format, connector, proxy)
244
if image is None:
245
raise NoImageResponseError()
246
return image
271
raise ValueError(f"Provider {provider} does not support image variation")
247
272
248
async def create_variation(
249
self,
250
image: ImageType,
251
model: str = None,
252
response_format: str = None,
253
connector: BaseConnector = None,
254
proxy: str = None,
255
**kwargs
256
):
257
provider = self.get_provider(model, provider)
258
result = None
259
if hasattr(provider, "create_async_generator"):
260
response = provider.create_async_generator(
261
"",
262
[{"role": "user", "content": "create a image like this"}],
263
stream=True,
264
image=image,
265
**filter_none(
266
response_format=response_format,
267
connector=connector,
268
proxy=self.client.get_proxy() if proxy is None else proxy,
269
),
270
**kwargs
271
)
272
result = iter_image_response(response, response_format, connector, proxy)
273
if result is None:
274
raise NoImageResponseError()
275
return result
273
raise NoImageResponseError("Failed to create image variation")