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

XFEstudio/gpt4free

feat: Enhance LMArena and PollinationsAI to support video models and improve model caching

68133b22
hlohaus <hlohaus@users.noreply.github.com>
提交于

代码差异

3 个文件 +63 -40
Modified g4f/Provider/PollinationsAI.py +1 -1
@@ -121,7 +121,7 @@ class PollinationsAI(AsyncGeneratorProvider, ProviderModelMixin):
121 121 image_url = cls.image_models_endpoint
122 122
123 123 if cls.current_models_endpoint != models_url:
124 path = Path(get_cookies_dir()) / "models" / datetime.today().strftime('%Y-%m-%d') / f"{secure_filename(models_url)}.json"
124 path = Path(get_cookies_dir()) / ".models" / datetime.today().strftime('%Y-%m-%d') / f"{secure_filename(models_url)}.json"
125 125 if path.exists():
126 126 try:
127 127 data = path.read_text()
Modified g4f/Provider/needs_auth/LMArena.py +59 -19
@@ -7,6 +7,8 @@ import os
7 7 import re
8 8 import secrets
9 9 import time
10 from datetime import datetime
11 from pathlib import Path
10 12 from typing import Dict
11 13 from urllib.parse import urlparse
12 14
@@ -30,6 +32,8 @@ except ImportError:
30 32 from ...typing import AsyncResult, Messages, MediaListType
31 33 from ...requests import get_args_from_nodriver, raise_for_status, merge_cookies
32 34 from ...requests import StreamSession
35 from ...cookies import get_cookies_dir
36 from ...tools.files import secure_filename
33 37 from ...errors import ModelNotFoundError, CloudflareError, MissingAuthError, MissingRequirementsError, \
34 38 RateLimitError
35 39 from ...providers.response import FinishReason, Usage, JsonConversation, ImageResponse, Reasoning, PlainTextResponse, \
@@ -66,6 +70,8 @@ text_models = {model["publicName"]: model["id"] for model in models if
66 70 "text" in model["capabilities"]["outputCapabilities"]}
67 71 image_models = {model["publicName"]: model["id"] for model in models if
68 72 "image" in model["capabilities"]["outputCapabilities"]}
73 video_models = {model["publicName"]: model["id"] for model in models if
74 "video" in model["capabilities"]["outputCapabilities"]}
69 75 vision_models = [model["publicName"] for model in models if "image" in model["capabilities"]["inputCapabilities"]]
70 76
71 77 if has_nodriver:
@@ -91,18 +97,20 @@ class LMArena(AsyncGeneratorProvider, ProviderModelMixin, AuthFileMixin):
91 97 share_url = None
92 98 create_evaluation = "https://arena.ai/nextjs-api/stream/create-evaluation"
93 99 post_to_evaluation = "https://arena.ai/nextjs-api/stream/post-to-evaluation/{id}"
100 models_url = "https://arena.ai/?mode=direct"
94 101 working = True
95 102 active_by_default = True
96 103 use_stream_timeout = False
97 104
98 105 default_model = list(text_models.keys())[0]
99 models = list(text_models) + list(image_models)
106 models = list(text_models) + list(image_models) + list(video_models)
100 107 model_aliases = {
101 108 "flux-kontext": "flux-1-kontext-pro",
102 109 }
103 110 image_models = image_models
104 111 text_models = text_models
105 112 vision_models = vision_models
113 video_models = video_models
106 114 looked = False
107 115 _models_loaded = False
108 116 image_cache = True
@@ -117,6 +125,19 @@ class LMArena(AsyncGeneratorProvider, ProviderModelMixin, AuthFileMixin):
117 125 @classmethod
118 126 def get_models(cls, timeout: int = None, **kwargs) -> list[str]:
119 127 if not cls._models_loaded and has_curl_cffi:
128 # Try to load models from cache
129 path = Path(get_cookies_dir()) / ".models" / datetime.today().strftime('%Y-%m-%d') / f"{secure_filename(cls.models_url)}.json"
130 if path.exists():
131 try:
132 data = path.read_text()
133 models_data = json.loads(data)
134 for key, value in models_data.items():
135 setattr(cls, key, value)
136 cls._models_loaded = True
137 return cls.models
138 except Exception as e:
139 debug.error(f"Failed to load cached models from {path}: {e}")
140 # Open auth file
120 141 cache_file = cls.get_cache_file()
121 142 args = {}
122 143 if cache_file.exists():
@@ -129,24 +150,41 @@ class LMArena(AsyncGeneratorProvider, ProviderModelMixin, AuthFileMixin):
129 150 args = {}
130 151 if not args:
131 152 return cls.models
132 response = curl_cffi.get(f"{cls.url}/?mode=direct", **args, timeout=timeout)
153 response = curl_cffi.get(cls.models_url, **args, timeout=timeout)
133 154 if response.ok:
134 155 for line in response.text.splitlines():
135 if "initialModels" in line:
136 line = line.split("initialModels", maxsplit=1)[-1].split("initialModelAId")[0][3:-3]
137 line = line.encode("utf-8").decode("unicode_escape")
138 models = json.loads(line)
139 cls.text_models = {model["publicName"]: model["id"] for model in models if
140 "text" in model["capabilities"]["outputCapabilities"]}
141 cls.image_models = {model["publicName"]: model["id"] for model in models if
142 "image" in model["capabilities"]["outputCapabilities"]}
143 cls.vision_models = [model["publicName"] for model in models if
144 "image" in model["capabilities"]["inputCapabilities"]]
145 cls.models = list(cls.text_models) + list(cls.image_models)
146 cls.default_model = list(cls.text_models.keys())[0]
147 cls._models_loaded = True
148 cls.live += 1
149 break
156 if "initialModels" not in line:
157 continue
158 line = line.split("initialModels", maxsplit=1)[-1].split("initialModelAId")[0][3:-3]
159 line = line.encode("utf-8").decode("unicode_escape")
160 models = json.loads(line)
161 cls.text_models = {model["publicName"]: model["id"] for model in models if
162 "text" in model["capabilities"]["outputCapabilities"]}
163 cls.image_models = {model["publicName"]: model["id"] for model in models if
164 "image" in model["capabilities"]["outputCapabilities"]}
165 cls.video_models = {model["publicName"]: model["id"] for model in models if
166 "video" in model["capabilities"]["outputCapabilities"]}
167 cls.vision_models = [model["publicName"] for model in models if
168 "image" in model["capabilities"]["inputCapabilities"]]
169 cls.models = list(cls.text_models) + list(cls.image_models) + list(cls.video_models)
170 cls.default_model = list(cls.text_models.keys())[0]
171 cls._models_loaded = True
172 cls.live += 1
173 break
174 # Cache the models to a file
175 try:
176 path.parent.mkdir(parents=True, exist_ok=True)
177 with open(path, "w") as f:
178 json.dump({
179 "text_models": cls.text_models,
180 "image_models": cls.image_models,
181 "video_models": cls.video_models,
182 "vision_models": cls.vision_models,
183 "models": cls.models,
184 "default_model": cls.default_model
185 }, f, indent=4)
186 except Exception as e:
187 debug.error(f"Failed to cache models to {path}: {e}")
150 188 else:
151 189 cls.live -= 1
152 190 debug.log(f"Failed to load models from {cls.url}: {response.status_code} {response.reason}")
@@ -156,7 +194,7 @@ class LMArena(AsyncGeneratorProvider, ProviderModelMixin, AuthFileMixin):
156 194 async def get_models_async(cls) -> list[str]:
157 195 if not cls._models_loaded:
158 196 async with StreamSession() as session:
159 async with session.get(f"{cls.url}/?mode=direct",) as response:
197 async with session.get(cls.models_url) as response:
160 198 await cls.__load_actions(await response.text())
161 199 return cls.models
162 200
@@ -298,9 +336,11 @@ class LMArena(AsyncGeneratorProvider, ProviderModelMixin, AuthFileMixin):
298 336 "text" in model["capabilities"]["outputCapabilities"]}
299 337 cls.image_models = {model["publicName"]: model["id"] for model in models if
300 338 "image" in model["capabilities"]["outputCapabilities"]}
339 cls.video_models = {model["publicName"]: model["id"] for model in models if
340 "video" in model["capabilities"]["outputCapabilities"]}
301 341 cls.vision_models = [model["publicName"] for model in models if
302 342 "image" in model["capabilities"]["inputCapabilities"]]
303 cls.models = list(cls.text_models) + list(cls.image_models)
343 cls.models = list(cls.text_models) + list(cls.image_models) + list(cls.video_models)
304 344 cls.default_model = list(cls.text_models.keys())[0]
305 345 cls._models_loaded = True
306 346 elif 'children' in json_data:
Modified g4f/gui/server/api.py +3 -20
@@ -31,29 +31,13 @@ from ... import debug
31 31 logger = logging.getLogger(__name__)
32 32
33 33 class Api:
34 @staticmethod
35 def get_models():
36 return [{
37 "name": model.name,
38 "image": isinstance(model, models.ImageModel),
39 "vision": isinstance(model, models.VisionModel),
40 "audio": isinstance(model, models.AudioModel),
41 "video": isinstance(model, models.VideoModel),
42 "providers": [
43 getattr(provider, "parent", provider.__name__)
44 for provider in providers
45 if provider.working
46 ]
47 }
48 for model, providers in models.__models__.values()]
49
50 34 @staticmethod
51 35 def get_provider_models(provider: str, api_key: str = None, ignored: list = None):
52 36 def get_model_data(provider: ProviderModelMixin, model: str, default: bool = False) -> dict:
53 37 model_id = model.get("id") if isinstance(model, dict) else model
54 38 return {
55 "model": model_id,
56 "label": model_id.split(":")[-1] if provider.__name__ == "AnyProvider" and not model_id.startswith("openrouter:") else model_id,
39 "id": model_id,
40 "label": model_id,
57 41 "default": default or model_id == provider.default_model,
58 42 "vision": model_id in provider.vision_models,
59 43 "audio": False if provider.audio_models is None else model_id in provider.audio_models,
@@ -183,8 +167,7 @@ class Api:
183 167 debug.logs.append(" ".join([str(value) for value in values]))
184 168 if debug.logging:
185 169 debug.log_handler(*values, file=file)
186 if "user" not in kwargs:
187 debug.log = decorated_log
170 debug.log = decorated_log
188 171 proxy = os.environ.get("G4F_PROXY")
189 172 try:
190 173 model, provider_handler = get_model_and_provider(