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

XFEstudio/gpt4free

feat: refactor provider create functions to class attributes and update calls

- Added `create_function` and `async_create_function` class attributes with default implementations in `base_provider.py` for `AbstractProvider`, `AsyncProvider`, and `AsyncGeneratorProvider` - Updated `get_create_function` and `get_async_create_function` methods to return these class attributes - Replaced calls to `provider.get_create_function()` and `provider.get_async_create_function()` with direct attribute access `provider.create_function` and `provider.async_create_function` across `g4f/__init__.py`, `g4f/client/__init__.py`, `g4f/providers/retry_provider.py`, and `g4f/tools/run_tools.py` - Removed redundant `get_create_function` and `get_async_create_function` methods from `providers/base_provider.py` and `providers/types.py` - Ensured all provider response calls now use the class attributes for creating completions asynchronously and synchronously as needed

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

代码差异

15 个文件 +142 -111
Modified g4f/Provider/Together.py +4 -4
@@ -10,20 +10,20 @@ from ..requests import StreamSession, raise_for_status
10 10 from ..errors import ModelNotFoundError
11 11 from .. import debug
12 12
13
14 13 class Together(OpenaiTemplate):
15 14 label = "Together"
16 15 url = "https://together.xyz"
16 login_url = "https://api.together.ai/"
17 17 api_base = "https://api.together.xyz/v1"
18 18 activation_endpoint = "https://www.codegeneration.ai/activate-v2"
19 19 models_endpoint = "https://api.together.xyz/v1/models"
20
20
21 21 working = True
22 22 needs_auth = False
23 23 supports_stream = True
24 24 supports_system_message = True
25 25 supports_message_history = True
26
26
27 27 default_model = 'meta-llama/Llama-4-Maverick-17B-128E-Instruct-FP8'
28 28 default_vision_model = default_model
29 29 default_image_model = 'black-forest-labs/FLUX.1.1-pro'
@@ -43,7 +43,7 @@ class Together(OpenaiTemplate):
43 43 model_configs = {} # Store model configurations including stop tokens
44 44 _models_cached = False
45 45 _api_key_cache = None
46
46
47 47 model_aliases = {
48 48 ### Models Chat/Language ###
49 49 # meta-llama
Modified g4f/Provider/needs_auth/OpenaiChat.py +6 -1
@@ -682,7 +682,12 @@ class OpenaiChat(AsyncAuthedProvider, ProviderModelMixin):
682 682 page.add_handler(nodriver.cdp.network.RequestWillBeSent, on_request)
683 683 page = await browser.get(cls.url)
684 684 user_agent = await page.evaluate("window.navigator.userAgent", return_by_value=True)
685 while not await page.evaluate("document.getElementById('prompt-textarea')?.id"):
685 textarea = None
686 while not textarea:
687 try:
688 textarea = await page.evaluate("document.getElementById('prompt-textarea')?.id")
689 except:
690 pass
686 691 await asyncio.sleep(1)
687 692 while not await page.evaluate("document.querySelector('[data-testid=\"send-button\"]')?.type"):
688 693 await asyncio.sleep(1)
Modified g4f/Provider/needs_auth/PuterJS.py +1 -0
@@ -15,6 +15,7 @@ from .. import debug
15 15
16 16 class PuterJS(AsyncGeneratorProvider, ProviderModelMixin):
17 17 label = "Puter.js"
18 parent = "Puter"
18 19 url = "https://docs.puter.com/playground"
19 20 login_url = "https://github.com/HeyPuter/puter-cli"
20 21 api_endpoint = "https://api.puter.com/drivers/call"
Modified g4f/Provider/template/OpenaiTemplate.py +12 -1
@@ -6,7 +6,7 @@ from ..helper import filter_none, format_media_prompt
6 6 from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin, RaiseErrorMixin
7 7 from ...typing import Union, AsyncResult, Messages, MediaListType
8 8 from ...requests import StreamSession, raise_for_status
9 from ...providers.response import FinishReason, ToolCalls, Usage, ImageResponse
9 from ...providers.response import FinishReason, ToolCalls, Usage, ImageResponse, ProviderInfo
10 10 from ...tools.media import render_messages
11 11 from ...errors import MissingAuthError, ResponseError
12 12 from ... import debug
@@ -93,6 +93,9 @@ class OpenaiTemplate(AsyncGeneratorProvider, ProviderModelMixin, RaiseErrorMixin
93 93 async with session.post(f"{api_base.rstrip('/')}/images/generations", json=data, ssl=cls.ssl) as response:
94 94 data = await response.json()
95 95 cls.raise_error(data, response.status)
96 model = data.get("model")
97 if model:
98 yield ProviderInfo(**cls.get_dict(), model=model)
96 99 await raise_for_status(response)
97 100 yield ImageResponse([image["url"] for image in data["data"]], prompt)
98 101 return
@@ -121,6 +124,9 @@ class OpenaiTemplate(AsyncGeneratorProvider, ProviderModelMixin, RaiseErrorMixin
121 124 data = await response.json()
122 125 cls.raise_error(data, response.status)
123 126 await raise_for_status(response)
127 model = data.get("model")
128 if model:
129 yield ProviderInfo(**cls.get_dict(), model=model)
124 130 choice = data["choices"][0]
125 131 if "content" in choice["message"] and choice["message"]["content"]:
126 132 yield choice["message"]["content"].strip()
@@ -134,8 +140,13 @@ class OpenaiTemplate(AsyncGeneratorProvider, ProviderModelMixin, RaiseErrorMixin
134 140 elif content_type.startswith("text/event-stream"):
135 141 await raise_for_status(response)
136 142 first = True
143 model_returned = False
137 144 async for data in response.sse():
138 145 cls.raise_error(data)
146 model = data.get("model")
147 if not model_returned and model:
148 yield ProviderInfo(**cls.get_dict(), model=model)
149 model_returned = True
139 150 choice = data["choices"][0]
140 151 if "content" in choice["delta"] and choice["delta"]["content"]:
141 152 delta = choice["delta"]["content"]
Modified g4f/__init__.py +2 -2
@@ -50,7 +50,7 @@ class ChatCompletion:
50 50 if ignore_stream:
51 51 kwargs["ignore_stream"] = True
52 52
53 result = provider.get_create_function()(model, messages, stream=stream, **kwargs)
53 result = provider.create_function(model, messages, stream=stream, **kwargs)
54 54
55 55 return result if stream or ignore_stream else concat_chunks(result)
56 56
@@ -76,7 +76,7 @@ class ChatCompletion:
76 76 if ignore_stream:
77 77 kwargs["ignore_stream"] = True
78 78
79 result = provider.get_async_create_function()(model, messages, stream=stream, **kwargs)
79 result = provider.async_create_function(model, messages, stream=stream, **kwargs)
80 80
81 81 if not stream and not ignore_stream:
82 82 if hasattr(result, "__aiter__"):
Modified g4f/api/__init__.py +1 -1
@@ -490,7 +490,7 @@ class Api:
490 490 config.provider = provider
491 491 if config.provider is None:
492 492 config.provider = AppConfig.media_provider
493 if credentials is not None and credentials.credentials != "secret":
493 if config.api_key is None and credentials is not None and credentials.credentials != "secret":
494 494 config.api_key = credentials.credentials
495 495 try:
496 496 response = await self.client.images.generate(
Modified g4f/api/stubs.py +3 -1
@@ -16,7 +16,7 @@ class RequestConfig(BaseModel):
16 16 top_p: Optional[float] = None
17 17 max_tokens: Optional[int] = None
18 18 stop: Union[list[str], str, None] = None
19 api_key: Optional[str] = None
19 api_key: Optional[Union[str, dict[str, str]]] = None
20 20 api_base: str = None
21 21 web_search: Optional[bool] = None
22 22 proxy: Optional[str] = None
@@ -70,6 +70,8 @@ class ImageGenerationConfig(BaseModel):
70 70 negative_prompt: Optional[str] = None
71 71 resolution: Optional[str] = None
72 72 audio: Optional[dict] = None
73 download_media: bool = True
74
73 75
74 76 @model_validator(mode='before')
75 77 def parse_size(cls, values):
Modified g4f/client/__init__.py +9 -4
@@ -375,7 +375,7 @@ class Completions:
375 375 kwargs["ignore_stream"] = True
376 376
377 377 response = iter_run_tools(
378 provider.get_create_function(),
378 provider.create_function,
379 379 model=model,
380 380 messages=messages,
381 381 stream=stream,
@@ -462,7 +462,7 @@ class Images:
462 462 if isinstance(provider_handler, IterListProvider):
463 463 for provider in provider_handler.providers:
464 464 try:
465 response = await self._generate_image_response(provider, provider.__name__, model, prompt, proxy=proxy, **kwargs)
465 response = await self._generate_image_response(provider, provider.__name__, model, prompt, proxy=proxy, api_key=api_key, **kwargs)
466 466 if response is not None:
467 467 provider_name = provider.__name__
468 468 break
@@ -485,21 +485,25 @@ class Images:
485 485
486 486 async def _generate_image_response(
487 487 self,
488 provider_handler,
489 provider_name,
488 provider_handler: ProviderType,
489 provider_name: str,
490 490 model: str,
491 491 prompt: str,
492 492 prompt_prefix: str = "Generate a image: ",
493 api_key: str = None,
493 494 **kwargs
494 495 ) -> MediaResponse:
495 496 messages = [{"role": "user", "content": f"{prompt_prefix}{prompt}"}]
496 497 items: list[MediaResponse] = []
498 if isinstance(api_key, dict):
499 api_key = api_key.get(provider_handler.get_parent())
497 500 if hasattr(provider_handler, "create_async_generator"):
498 501 async for item in provider_handler.create_async_generator(
499 502 model,
500 503 messages,
501 504 stream=True,
502 505 prompt=prompt,
506 api_key=api_key,
503 507 **kwargs
504 508 ):
505 509 if isinstance(item, (MediaResponse, AudioResponse)):
@@ -510,6 +514,7 @@ class Images:
510 514 messages,
511 515 True,
512 516 prompt=prompt,
517 api_key=api_key,
513 518 **kwargs
514 519 ):
515 520 if isinstance(item, (MediaResponse, AudioResponse)):
Modified g4f/image/copy_images.py +1 -1
@@ -157,7 +157,7 @@ async def copy_media(
157 157 if media_type not in ("application/octet-stream", "binary/octet-stream"):
158 158 if media_type not in MEDIA_TYPE_MAP:
159 159 raise ValueError(f"Unsupported media type: {media_type}")
160 if not media_extension:
160 if target is None and not media_extension:
161 161 media_extension = f".{MEDIA_TYPE_MAP[media_type]}"
162 162 target_path = f"{target_path}{media_extension}"
163 163 with open(target_path, "wb") as f:
Modified g4f/providers/any_provider.py +53 -33
@@ -1,22 +1,32 @@
1 1 from __future__ import annotations
2 2
3 3 import re
4 from typing import Dict, List, Set, Optional, Tuple, Any
5 from ..typing import AsyncResult, Messages, MediaListType
4 from ..typing import AsyncResult, Messages, MediaListType, Union
6 5 from ..errors import ModelNotFoundError
7 6 from ..image import is_data_an_audio
8 7 from ..providers.retry_provider import IterListProvider
9 8 from ..providers.types import ProviderType
10 9 from ..Provider.needs_auth import OpenaiChat, CopilotAccount
11 10 from ..Provider.hf_space import HuggingSpace
12 from ..Provider import Cloudflare, Gemini, Grok, DeepSeekAPI, PerplexityLabs, LambdaChat, PollinationsAI
11 from ..Provider import __map__
12 from ..Provider import Cloudflare, Gemini, Grok, DeepSeekAPI, PerplexityLabs, LambdaChat, PollinationsAI, PuterJS
13 13 from ..Provider import Microsoft_Phi_4_Multimodal, DeepInfraChat, Blackbox, OIVSCodeSer2, OIVSCodeSer0501, TeachAnything, Together, WeWordle, Yqcloud, Chatai, Free2GPT, ARTA, ImageLabs, LegacyLMArena
14 from ..Provider import EdgeTTS, gTTS, MarkItDown
14 from ..Provider import EdgeTTS, gTTS, MarkItDown, OpenAIFM
15 15 from ..Provider import HarProvider, HuggingFace, HuggingFaceMedia
16 16 from .base_provider import AsyncGeneratorProvider, ProviderModelMixin
17 17 from .. import Provider
18 18 from .. import models
19 19
20 MAIN_PROVIERS = [
21 OpenaiChat, Cloudflare, HarProvider, PerplexityLabs, Gemini, Grok, DeepSeekAPI, Blackbox,
22 OIVSCodeSer2, OIVSCodeSer0501, TeachAnything, Together, WeWordle, Yqcloud, Chatai, Free2GPT, ARTA, ImageLabs, LegacyLMArena,
23 HuggingSpace, LambdaChat, CopilotAccount, PollinationsAI, DeepInfraChat, HuggingFace, HuggingFaceMedia
24 ]
25
26 SPECIAL_PROVIDERS = [OpenaiChat, CopilotAccount, PollinationsAI, HuggingSpace, Cloudflare, PerplexityLabs, Gemini, Grok, LegacyLMArena, ARTA]
27
28 SPECIAL_PROVIDERS2 = [HarProvider, LambdaChat, DeepInfraChat, HuggingFace, HuggingFaceMedia, PuterJS]
29
20 30 LABELS = {
21 31 "default": "Default",
22 32 "openai": "OpenAI: ChatGPT",
@@ -31,6 +41,7 @@ LABELS = {
31 41 "mistral": "Mistral",
32 42 "PollinationsAI": "Pollinations AI",
33 43 "perplexity": "Perplexity Labs",
44 "openrouter": "OpenRouter",
34 45 "video": "Video Generation",
35 46 "image": "Image Generation",
36 47 "other": "Other Models",
@@ -45,16 +56,16 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
45 56 def get_grouped_models(cls, ignored: list[str] = []) -> dict[str, list[str]]:
46 57 unsorted_models = cls.get_models(ignored=ignored)
47 58 groups = {key: [] for key in LABELS.keys()}
48
59
49 60 # Always add default first
50 61 groups["default"].append("default")
51
62
52 63 for model in unsorted_models:
53 64 if model == "default":
54 65 continue # Already added
55
66
56 67 added = False
57
68
58 69 # Check for PollinationsAI models (with prefix)
59 70 if model.startswith("PollinationsAI:"):
60 71 groups["PollinationsAI"].append(model)
@@ -109,6 +120,10 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
109 120 elif model.startswith(("gpt-", "chatgpt-", "o1", "o1-", "o3-", "o4-")) or model in ("auto", "dall-e-3", "searchgpt"):
110 121 groups["openai"].append(model)
111 122 added = True
123 # Check for openrouter models
124 elif model.startswith(("openrouter:")):
125 groups["openrouter"].append(model)
126 added = True
112 127 # Check for video models
113 128 elif model in cls.video_models:
114 129 groups["video"].append(model)
@@ -117,7 +132,7 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
117 132 elif model in cls.image_models or "flux" in model.lower() or "stable-diffusion" in model.lower() or "sdxl" in model.lower() or "gpt-image" in model.lower():
118 133 groups["image"].append(model)
119 134 added = True
120
135
121 136 # If not categorized, check for special cases then put in other
122 137 if not added:
123 138 # CodeLlama is Meta's model
@@ -128,7 +143,7 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
128 143 groups["phi"].append(model)
129 144 else:
130 145 groups["other"].append(model)
131
146
132 147 return [
133 148 {"group": LABELS[group], "models": names} for group, names in groups.items()
134 149 ]
@@ -157,9 +172,9 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
157 172 model: len(providers) for model, providers in model_with_providers.items() if len(providers) > 1
158 173 }
159 174 all_models = [cls.default_model] + list(model_with_providers.keys())
160
175
161 176 # Process special providers
162 for provider in [OpenaiChat, CopilotAccount, PollinationsAI, HuggingSpace, Cloudflare, PerplexityLabs, Gemini, Grok, LegacyLMArena, ARTA]:
177 for provider in SPECIAL_PROVIDERS:
163 178 provider: ProviderType = provider
164 179 if not provider.working or provider.get_parent() in ignored:
165 180 continue
@@ -186,7 +201,7 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
186 201 cls.image_models.extend(arta_models)
187 202 else:
188 203 all_models.extend(provider.get_models())
189
204
190 205 # Update special model lists
191 206 if hasattr(provider, 'image_models'):
192 207 cls.image_models.extend(provider.image_models)
@@ -194,7 +209,7 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
194 209 cls.vision_models.extend(provider.vision_models)
195 210 if hasattr(provider, 'video_models'):
196 211 cls.video_models.extend(provider.video_models)
197
212
198 213 # Clean model names function
199 214 def clean_name(name: str) -> str:
200 215 name = name.split("/")[-1].split(":")[0].lower()
@@ -212,24 +227,24 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
212 227 name = name.replace("llama3", "llama-3")
213 228 name = name.replace("flux.1-", "flux-")
214 229 return name
215
230
216 231 # Process HAR providers
217 for provider in [HarProvider, LambdaChat, DeepInfraChat, HuggingFace, HuggingFaceMedia]:
232 for provider in SPECIAL_PROVIDERS2:
218 233 if not provider.working or provider.get_parent() in ignored:
219 234 continue
220 235 new_models = provider.get_models()
221 236 if provider == HuggingFaceMedia:
222 237 new_models = provider.video_models
223
238
224 239 # Add original models too, not just cleaned names
225 240 all_models.extend(new_models)
226
227 model_map = {clean_name(model): model for model in new_models}
241
242 model_map = {model if model.startswith("openrouter:") else clean_name(model): model for model in new_models}
228 243 if not provider.model_aliases:
229 244 provider.model_aliases = {}
230 245 provider.model_aliases.update(model_map)
231 246 all_models.extend(list(model_map.keys()))
232
247
233 248 # Update special model lists with both original and cleaned names
234 249 if hasattr(provider, 'image_models'):
235 250 cls.image_models.extend(provider.image_models)
@@ -240,18 +255,18 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
240 255 if hasattr(provider, 'video_models'):
241 256 cls.video_models.extend(provider.video_models)
242 257 cls.video_models.extend([clean_name(model) for model in provider.video_models])
243
258
244 259 # Process audio providers
245 260 for provider in [Microsoft_Phi_4_Multimodal, PollinationsAI]:
246 261 if provider.working and provider.get_parent() not in ignored:
247 262 cls.audio_models.update(provider.audio_models)
248
263
249 264 # Update model counts
250 265 cls.models_count.update({model: all_models.count(model) for model in all_models if all_models.count(model) > cls.models_count.get(model, 0)})
251
266
252 267 # Deduplicate and store
253 268 cls.models_storage[ignored_key] = list(dict.fromkeys([model if model else cls.default_model for model in all_models]))
254
269
255 270 return cls.models_storage[ignored_key]
256 271
257 272 @classmethod
@@ -262,6 +277,7 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
262 277 stream: bool = True,
263 278 media: MediaListType = None,
264 279 ignored: list[str] = [],
280 api_key: Union[str, dict[str, str]] = None,
265 281 **kwargs
266 282 ) -> AsyncResult:
267 283 cls.get_models(ignored=ignored)
@@ -284,7 +300,10 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
284 300 if "tools" in kwargs:
285 301 providers = [PollinationsAI]
286 302 elif "audio" in kwargs or "audio" in kwargs.get("modalities", []):
287 providers = [PollinationsAI, EdgeTTS, gTTS]
303 if kwargs.get("audio", {}).get("language") is None:
304 providers = [PollinationsAI, OpenAIFM, Gemini]
305 else:
306 providers = [PollinationsAI, OpenAIFM, EdgeTTS, gTTS]
288 307 elif has_audio:
289 308 providers = [PollinationsAI, Microsoft_Phi_4_Multimodal, MarkItDown]
290 309 elif has_image:
@@ -297,29 +316,30 @@ class AnyProvider(AsyncGeneratorProvider, ProviderModelMixin):
297 316 model = None
298 317 providers.append(provider)
299 318 else:
300 for provider in [
301 OpenaiChat, Cloudflare, HarProvider, PerplexityLabs, Gemini, Grok, DeepSeekAPI, Blackbox,
302 OIVSCodeSer2, OIVSCodeSer0501, TeachAnything, Together, WeWordle, Yqcloud, Chatai, Free2GPT, ARTA, ImageLabs, LegacyLMArena,
303 HuggingSpace, LambdaChat, CopilotAccount, PollinationsAI, DeepInfraChat, HuggingFace, HuggingFaceMedia,
304 ]:
319 extra_providers = []
320 if isinstance(api_key, dict):
321 for provider in api_key:
322 if provider in __map__ and __map__[provider] not in MAIN_PROVIERS:
323 extra_providers.append(__map__[provider])
324 for provider in MAIN_PROVIERS + extra_providers:
305 325 if provider.working:
306 326 if not model or model in provider.get_models() or model in provider.model_aliases:
307 327 providers.append(provider)
308 328 if model in models.__models__:
309 329 for provider in models.__models__[model][1]:
310 330 providers.append(provider)
311
312 331 providers = [provider for provider in providers if provider.working and provider.get_parent() not in ignored]
313 332 providers = list({provider.__name__: provider for provider in providers}.values())
314
333
315 334 if len(providers) == 0:
316 335 raise ModelNotFoundError(f"AnyProvider: Model {model} not found in any provider.")
317
336
318 337 async for chunk in IterListProvider(providers).create_async_generator(
319 338 model,
320 339 messages,
321 340 stream=stream,
322 341 media=media,
342 api_key=api_key,
323 343 **kwargs
324 344 ):
325 345 yield chunk
Modified g4f/providers/base_provider.py +32 -25
@@ -126,12 +126,30 @@ class AbstractProvider(BaseProvider):
126 126 )
127 127
128 128 @classmethod
129 def get_create_function(cls) -> callable:
130 return cls.create_completion
129 def create_function(cls, *args, **kwargs) -> CreateResult:
130 """
131 Creates a completion using the synchronous method.
132
133 Args:
134 **kwargs: Additional keyword arguments.
135
136 Returns:
137 CreateResult: The result of the completion creation.
138 """
139 return cls.create_completion(*args, **kwargs)
131 140
132 141 @classmethod
133 def get_async_create_function(cls) -> callable:
134 return cls.create_async
142 def async_create_function(cls, *args, **kwargs) -> AsyncResult:
143 """
144 Creates a completion using the synchronous method.
145
146 Args:
147 **kwargs: Additional keyword arguments.
148
149 Returns:
150 CreateResult: The result of the completion creation.
151 """
152 return cls.create_async(*args, **kwargs)
135 153
136 154 @classmethod
137 155 def get_parameters(cls, as_json: bool = False) -> dict[str, Parameter]:
@@ -264,14 +282,6 @@ class AsyncProvider(AbstractProvider):
264 282 """
265 283 raise NotImplementedError()
266 284
267 @classmethod
268 def get_create_function(cls) -> callable:
269 return cls.create_completion
270
271 @classmethod
272 def get_async_create_function(cls) -> callable:
273 return cls.create_async
274
275 285 class AsyncGeneratorProvider(AbstractProvider):
276 286 """
277 287 Provides asynchronous generator functionality for streaming results.
@@ -331,12 +341,17 @@ class AsyncGeneratorProvider(AbstractProvider):
331 341 raise NotImplementedError()
332 342
333 343 @classmethod
334 def get_create_function(cls) -> callable:
335 return cls.create_completion
344 def async_create_function(cls, *args, **kwargs) -> AsyncResult:
345 """
346 Creates a completion using the synchronous method.
336 347
337 @classmethod
338 def get_async_create_function(cls) -> callable:
339 return cls.create_async_generator
348 Args:
349 **kwargs: Additional keyword arguments.
350
351 Returns:
352 CreateResult: The result of the completion creation.
353 """
354 return cls.create_async_generator(*args, **kwargs)
340 355
341 356 class ProviderModelMixin:
342 357 default_model: str = None
@@ -417,14 +432,6 @@ class AsyncAuthedProvider(AsyncGeneratorProvider, AuthFileMixin):
417 432 return to_sync_generator(auth_result)
418 433 return asyncio.run(auth_result)
419 434
420 @classmethod
421 def get_create_function(cls) -> callable:
422 return cls.create_completion
423
424 @classmethod
425 def get_async_create_function(cls) -> callable:
426 return cls.create_async_generator
427
428 435 @classmethod
429 436 def write_cache_file(cls, cache_file: Path, auth_result: AuthResult = None):
430 437 if auth_result is not None:
Modified g4f/providers/retry_provider.py +14 -15
@@ -26,7 +26,6 @@ class IterListProvider(BaseRetryProvider):
26 26 self.shuffle = shuffle
27 27 self.working = True
28 28 self.last_provider: Type[BaseProvider] = None
29 self.add_api_key = False
30 29
31 30 def create_completion(
32 31 self,
@@ -35,6 +34,7 @@ class IterListProvider(BaseRetryProvider):
35 34 stream: bool = False,
36 35 ignore_stream: bool = False,
37 36 ignored: list[str] = [],
37 api_key: str = None,
38 38 **kwargs,
39 39 ) -> CreateResult:
40 40 """
@@ -55,8 +55,11 @@ class IterListProvider(BaseRetryProvider):
55 55 self.last_provider = provider
56 56 debug.log(f"Using {provider.__name__} provider")
57 57 yield ProviderInfo(**provider.get_dict(), model=model if model else getattr(provider, "default_model"))
58 extra_body = kwargs.copy()
59 if isinstance(api_key, dict):
60 extra_body["api_key"] = api_key.get(provider.get_parent())
58 61 try:
59 response = provider.get_create_function()(model, messages, stream=stream, **kwargs)
62 response = provider.create_function(model, messages, stream=stream, **extra_body)
60 63 for chunk in response:
61 64 if chunk:
62 65 yield chunk
@@ -66,7 +69,7 @@ class IterListProvider(BaseRetryProvider):
66 69 return
67 70 except Exception as e:
68 71 exceptions[provider.__name__] = e
69 debug.error(f"{provider.__name__} {type(e).__name__}: {e}")
72 debug.error(f"{provider.__name__}:", e)
70 73 if started:
71 74 raise e
72 75 yield e
@@ -92,12 +95,12 @@ class IterListProvider(BaseRetryProvider):
92 95 debug.log(f"Using {provider.__name__} provider" + (f" and {model} model" if model else ""))
93 96 yield ProviderInfo(**provider.get_dict(), model=model if model else getattr(provider, "default_model"))
94 97 extra_body = kwargs.copy()
95 if self.add_api_key or provider.__name__ in ["HuggingFace", "HuggingFaceMedia"]:
96 extra_body["api_key"] = api_key
98 if isinstance(api_key, dict):
99 extra_body["api_key"] = api_key.get(provider.get_parent())
97 100 if conversation is not None and hasattr(conversation, provider.__name__):
98 101 extra_body["conversation"] = JsonConversation(**getattr(conversation, provider.__name__))
99 102 try:
100 response = provider.get_async_create_function()(model, messages, stream=stream, **extra_body)
103 response = provider.async_create_function(model, messages, stream=stream, **extra_body)
101 104 if hasattr(response, "__aiter__"):
102 105 async for chunk in response:
103 106 if isinstance(chunk, JsonConversation):
@@ -118,18 +121,15 @@ class IterListProvider(BaseRetryProvider):
118 121 return
119 122 except Exception as e:
120 123 exceptions[provider.__name__] = e
121 debug.error(f"{provider.__name__} {type(e).__name__}: {e}")
124 debug.error(f"{provider.__name__}:", e)
122 125 if started:
123 126 raise e
124 127 yield e
125 128
126 129 raise_exceptions(exceptions)
127 130
128 def get_create_function(self) -> callable:
129 return self.create_completion
130
131 def get_async_create_function(self) -> callable:
132 return self.create_async_generator
131 create_function = create_completion
132 async_create_function = create_async_generator
133 133
134 134 def get_providers(self, stream: bool, ignored: list[str]) -> list[ProviderType]:
135 135 providers = [p for p in self.providers if (p.supports_stream or not stream) and p.__name__ not in ignored]
@@ -156,7 +156,6 @@ class RetryProvider(IterListProvider):
156 156 super().__init__(providers, shuffle)
157 157 self.single_provider_retry = single_provider_retry
158 158 self.max_retries = max_retries
159 self.add_api_key = True
160 159
161 160 def create_completion(
162 161 self,
@@ -185,7 +184,7 @@ class RetryProvider(IterListProvider):
185 184 try:
186 185 if debug.logging:
187 186 print(f"Using {provider.__name__} provider (attempt {attempt + 1})")
188 response = provider.get_create_function()(model, messages, stream=stream, **kwargs)
187 response = provider.create_function(model, messages, stream=stream, **kwargs)
189 188 for chunk in response:
190 189 yield chunk
191 190 if is_content(chunk):
@@ -218,7 +217,7 @@ class RetryProvider(IterListProvider):
218 217 for attempt in range(self.max_retries):
219 218 try:
220 219 debug.log(f"Using {provider.__name__} provider (attempt {attempt + 1})")
221 response = provider.get_async_create_function()(model, messages, stream=stream, **kwargs)
220 response = provider.async_create_function(model, messages, stream=stream, **kwargs)
222 221 if hasattr(response, "__aiter__"):
223 222 async for chunk in response:
224 223 yield chunk
Modified g4f/providers/tool_support.py +1 -1
Modified g4f/providers/types.py +2 -20
Modified g4f/tools/run_tools.py +1 -2