返回提交历史
Modified
g4f/Provider/PollinationsAI.py
+1
-1
Modified
g4f/Provider/needs_auth/LMArena.py
+59
-19
Modified
g4f/gui/server/api.py
+3
-20
XFEstudio/gpt4free
feat: Enhance LMArena and PollinationsAI to support video models and improve model caching
68133b22
代码差异
3 个文件
+63
-40
@@ -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()
@@ -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:
@@ -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(