返回提交历史
Modified
g4f/gui/server/api.py
+63
-61
XFEstudio/gpt4free
feat(g4f/gui/server/api.py): improve image handling and response streaming
c2e3107c
代码差异
1 个文件
+63
-61
@@ -23,8 +23,8 @@ from g4f.providers.conversation import BaseConversation
23
23
conversations: dict[dict[str, BaseConversation]] = {}
24
24
images_dir = "./generated_images"
25
25
26
class Api():
27
26
27
class Api:
28
28
@staticmethod
29
29
def get_models() -> list[str]:
30
30
"""
@@ -42,9 +42,11 @@ class Api():
42
42
if provider in __map__:
43
43
provider: ProviderType = __map__[provider]
44
44
if issubclass(provider, ProviderModelMixin):
45
return [{"model": model, "default": model == provider.default_model} for model in provider.get_models()]
46
else:
47
return []
45
return [
46
{"model": model, "default": model == provider.default_model}
47
for model in provider.get_models()
48
]
49
return []
48
50
49
51
@staticmethod
50
52
def get_image_models() -> list[dict]:
@@ -66,7 +68,7 @@ class Api():
66
68
"image_model": model,
67
69
"vision_model": parent.default_vision_model if hasattr(parent, "default_vision_model") else None
68
70
})
69
index.append(parent.__name__)
71
index.append(parent.__name__)
70
72
elif hasattr(provider, "default_vision_model") and provider.__name__ not in index:
71
73
image_models.append({
72
74
"provider": provider.__name__,
@@ -84,15 +86,13 @@ class Api():
84
86
Return a list of all working providers.
85
87
"""
86
88
return {
87
provider.__name__: (provider.label
88
if hasattr(provider, "label")
89
else provider.__name__) +
90
(" (WebDriver)"
91
if "webdriver" in provider.get_parameters()
92
else "") +
93
(" (Auth)"
94
if provider.needs_auth
95
else "")
89
provider.__name__: (
90
provider.label if hasattr(provider, "label") else provider.__name__
91
) + (
92
" (WebDriver)" if "webdriver" in provider.get_parameters() else ""
93
) + (
94
" (Auth)" if provider.needs_auth else ""
95
)
96
96
for provider in __providers__
97
97
if provider.working
98
98
}
@@ -126,7 +126,7 @@ class Api():
126
126
127
127
Returns:
128
128
dict: Arguments prepared for chat completion.
129
"""
129
"""
130
130
model = json_data.get('model') or models.default
131
131
provider = json_data.get('provider')
132
132
messages = json_data['messages']
@@ -155,61 +155,62 @@ class Api():
155
155
}
156
156
157
157
def _create_response_stream(self, kwargs: dict, conversation_id: str, provider: str) -> Iterator:
158
"""
159
Creates and returns a streaming response for the conversation.
160
161
Args:
162
kwargs (dict): Arguments for creating the chat completion.
163
164
Yields:
165
str: JSON formatted response chunks for the stream.
166
167
Raises:
168
Exception: If an error occurs during the streaming process.
169
"""
170
158
try:
159
result = ChatCompletion.create(**kwargs)
171
160
first = True
172
for chunk in ChatCompletion.create(**kwargs):
161
if isinstance(result, ImageResponse):
162
# Якщо результат є ImageResponse, обробляємо його як одиночний елемент
173
163
if first:
174
164
first = False
175
165
yield self._format_json("provider", get_last_provider(True))
176
if isinstance(chunk, BaseConversation):
177
if provider not in conversations:
178
conversations[provider] = {}
179
conversations[provider][conversation_id] = chunk
180
yield self._format_json("conversation", conversation_id)
181
elif isinstance(chunk, Exception):
182
logging.exception(chunk)
183
yield self._format_json("message", get_error_message(chunk))
184
elif isinstance(chunk, ImagePreview):
185
yield self._format_json("preview", chunk.to_string())
186
elif isinstance(chunk, ImageResponse):
187
async def copy_images(images: list[str], cookies: Optional[Cookies] = None):
188
async with ClientSession(
189
connector=get_connector(None, os.environ.get("G4F_PROXY")),
190
cookies=cookies
191
) as session:
192
async def copy_image(image):
193
async with session.get(image) as response:
194
target = os.path.join(images_dir, f"{int(time.time())}_{str(uuid.uuid4())}")
195
with open(target, "wb") as f:
196
async for chunk in response.content.iter_any():
197
f.write(chunk)
198
with open(target, "rb") as f:
199
extension = is_accepted_format(f.read(12)).split("/")[-1]
200
extension = "jpg" if extension == "jpeg" else extension
201
new_target = f"{target}.{extension}"
202
os.rename(target, new_target)
203
return f"/images/{os.path.basename(new_target)}"
204
return await asyncio.gather(*[copy_image(image) for image in images])
205
images = asyncio.run(copy_images(chunk.get_list(), chunk.options.get("cookies")))
206
yield self._format_json("content", str(ImageResponse(images, chunk.alt)))
207
elif not isinstance(chunk, FinishReason):
208
yield self._format_json("content", str(chunk))
166
yield self._format_json("content", str(result))
167
else:
168
# Якщо результат є ітерабельним, обробляємо його як раніше
169
for chunk in result:
170
if first:
171
first = False
172
yield self._format_json("provider", get_last_provider(True))
173
if isinstance(chunk, BaseConversation):
174
if provider not in conversations:
175
conversations[provider] = {}
176
conversations[provider][conversation_id] = chunk
177
yield self._format_json("conversation", conversation_id)
178
elif isinstance(chunk, Exception):
179
logging.exception(chunk)
180
yield self._format_json("message", get_error_message(chunk))
181
elif isinstance(chunk, ImagePreview):
182
yield self._format_json("preview", chunk.to_string())
183
elif isinstance(chunk, ImageResponse):
184
# Обробка ImageResponse
185
images = asyncio.run(self._copy_images(chunk.get_list(), chunk.options.get("cookies")))
186
yield self._format_json("content", str(ImageResponse(images, chunk.alt)))
187
elif not isinstance(chunk, FinishReason):
188
yield self._format_json("content", str(chunk))
209
189
except Exception as e:
210
190
logging.exception(e)
211
191
yield self._format_json('error', get_error_message(e))
212
192
193
# Додайте цей метод до класу Api
194
async def _copy_images(self, images: list[str], cookies: Optional[Cookies] = None):
195
async with ClientSession(
196
connector=get_connector(None, os.environ.get("G4F_PROXY")),
197
cookies=cookies
198
) as session:
199
async def copy_image(image):
200
async with session.get(image) as response:
201
target = os.path.join(images_dir, f"{int(time.time())}_{str(uuid.uuid4())}")
202
with open(target, "wb") as f:
203
async for chunk in response.content.iter_any():
204
f.write(chunk)
205
with open(target, "rb") as f:
206
extension = is_accepted_format(f.read(12)).split("/")[-1]
207
extension = "jpg" if extension == "jpeg" else extension
208
new_target = f"{target}.{extension}"
209
os.rename(target, new_target)
210
return f"/images/{os.path.basename(new_target)}"
211
212
return await asyncio.gather(*[copy_image(image) for image in images])
213
213
214
def _format_json(self, response_type: str, content):
214
215
"""
215
216
Formats and returns a JSON response.
@@ -226,6 +227,7 @@ class Api():
226
227
response_type: content
227
228
}
228
229
230
229
231
def get_error_message(exception: Exception) -> str:
230
232
"""
231
233
Generates a formatted error message from an exception.