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

XFEstudio/gpt4free

Refactor PollinationsAI model handling and improve API response structure

43f010fb
hlohaus <hlohaus@users.noreply.github.com>
提交于

代码差异

4 个文件 +37 -33
Modified g4f/Provider/PollinationsAI.py +27 -26
@@ -55,7 +55,7 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
55 55 default_image_model = "flux"
56 56 default_vision_model = default_model
57 57 default_voice = "alloy"
58 text_models = [default_model]
58 text_models = {default_model: {"id": default_model}}
59 59 image_models = [default_image_model, "turbo", "kontext"]
60 60 audio_models = {}
61 61 vision_models = [default_vision_model]
@@ -78,7 +78,14 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
78 78 current_models_endpoint: Optional[str] = None
79 79
80 80 @classmethod
81 def get_balance(cls, api_key: str, timeout: Optional[float] = None) -> Optional[float]:
81 async def get_quota(cls, api_key: Optional[str] = None, timeout: Optional[float] = None) -> dict:
82 balance = cls.get_balance(api_key, timeout)
83 if balance is not None:
84 return {"balance": balance}
85 return None
86
87 @classmethod
88 def get_balance(cls, api_key: Optional[str] = None, timeout: Optional[float] = None) -> Optional[float]:
82 89 try:
83 90 headers = None
84 91 if api_key:
@@ -96,6 +103,8 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
96 103 @classmethod
97 104 def get_models(cls, api_key: Optional[str] = None, timeout: Optional[float] = None, **kwargs):
98 105 def get_alias(model: dict) -> str:
106 if isinstance(model, str):
107 return model
99 108 alias = model.get("name")
100 109 if (model.get("aliases")):
101 110 alias = model.get("aliases")[0]
@@ -141,16 +150,20 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
141 150
142 151 # Add extra image models if not already in the list
143 152 for model in new_image_models:
144 alias = get_alias(model) if isinstance(model, dict) else model
153 alias = get_alias(model)
154 model["label"] = alias
145 155 if model not in image_models:
146 if isinstance(model, str) or "image" in model.get("output_modalities", []):
147 image_models.append(alias)
148 if isinstance(model, dict) and alias != model.get("name"):
156 if "image" in model.get("output_modalities", []):
157 if model.get("name") not in image_models:
158 image_models.append(model.get("name"))
159 if alias not in image_models:
160 image_models.append(alias)
161 for alias in model.get("aliases", []):
149 162 cls.model_aliases[alias] = model.get("name")
150 163
151 164 cls.image_models = image_models
152 cls.video_models = [get_alias(model) for model in new_image_models if isinstance(model, dict) and "video" in model.get("output_modalities", [])]
153
165 cls.video_models = [model.get("name") if isinstance(model, dict) else model for model in new_image_models if isinstance(model, dict) and "video" in model.get("output_modalities", [])]
166 cls.video_models = [get_alias(model) for model in cls.video_models if get_alias(model) != model]
154 167 text_response = requests.get(cls.text_models_endpoint, timeout=timeout)
155 168 if not text_response.ok:
156 169 text_response = requests.get(cls.text_models_endpoint, timeout=timeout)
@@ -167,30 +180,18 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
167 180 if model in cls.audio_models and alias not in cls.audio_models:
168 181 cls.audio_models.update({alias: {}})
169 182
170 cls.vision_models.extend([
171 get_alias(model)
172 for model in models
173 if model.get("vision") and get_alias(model) not in cls.vision_models
174 ])
175
183 cls.vision_models = [model.get("name") for model in models if "image" in model.get("input_modalities", [])]
184 cls.vision_models.extend([get_alias(model) for model in cls.vision_models if get_alias(model) != model])
176 185 for model in models:
177 alias = get_alias(model)
178 if alias != model.get("name"):
186 for alias in model.get("aliases", []):
179 187 cls.model_aliases[alias] = model.get("name")
180 if alias not in cls.text_models:
181 cls.text_models.append(alias)
182 elif model.get("name") not in cls.text_models:
183 cls.text_models.append(model.get("name"))
184 188 cls.live += 1
185 189 cls.swap_model_aliases = {v: k for k, v in cls.model_aliases.items()}
186
190 cls.text_models = {model.get("name"): {"id": model.get("name"), "label": get_alias(model), **model} for model in models}
191 cls.models = cls.text_models.copy()
192 cls.models.update({model.get("name"): {"id": model.get("name"), "label": get_alias(model), **model} for model in new_image_models})
187 193 finally:
188 194 cls.current_models_endpoint = models_url
189 # Return unique models across all categories
190 all_models = cls.text_models.copy()
191 all_models.extend(cls.image_models)
192 all_models.extend(cls.audio_models.keys())
193 cls.models = all_models
194 195 # Cache the models to a file
195 196 try:
196 197 path.parent.mkdir(parents=True, exist_ok=True)
Modified g4f/gui/server/api.py +2 -2
@@ -76,8 +76,8 @@ class Api:
76 76 models = method()
77 77 if has_grouped_models:
78 78 return [{
79 "group": model["group"],
80 "models": [get_model_data(provider, name) for name in model["models"]]
79 "group": model.get("group"),
80 "models": [get_model_data(provider, name) for name in (model.get("models", {}).values() if isinstance(model.get("models"), dict) else model.get("models", []))]
81 81 } for model in models]
82 82 return [
83 83 get_model_data(provider, model)
Modified g4f/providers/asyncio.py +4 -1
@@ -35,7 +35,10 @@ def get_running_loop(check_nested: bool) -> Optional[AbstractEventLoop]:
35 35
36 36 # Fix for RuntimeError: async generator ignored GeneratorExit
37 37 async def await_callback(callback: Callable, timeout: Optional[int] = None) -> any:
38 return await asyncio.wait_for(callback(), timeout) if timeout is not None else await callback()
38 try:
39 return await asyncio.wait_for(callback(), timeout) if timeout is not None else await callback()
40 except TimeoutError as e:
41 raise TimeoutError("The operation timed out after {} seconds".format(timeout)) from e
39 42
40 43 async def async_generator_to_list(generator: AsyncIterator) -> list:
41 44 return [item async for item in generator]
Modified g4f/providers/base_provider.py +4 -4
@@ -16,7 +16,7 @@ except ImportError:
16 16
17 17 from ..typing import CreateResult, AsyncResult, Messages
18 18 from .types import BaseProvider
19 from .asyncio import get_running_loop, to_sync_generator, to_async_iterator
19 from .asyncio import get_running_loop, to_sync_generator, to_async_iterator, await_callback
20 20 from .response import BaseConversation, AuthResult
21 21 from .helper import concat_chunks
22 22 from ..cookies import get_cookies_dir
@@ -120,7 +120,7 @@ class AbstractProvider(BaseProvider):
120 120 def create_func() -> str:
121 121 return concat_chunks(cls.create_completion(model, messages, **kwargs))
122 122
123 return await asyncio.wait_for(
123 return await await_callback(
124 124 loop.run_in_executor(executor, create_func),
125 125 timeout=timeout
126 126 )
@@ -352,7 +352,7 @@ class AsyncGeneratorProvider(AbstractProvider):
352 352 if "stream_timeout" in kwargs or "timeout" in kwargs:
353 353 while True:
354 354 try:
355 yield await asyncio.wait_for(
355 yield await await_callback(
356 356 response.__anext__(),
357 357 timeout=kwargs.get("stream_timeout") if cls.use_stream_timeout else kwargs.get("timeout")
358 358 )
@@ -524,7 +524,7 @@ class AsyncAuthedProvider(AsyncGeneratorProvider, AuthFileMixin):
524 524 if "stream_timeout" in kwargs or "timeout" in kwargs:
525 525 while True:
526 526 try:
527 yield await asyncio.wait_for(
527 yield await await_callback(
528 528 response.__anext__(),
529 529 timeout=kwargs.get("stream_timeout") if cls.use_stream_timeout else kwargs.get("timeout")
530 530 )