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

XFEstudio/gpt4free

feat: Refactor extra_body handling and update model error handling

- Changed the default value of `extra_body` from an empty dictionary to `None` in `ImageLabs` and `PollinationsAI` classes. - Added a check to initialize `extra_body` to an empty dictionary if it is `None` in the `ImageLabs` class. - Removed the `extra_image_models` list from the `PollinationsAI` class. - Updated the way image models are combined in the `PollinationsAI` class to avoid duplicates. - Changed the error handling for unsupported models from `ModelNotSupportedError` to `ModelNotFoundError` in multiple classes including `OpenaiChat`, `HuggingFaceAPI`, and `HuggingFaceInference`. - Updated the `save_response_media` function to handle both string and bytes responses. - Adjusted the handling of audio data in the `PollinationsAI` class to ensure proper processing of audio responses.

bf4ed09a
hlohaus <983577+hlohaus@users.noreply.github.com>
提交于

代码差异

20 个文件 +164 -150
Modified g4f/Provider/ImageLabs.py +3 -1
@@ -36,9 +36,11 @@ class ImageLabs(AsyncGeneratorProvider, ProviderModelMixin):
36 36 aspect_ratio: str = "1:1",
37 37 width: int = None,
38 38 height: int = None,
39 extra_body: dict = {},
39 extra_body: dict = None,
40 40 **kwargs
41 41 ) -> AsyncResult:
42 if extra_body is None:
43 extra_body = {}
42 44 extra_body = use_aspect_ratio({
43 45 "width": width,
44 46 "height": height,
Modified g4f/Provider/LegacyLMArena.py +3 -2
@@ -3,7 +3,6 @@ from __future__ import annotations
3 3 import random
4 4 import json
5 5 import uuid
6 import sys
7 6 import asyncio
8 7
9 8
@@ -14,7 +13,7 @@ from ..tools.media import merge_media
14 13 from ..image import to_bytes, is_accepted_format
15 14 from .base_provider import AsyncGeneratorProvider, ProviderModelMixin
16 15 from .helper import get_last_user_message
17 from ..errors import ModelNotFoundError
16 from ..errors import ModelNotFoundError, ResponseError
18 17 from .. import debug
19 18
20 19 class LegacyLMArena(AsyncGeneratorProvider, ProviderModelMixin):
@@ -460,6 +459,8 @@ class LegacyLMArena(AsyncGeneratorProvider, ProviderModelMixin):
460 459 content = data
461 460
462 461 if content:
462 if "**NETWORK ERROR DUE TO HIGH TRAFFIC." in content:
463 raise ResponseError(data)
463 464 # Clean up content
464 465 if isinstance(content, str):
465 466 if content.endswith("▌"):
Modified g4f/Provider/PollinationsAI.py +67 -34
@@ -19,7 +19,7 @@ from ..requests.raise_for_status import raise_for_status
19 19 from ..requests.aiohttp import get_connector
20 20 from ..image.copy_images import save_response_media
21 21 from ..image import use_aspect_ratio
22 from ..providers.response import FinishReason, Usage, ToolCalls, ImageResponse, Reasoning, TitleGeneration, SuggestedFollowups
22 from ..providers.response import FinishReason, Usage, ToolCalls, ImageResponse, Reasoning, TitleGeneration, SuggestedFollowups, ProviderInfo
23 23 from ..tools.media import render_messages
24 24 from ..constants import STATIC_URL
25 25 from .. import debug
@@ -84,7 +84,6 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
84 84 text_models = [default_model, "evil"]
85 85 image_models = [default_image_model]
86 86 audio_models = {default_audio_model: []}
87 extra_image_models = ["flux-pro", "flux-dev", "flux-schnell", "dall-e-3", "turbo"]
88 87 vision_models = [default_vision_model, "gpt-4o-mini", "openai", "openai-large", "openai-reasoning", "searchgpt"]
89 88 _models_loaded = False
90 89 # https://github.com/pollinations/pollinations/blob/master/text.pollinations.ai/generateTextPortkey.js#L15
@@ -133,6 +132,9 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
133 132 ### Image Models ###
134 133 "sdxl-turbo": "turbo",
135 134 "gpt-image": "gptimage",
135 "flux-pro": "flux",
136 "flux-dev": "flux",
137 "flux-schnell": "flux"
136 138 }
137 139
138 140 @classmethod
@@ -164,14 +166,14 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
164 166 new_image_models = []
165 167
166 168 # Combine image models without duplicates
167 all_image_models = [cls.default_image_model] # Start with default model
169 image_models = [cls.default_image_model] # Start with default model
168 170
169 171 # Add extra image models if not already in the list
170 for model in cls.extra_image_models + new_image_models:
171 if model not in all_image_models:
172 all_image_models.append(model)
172 for model in new_image_models:
173 if model not in image_models:
174 image_models.append(model)
173 175
174 cls.image_models = all_image_models
176 cls.image_models = image_models
175 177
176 178 text_response = requests.get("https://text.pollinations.ai/models")
177 179 text_response.raise_for_status()
@@ -194,19 +196,19 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
194 196 cls.vision_models.append(alias)
195 197
196 198 # Create a set of unique text models starting with default model
197 unique_text_models = cls.text_models.copy()
199 text_models = cls.text_models.copy()
198 200
199 201 # Add models from vision_models
200 unique_text_models.extend(cls.vision_models)
202 text_models.extend(cls.vision_models)
201 203
202 204 # Add models from the API response
203 205 for model in models:
204 206 model_name = model.get("name")
205 207 if model_name and "input_modalities" in model and "text" in model["input_modalities"]:
206 unique_text_models.append(model_name)
208 text_models.append(model_name)
207 209
208 210 # Convert to list and update text_models
209 cls.text_models = list(dict.fromkeys(unique_text_models))
211 cls.text_models = list(dict.fromkeys(text_models))
210 212
211 213 cls._models_loaded = True
212 214
@@ -243,10 +245,10 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
243 245 messages: Messages,
244 246 stream: bool = True,
245 247 proxy: str = None,
246 cache: bool = False,
248 cache: bool = None,
247 249 referrer: str = STATIC_URL,
248 250 api_key: str = None,
249 extra_body: dict = {},
251 extra_body: dict = None,
250 252 # Image generation parameters
251 253 prompt: str = None,
252 254 aspect_ratio: str = "1:1",
@@ -268,6 +270,10 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
268 270 extra_parameters: list[str] = ["tools", "parallel_tool_calls", "tool_choice", "reasoning_effort", "logit_bias", "voice", "modalities", "audio"],
269 271 **kwargs
270 272 ) -> AsyncResult:
273 if cache is None:
274 cache = kwargs.get("action") == "next"
275 if extra_body is None:
276 extra_body = {}
271 277 # Load model list
272 278 cls.get_models()
273 279 if not model:
@@ -363,8 +369,11 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
363 369 "safe": str(safe).lower(),
364 370 }, aspect_ratio)
365 371 query = "&".join(f"{k}={quote_plus(str(v))}" for k, v in params.items() if v is not None)
366 prompt = quote_plus(prompt)[:2048-len(cls.image_api_endpoint)-len(query)-8]
367 url = f"{cls.image_api_endpoint}prompt/{prompt}?{query}"
372 encoded_prompt = prompt
373 if model == "gptimage" and aspect_ratio != "1:1":
374 encoded_prompt = f"{encoded_prompt} aspect-ratio: {aspect_ratio}"
375 encoded_prompt = quote_plus(encoded_prompt)[:2048-len(cls.image_api_endpoint)-len(query)-8]
376 url = f"{cls.image_api_endpoint}prompt/{encoded_prompt}?{query}"
368 377 def get_image_url(i: int, seed: Optional[int] = None):
369 378 if i == 0:
370 379 if not cache and seed is None:
@@ -374,10 +383,10 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
374 383 return f"{url}&seed={seed}" if seed else url
375 384 headers = {"referer": referrer}
376 385 if api_key:
377 headers["Authorization"] = f"Bearer {api_key}"
386 headers["authorization"] = f"Bearer {api_key}"
378 387 async with ClientSession(headers=DEFAULT_HEADERS, connector=get_connector(proxy=proxy)) as session:
379 388 responses = set()
380 responses.add(Reasoning(status=f"Generating {n} {'image' if n == 1 else 'images'}..."))
389 responses.add(Reasoning(status=f"Generate {n} {'image' if n == 1 else 'images'}..."))
381 390 finished = 0
382 391 start = time.time()
383 392 async def get_image(responses: set, i: int, seed: Optional[int] = None):
@@ -386,8 +395,11 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
386 395 try:
387 396 await raise_for_status(response)
388 397 except Exception as e:
398 if response.status == 500:
399 responses.add(e)
400 return
389 401 debug.error(f"Error fetching image: {e}")
390 responses.add(ImageResponse(str(response.url), prompt))
402 responses.add(ImageResponse(str(response.url), prompt, {"headers": headers}))
391 403 finished += 1
392 404 responses.add(Reasoning(status=f"Image {finished}/{n} generated in {time.time() - start:.2f}s"))
393 405 tasks = []
@@ -395,7 +407,12 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
395 407 tasks.append(asyncio.create_task(get_image(responses, i, seed)))
396 408 while finished < n or len(responses) > 0:
397 409 while len(responses) > 0:
398 yield responses.pop()
410 item = responses.pop()
411 if isinstance(item, Exception):
412 for task in tasks:
413 task.cancel()
414 raise item
415 yield item
399 416 await asyncio.sleep(0.1)
400 417 await asyncio.gather(*tasks)
401 418
@@ -424,14 +441,17 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
424 441 seed = random.randint(0, 2**32)
425 442
426 443 async with ClientSession(headers=DEFAULT_HEADERS, connector=get_connector(proxy=proxy)) as session:
444 extra_body.update({param: kwargs[param] for param in extra_parameters if param in kwargs})
427 445 if model in cls.audio_models:
428 if "audio" in kwargs and kwargs.get("audio", {}).get("voice") is None:
446 if "audio" in extra_body and extra_body.get("audio", {}).get("voice") is None:
429 447 kwargs["audio"]["voice"] = cls.audio_models[model][0]
430 url = cls.text_api_endpoint
448 elif "audio" not in extra_body:
449 extra_body["audio"] = {"voice": cls.audio_models[model][0]}
450 if extra_body.get("audio", {}).get("format") is None:
451 extra_body["audio"]["format"] = "mp3"
452 if "modalities" not in extra_body:
453 extra_body["modalities"] = ["text", "audio"]
431 454 stream = False
432 else:
433 url = cls.openai_endpoint
434 extra_body.update({param: kwargs[param] for param in extra_parameters if param in kwargs})
435 455 data = filter_none(
436 456 messages=list(render_messages(messages, media)),
437 457 model=model,
@@ -447,19 +467,23 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
447 467 )
448 468 headers = {"referer": referrer}
449 469 if api_key:
450 headers["Authorization"] = f"Bearer {api_key}"
451 async with session.post(url, json=data, headers=headers) as response:
452 if response.status == 400:
453 debug.error(f"Error: 400 - Bad Request: {data}")
470 headers["authorization"] = f"Bearer {api_key}"
471 async with session.post(cls.openai_endpoint, json=data, headers=headers) as response:
472 if response.status in (400, 500):
473 debug.error(f"Error: {response.status} - Bad Request: {data}")
454 474 await raise_for_status(response)
455 475 if response.headers["content-type"].startswith("text/plain"):
456 476 yield await response.text()
457 477 return
458 478 elif response.headers["content-type"].startswith("text/event-stream"):
459 479 reasoning = False
480 model_returned = False
460 481 async for result in see_stream(response.content):
461 482 if "error" in result:
462 483 raise ResponseError(result["error"].get("message", result["error"]))
484 if not model_returned and result.get("model"):
485 yield ProviderInfo(**cls.get_dict(), model=result.get("model"))
486 model_returned = True
463 487 if result.get("usage") is not None:
464 488 yield Usage(**result["usage"])
465 489 choices = result.get("choices", [{}])
@@ -478,15 +502,15 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
478 502 if finish_reason:
479 503 yield FinishReason(finish_reason)
480 504 if reasoning:
481 yield Reasoning(status="Done")
505 yield Reasoning(status="")
482 506 if kwargs.get("action") == "next":
483 507 data = {
484 508 "model": "openai",
485 "messages": messages + FOLLOWUPS_DEVELOPER_MESSAGE,
509 "messages": [m for m in messages if m.get("role") == "user"] + FOLLOWUPS_DEVELOPER_MESSAGE,
486 510 "tool_choice": "required",
487 511 "tools": FOLLOWUPS_TOOLS
488 512 }
489 async with session.post(url, json=data, headers=headers) as response:
513 async with session.post(cls.openai_endpoint, json=data, headers=headers) as response:
490 514 try:
491 515 await raise_for_status(response)
492 516 tool_calls = (await response.json()).get("choices", [{}])[0].get("message", {}).get("tool_calls", [])
@@ -500,7 +524,10 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
500 524 debug.error("Error generating title and followups")
501 525 debug.error(e)
502 526 elif response.headers["content-type"].startswith("application/json"):
527 prompt = format_image_prompt(messages)
503 528 result = await response.json()
529 if result.get("model"):
530 yield ProviderInfo(**cls.get_dict(), model=result.get("model"))
504 531 if "choices" in result:
505 532 choice = result["choices"][0]
506 533 message = choice.get("message", {})
@@ -509,6 +536,13 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
509 536 yield content
510 537 if "tool_calls" in message:
511 538 yield ToolCalls(message["tool_calls"])
539 audio = message.get("audio", {})
540 if "data" in audio:
541 async for chunk in save_response_media(audio["data"], prompt, [model, extra_body.get("audio", {}).get("voice")]):
542 yield chunk
543 if "transcript" in audio:
544 yield "\n\n"
545 yield audio["transcript"]
512 546 else:
513 547 raise ResponseError(result)
514 548 if result.get("usage") is not None:
@@ -517,6 +551,5 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
517 551 if finish_reason:
518 552 yield FinishReason(finish_reason)
519 553 else:
520 async for chunk in save_response_media(response, format_image_prompt(messages), [model, extra_body.get("audio", {}).get("voice")]):
521 yield chunk
522 return
554 async for chunk in save_response_media(response, prompt, [model, extra_body.get("audio", {}).get("voice")]):
555 yield chunk
Modified g4f/Provider/PollinationsImage.py +11 -14
@@ -14,23 +14,20 @@ class PollinationsImage(PollinationsAI):
14 14 default_vision_model = None
15 15 default_image_model = default_model
16 16 audio_models = {}
17 image_models = [default_image_model] # Default models
18 _models_loaded = False # Add a checkbox for synchronization
19 17
20 18 @classmethod
21 19 def get_models(cls, **kwargs):
22 if not cls._models_loaded:
23 # Calling the parent method to load models
24 super().get_models()
25 # Combine models from the parent class and additional ones
26 all_image_models = list(dict.fromkeys(
27 cls.image_models +
28 PollinationsAI.image_models +
29 cls.extra_image_models
30 ))
31 cls.image_models = all_image_models
32 cls._models_loaded = True
33 return cls.image_models
20 PollinationsAI.get_models()
21 cls.image_models = PollinationsAI.image_models
22 cls.models = cls.image_models
23 return cls.models
24
25 @classmethod
26 def get_grouped_models(cls) -> dict[str, list[str]]:
27 PollinationsAI.get_models()
28 return [
29 {"group": "Image Generation", "models": PollinationsAI.image_models},
30 ]
34 31
35 32 @classmethod
36 33 async def create_async_generator(
Modified g4f/Provider/audio/OpenAIFM.py +13 -40
@@ -1,34 +1,27 @@
1 1 from __future__ import annotations
2 2
3 try:
4 has_openaifm = True
5 except ImportError:
6 has_openaifm = False
7
8 3 from aiohttp import ClientSession
9 from urllib.parse import urlencode
10 import json
11 4
12 5 from ...typing import AsyncResult, Messages
13 6 from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin
14 from ..helper import get_last_message
7 from ..helper import get_last_user_message, get_system_prompt
15 8 from ...image.copy_images import save_response_media
16
9 from ...requests.raise_for_status import raise_for_status
10 from ...requests.aiohttp import get_connector
11 from ...requests import DEFAULT_HEADERS
17 12
18 13 class OpenAIFM(AsyncGeneratorProvider, ProviderModelMixin):
19 14 label = "OpenAI.fm"
20 15 url = "https://www.openai.fm"
21 16 api_endpoint = "https://www.openai.fm/api/generate"
22
23 working = has_openaifm
17 working = True
24 18
25 19 default_model = 'gpt-4o-mini-tts'
26 20 default_audio_model = default_model
27 21 default_voice = 'coral'
28 22 voices = ['alloy', 'ash', 'ballad', default_voice, 'echo', 'fable', 'onyx', 'nova', 'sage', 'shimmer', 'verse']
29 23 audio_models = {default_audio_model: voices}
30 models = [default_audio_model]
31
24 models = voices
32 25
33 26 friendly = """Affect/personality: A cheerful guide
34 27
@@ -106,44 +99,24 @@ Emotion: Restrained enthusiasm for discoveries and findings, conveying intellect
106 99 audio: dict = {},
107 100 **kwargs
108 101 ) -> AsyncResult:
109
110 102 # Retrieve parameters from the audio dictionary
111 103 voice = audio.get("voice", kwargs.get("voice", cls.default_voice))
112 instructions = audio.get("instructions", kwargs.get("instructions", cls.friendly))
113
104 instructions = audio.get("instructions", kwargs.get("instructions", get_system_prompt(messages) or cls.friendly))
114 105 headers = {
115 "accept": "*/*",
116 "accept-language": "en-US,en;q=0.9",
117 "cache-control": "no-cache",
118 "pragma": "no-cache",
119 "sec-fetch-dest": "audio",
120 "sec-fetch-mode": "no-cors",
121 "sec-fetch-site": "same-origin",
122 "user-agent": "Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/136.0.0.0 Safari/537.36",
123 "referer": cls.url
106 **DEFAULT_HEADERS,
107 "referer": f"{cls.url}/"
124 108 }
125
126 # Using prompts or formatting messages
127 text = get_last_message(messages, prompt)
128
109 text = get_last_user_message(messages, prompt)
129 110 params = {
130 111 "input": text,
131 112 "prompt": instructions,
132 113 "voice": voice
133 114 }
134
135 async with ClientSession(headers=headers) as session:
136
137 # Print the full URL with parameters
138 full_url = f"{cls.api_endpoint}?{urlencode(params)}"
139
115 async with ClientSession(headers=headers, connector=get_connector(proxy=proxy)) as session:
140 116 async with session.get(
141 117 cls.api_endpoint,
142 params=params,
143 proxy=proxy
118 params=params
144 119 ) as response:
145
146 response.raise_for_status()
147
120 await raise_for_status(response)
148 121 async for chunk in save_response_media(response, text, [model, voice]):
149 122 yield chunk
Modified g4f/Provider/hf_space/Qwen_Qwen_3.py +2 -2
@@ -7,7 +7,7 @@ import uuid
7 7 from ...typing import AsyncResult, Messages
8 8 from ...providers.response import Reasoning, JsonConversation
9 9 from ...requests.raise_for_status import raise_for_status
10 from ...errors import ModelNotSupportedError
10 from ...errors import ModelNotFoundError
11 11 from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin
12 12 from ..helper import get_last_user_message
13 13 from ... import debug
@@ -55,7 +55,7 @@ class Qwen_Qwen_3(AsyncGeneratorProvider, ProviderModelMixin):
55 55 ) -> AsyncResult:
56 56 try:
57 57 model = cls.get_model(model)
58 except ModelNotSupportedError:
58 except ModelNotFoundError:
59 59 pass
60 60 if conversation is None:
61 61 conversation = JsonConversation(session_hash=str(uuid.uuid4()).replace('-', ''))
Modified g4f/Provider/needs_auth/OpenaiChat.py +2 -2
@@ -23,7 +23,7 @@ from ...requests.raise_for_status import raise_for_status
23 23 from ...requests import StreamSession
24 24 from ...requests import get_nodriver
25 25 from ...image import ImageRequest, to_image, to_bytes, is_accepted_format
26 from ...errors import MissingAuthError, NoValidHarFileError, ModelNotSupportedError
26 from ...errors import MissingAuthError, NoValidHarFileError, ModelNotFoundError
27 27 from ...providers.response import JsonConversation, FinishReason, SynthesizeData, AuthResult, ImageResponse, ImagePreview
28 28 from ...providers.response import Sources, TitleGeneration, RequestLogin, Reasoning
29 29 from ...tools.media import merge_media
@@ -358,7 +358,7 @@ class OpenaiChat(AsyncAuthedProvider, ProviderModelMixin):
358 358 debug.error(e)
359 359 try:
360 360 model = cls.get_model(model)
361 except ModelNotSupportedError:
361 except ModelNotFoundError:
362 362 pass
363 363 if conversation is None:
364 364 conversation = Conversation(None, str(uuid.uuid4()), getattr(auth_result, "cookies", {}).get("oai-did"))
Modified g4f/Provider/needs_auth/hf/HuggingFaceAPI.py +4 -4
@@ -5,7 +5,7 @@ import requests
5 5 from ....providers.types import Messages
6 6 from ....typing import MediaListType
7 7 from ....requests import StreamSession, raise_for_status
8 from ....errors import ModelNotSupportedError, PaymentRequiredError
8 from ....errors import ModelNotFoundError, PaymentRequiredError
9 9 from ....providers.response import ProviderInfo
10 10 from ...template.OpenaiTemplate import OpenaiTemplate
11 11 from .models import model_aliases, vision_models, default_llama_model, default_vision_model, text_models
@@ -34,7 +34,7 @@ class HuggingFaceAPI(OpenaiTemplate):
34 34 def get_model(cls, model: str, **kwargs) -> str:
35 35 try:
36 36 return super().get_model(model, **kwargs)
37 except ModelNotSupportedError:
37 except ModelNotFoundError:
38 38 return model
39 39
40 40 @classmethod
@@ -87,14 +87,14 @@ class HuggingFaceAPI(OpenaiTemplate):
87 87 model = cls.get_model(model)
88 88 provider_mapping = await cls.get_mapping(model, api_key)
89 89 if not provider_mapping:
90 raise ModelNotSupportedError(f"Model is not supported: {model} in: {cls.__name__}")
90 raise ModelNotFoundError(f"Model is not supported: {model} in: {cls.__name__}")
91 91 error = None
92 92 for provider_key in provider_mapping:
93 93 api_path = provider_key if provider_key == "novita" else f"{provider_key}/v1"
94 94 api_base = f"https://router.huggingface.co/{api_path}"
95 95 task = provider_mapping[provider_key]["task"]
96 96 if task != "conversational":
97 raise ModelNotSupportedError(f"Model is not supported: {model} in: {cls.__name__} task: {task}")
97 raise ModelNotFoundError(f"Model is not supported: {model} in: {cls.__name__} task: {task}")
98 98 model = provider_mapping[provider_key]["providerId"]
99 99 yield ProviderInfo(**{**cls.get_dict(), "label": f"HuggingFace ({provider_key})"})
100 100 # start = calculate_lenght(messages)
Modified g4f/Provider/needs_auth/hf/HuggingFaceInference.py +10 -8
@@ -7,7 +7,7 @@ import requests
7 7
8 8 from ....typing import AsyncResult, Messages
9 9 from ...base_provider import AsyncGeneratorProvider, ProviderModelMixin, format_prompt
10 from ....errors import ModelNotSupportedError, ResponseError
10 from ....errors import ModelNotFoundError, ResponseError
11 11 from ....requests import StreamSession, raise_for_status
12 12 from ....providers.response import FinishReason, ImageResponse
13 13 from ....image.copy_images import save_response_media
@@ -58,7 +58,7 @@ class HuggingFaceInference(AsyncGeneratorProvider, ProviderModelMixin):
58 58 return cls.model_data[model]
59 59 async with session.get(f"https://huggingface.co/api/models/{model}") as response:
60 60 if response.status == 404:
61 raise ModelNotSupportedError(f"Model is not supported: {model} in: {cls.__name__}")
61 raise ModelNotFoundError(f"Model not found: {model} in: {cls.__name__}")
62 62 await raise_for_status(response)
63 63 cls.model_data[model] = await response.json()
64 64 return cls.model_data[model]
@@ -77,7 +77,7 @@ class HuggingFaceInference(AsyncGeneratorProvider, ProviderModelMixin):
77 77 temperature: float = None,
78 78 prompt: str = None,
79 79 action: str = None,
80 extra_body: dict = {},
80 extra_body: dict = None,
81 81 seed: int = None,
82 82 aspect_ratio: str = None,
83 83 width: int = None,
@@ -86,7 +86,7 @@ class HuggingFaceInference(AsyncGeneratorProvider, ProviderModelMixin):
86 86 ) -> AsyncResult:
87 87 try:
88 88 model = cls.get_model(model)
89 except ModelNotSupportedError:
89 except ModelNotFoundError:
90 90 pass
91 91 headers = {
92 92 'Accept-Encoding': 'gzip, deflate',
@@ -94,6 +94,8 @@ class HuggingFaceInference(AsyncGeneratorProvider, ProviderModelMixin):
94 94 }
95 95 if api_key is not None:
96 96 headers["Authorization"] = f"Bearer {api_key}"
97 if extra_body is None:
98 extra_body = {}
97 99 image_extra_body = use_aspect_ratio({
98 100 "width": width,
99 101 "height": height,
@@ -114,12 +116,12 @@ class HuggingFaceInference(AsyncGeneratorProvider, ProviderModelMixin):
114 116 }
115 117 async with session.post(provider_together_urls[model], json=data) as response:
116 118 if response.status == 404:
117 raise ModelNotSupportedError(f"Model is not supported: {model}")
119 raise ModelNotFoundError(f"Model not found: {model}")
118 120 await raise_for_status(response)
119 121 result = await response.json()
120 122 yield ImageResponse([item["url"] for item in result["data"]], data["prompt"])
121 123 return
122 except ModelNotSupportedError:
124 except ModelNotFoundError:
123 125 pass
124 126 payload = None
125 127 params = {
@@ -156,11 +158,11 @@ class HuggingFaceInference(AsyncGeneratorProvider, ProviderModelMixin):
156 158 params["seed"] = seed
157 159 payload = {"inputs": inputs, "parameters": params, "stream": stream}
158 160 else:
159 raise ModelNotSupportedError(f"Model is not supported: {model} in: {cls.__name__} pipeline_tag: {pipeline_tag}")
161 raise ModelNotFoundError(f"Model is not supported: {model} in: {cls.__name__} pipeline_tag: {pipeline_tag}")
160 162
161 163 async with session.post(f"{api_base.rstrip('/')}/models/{model}", json=payload) as response:
162 164 if response.status == 404:
163 raise ModelNotSupportedError(f"Model is not supported: {model}")
165 raise ModelNotFoundError(f"Model not found: {model}")
164 166 await raise_for_status(response)
165 167 if stream:
166 168 first = True
Modified g4f/Provider/needs_auth/hf/HuggingFaceMedia.py +7 -5
@@ -7,7 +7,7 @@ import requests
7 7
8 8 from ....providers.types import Messages
9 9 from ....requests import StreamSession, raise_for_status
10 from ....errors import ModelNotSupportedError
10 from ....errors import ModelNotFoundError
11 11 from ....providers.helper import format_image_prompt
12 12 from ....providers.base_provider import AsyncGeneratorProvider, ProviderModelMixin
13 13 from ....providers.response import ProviderInfo, ImageResponse, VideoResponse, Reasoning
@@ -98,7 +98,7 @@ class HuggingFaceMedia(AsyncGeneratorProvider, ProviderModelMixin):
98 98 model: str,
99 99 messages: Messages,
100 100 api_key: str = None,
101 extra_body: dict = {},
101 extra_body: dict = None,
102 102 prompt: str = None,
103 103 proxy: str = None,
104 104 timeout: int = 0,
@@ -112,6 +112,8 @@ class HuggingFaceMedia(AsyncGeneratorProvider, ProviderModelMixin):
112 112 resolution: str = "480p",
113 113 **kwargs
114 114 ):
115 if extra_body is None:
116 extra_body = {}
115 117 selected_provider = None
116 118 if model and ":" in model:
117 119 model, selected_provider = model.split(":", 1)
@@ -130,7 +132,7 @@ class HuggingFaceMedia(AsyncGeneratorProvider, ProviderModelMixin):
130 132 }
131 133 provider_mapping = {**new_mapping, **provider_mapping}
132 134 if not provider_mapping:
133 raise ModelNotSupportedError(f"Model is not supported: {model} in: {cls.__name__}")
135 raise ModelNotFoundError(f"Model is not supported: {model} in: {cls.__name__}")
134 136 async def generate(extra_body: dict, aspect_ratio: str = None):
135 137 last_response = None
136 138 for provider_key, provider in provider_mapping.items():
@@ -142,7 +144,7 @@ class HuggingFaceMedia(AsyncGeneratorProvider, ProviderModelMixin):
142 144 task = provider["task"]
143 145 provider_id = provider["providerId"]
144 146 if task not in cls.tasks:
145 raise ModelNotSupportedError(f"Model is not supported: {model} in: {cls.__name__} task: {task}")
147 raise ModelNotFoundError(f"Model is not supported: {model} in: {cls.__name__} task: {task}")
146 148
147 149 if aspect_ratio is None:
148 150 aspect_ratio = "1:1" if task == "text-to-image" else "16:9"
@@ -209,7 +211,7 @@ class HuggingFaceMedia(AsyncGeneratorProvider, ProviderModelMixin):
209 211 debug.error(f"{cls.__name__}: Error {response.status} with {provider_key} and {provider_id}")
210 212 continue
211 213 if response.status == 404:
212 raise ModelNotSupportedError(f"Model is not supported: {model}")
214 raise ModelNotFoundError(f"Model not found: {model}")
213 215 await raise_for_status(response)
214 216 if response.headers.get("Content-Type", "").startswith("application/json"):
215 217 result = await response.json()
Modified g4f/Provider/needs_auth/hf/__init__.py +3 -3
@@ -4,7 +4,7 @@ import random
4 4
5 5 from ....typing import AsyncResult, Messages
6 6 from ....providers.response import ImageResponse
7 from ....errors import ModelNotSupportedError, MissingAuthError
7 from ....errors import ModelNotFoundError, MissingAuthError
8 8 from ...base_provider import AsyncGeneratorProvider, ProviderModelMixin
9 9 from .HuggingChat import HuggingChat
10 10 from .HuggingFaceAPI import HuggingFaceAPI
@@ -58,7 +58,7 @@ class HuggingFace(AsyncGeneratorProvider, ProviderModelMixin):
58 58 async for chunk in HuggingFaceMedia.create_async_generator(model, messages, **kwargs):
59 59 yield chunk
60 60 return
61 except ModelNotSupportedError:
61 except ModelNotFoundError:
62 62 pass
63 63 if model in cls.image_models:
64 64 if "api_key" not in kwargs:
@@ -71,6 +71,6 @@ class HuggingFace(AsyncGeneratorProvider, ProviderModelMixin):
71 71 try:
72 72 async for chunk in HuggingFaceAPI.create_async_generator(model, messages, **kwargs):
73 73 yield chunk
74 except (ModelNotSupportedError, MissingAuthError):
74 except (ModelNotFoundError, MissingAuthError):
75 75 async for chunk in HuggingFaceInference.create_async_generator(model, messages, **kwargs):
76 76 yield chunk
Modified g4f/Provider/template/OpenaiTemplate.py +3 -1
@@ -66,7 +66,7 @@ class OpenaiTemplate(AsyncGeneratorProvider, ProviderModelMixin, RaiseErrorMixin
66 66 headers: dict = None,
67 67 impersonate: str = None,
68 68 extra_parameters: list[str] = ["tools", "parallel_tool_calls", "tool_choice", "reasoning_effort", "logit_bias", "modalities", "audio"],
69 extra_body: dict = {},
69 extra_body: dict = None,
70 70 **kwargs
71 71 ) -> AsyncResult:
72 72 if api_key is None and cls.api_key is not None:
@@ -98,6 +98,8 @@ class OpenaiTemplate(AsyncGeneratorProvider, ProviderModelMixin, RaiseErrorMixin
98 98 return
99 99
100 100 extra_parameters = {key: kwargs[key] for key in extra_parameters if key in kwargs}
101 if extra_body is None:
102 extra_body = {}
101 103 data = filter_none(
102 104 messages=list(render_messages(messages, media)),
103 105 model=model,
Modified g4f/constants.py +3 -2
Modified g4f/errors.py +0 -3
Modified g4f/gui/server/api.py +1 -1
Modified g4f/image/__init__.py +1 -1
Modified g4f/image/copy_images.py +19 -7
Modified g4f/providers/base_provider.py +2 -2
Modified g4f/requests/raise_for_status.py +6 -3
Modified setup.py +4 -15