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

XFEstudio/gpt4free

feat(g4f/client/async_client.py): enhance image and chat response handling

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

代码差异

1 个文件 +160 -162
Modified g4f/client/async_client.py +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")