返回提交历史
Modified
g4f/Provider/Cloudflare.py
+4
-7
Modified
g4f/Provider/PollinationsAI.py
+1
-1
Modified
g4f/Provider/needs_auth/DeepInfra.py
+1
-1
Modified
g4f/Provider/needs_auth/GeminiPro.py
+29
-7
Modified
g4f/Provider/needs_auth/OpenaiAPI.py
+4
-4
Modified
g4f/providers/base_provider.py
+5
-4
XFEstudio/gpt4free
Add get_models to GeminiPro provider
ec9df598
代码差异
6 个文件
+44
-24
@@ -2,7 +2,6 @@ from __future__ import annotations
2
2
3
3
import asyncio
4
4
import json
5
import uuid
6
5
7
6
from ..typing import AsyncResult, Messages, Cookies
8
7
from .base_provider import AsyncGeneratorProvider, ProviderModelMixin, get_running_loop
@@ -37,18 +36,16 @@ class Cloudflare(AsyncGeneratorProvider, ProviderModelMixin):
37
36
if not cls.models:
38
37
if cls._args is None:
39
38
get_running_loop(check_nested=True)
40
args = get_args_from_nodriver(cls.url, cookies={
41
'__cf_bm': uuid.uuid4().hex,
42
})
39
args = get_args_from_nodriver(cls.url)
43
40
cls._args = asyncio.run(args)
44
41
with Session(**cls._args) as session:
45
42
response = session.get(cls.models_url)
46
43
cls._args["cookies"] = merge_cookies(cls._args["cookies"] , response)
47
44
try:
48
45
raise_for_status(response)
49
except ResponseStatusError as e:
46
except ResponseStatusError:
50
47
cls._args = None
51
raise e
48
raise
52
49
json_data = response.json()
53
50
cls.models = [model.get("name") for model in json_data.get("models")]
54
51
return cls.models
@@ -64,9 +61,9 @@ class Cloudflare(AsyncGeneratorProvider, ProviderModelMixin):
64
61
timeout: int = 300,
65
62
**kwargs
66
63
) -> AsyncResult:
67
model = cls.get_model(model)
68
64
if cls._args is None:
69
65
cls._args = await get_args_from_nodriver(cls.url, proxy, timeout, cookies)
66
model = cls.get_model(model)
70
67
data = {
71
68
"messages": messages,
72
69
"lora": None,
@@ -40,7 +40,7 @@ class PollinationsAI(OpenaiAPI):
40
40
}
41
41
42
42
@classmethod
43
def get_models(cls):
43
def get_models(cls, **kwargs):
44
44
if not hasattr(cls, 'image_models'):
45
45
cls.image_models = []
46
46
if not cls.image_models:
@@ -14,7 +14,7 @@ class DeepInfra(OpenaiAPI):
14
14
default_model = "meta-llama/Meta-Llama-3.1-70B-Instruct"
15
15
16
16
@classmethod
17
def get_models(cls):
17
def get_models(cls, **kwargs):
18
18
if not cls.models:
19
19
url = 'https://api.deepinfra.com/models/featured'
20
20
models = requests.get(url).json()
@@ -2,30 +2,52 @@ from __future__ import annotations
2
2
3
3
import base64
4
4
import json
5
import requests
5
6
from aiohttp import ClientSession, BaseConnector
6
7
7
8
from ...typing import AsyncResult, Messages, ImagesType
8
from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin
9
9
from ...image import to_bytes, is_accepted_format
10
10
from ...errors import MissingAuthError
11
from ...requests.raise_for_status import raise_for_status
12
from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin
11
13
from ..helper import get_connector
14
from ... import debug
12
15
13
16
class GeminiPro(AsyncGeneratorProvider, ProviderModelMixin):
14
17
label = "Google Gemini API"
15
18
url = "https://ai.google.dev"
16
19
api_base = "https://generativelanguage.googleapis.com/v1beta"
20
17
21
working = True
18
22
supports_message_history = True
19
23
needs_auth = True
20
24
21
25
default_model = "gemini-1.5-pro"
22
26
default_vision_model = default_model
23
models = [default_model, "gemini-pro", "gemini-1.5-flash", "gemini-1.5-flash-8b"]
27
fallback_models = [default_model, "gemini-pro", "gemini-1.5-flash", "gemini-1.5-flash-8b"]
24
28
model_aliases = {
25
29
"gemini-flash": "gemini-1.5-flash",
26
30
"gemini-flash": "gemini-1.5-flash-8b",
27
31
}
28
32
33
@classmethod
34
def get_models(cls, api_key: str = None, api_base: str = api_base) -> list[str]:
35
if not cls.models:
36
try:
37
response = requests.get(f"{api_base}/models?key={api_key}")
38
raise_for_status(response)
39
data = response.json()
40
cls.models = [
41
model.get("name").split("/").pop()
42
for model in data.get("models")
43
if "generateContent" in model.get("supportedGenerationMethods")
44
]
45
cls.models.sort()
46
except Exception as e:
47
debug.log(e)
48
cls.models = cls.fallback_models
49
return cls.models
50
29
51
@classmethod
30
52
async def create_async_generator(
31
53
cls,
@@ -34,17 +56,17 @@ class GeminiPro(AsyncGeneratorProvider, ProviderModelMixin):
34
56
stream: bool = False,
35
57
proxy: str = None,
36
58
api_key: str = None,
37
api_base: str = "https://generativelanguage.googleapis.com/v1beta",
59
api_base: str = api_base,
38
60
use_auth_header: bool = False,
39
61
images: ImagesType = None,
40
62
connector: BaseConnector = None,
41
63
**kwargs
42
64
) -> AsyncResult:
43
model = cls.get_model(model)
44
45
65
if not api_key:
46
66
raise MissingAuthError('Add a "api_key"')
47
67
68
model = cls.get_model(model, api_key=api_key, api_base=api_base)
69
48
70
headers = params = None
49
71
if use_auth_header:
50
72
headers = {"Authorization": f"Bearer {api_key}"}
@@ -23,13 +23,13 @@ class OpenaiAPI(AsyncGeneratorProvider, ProviderModelMixin):
23
23
fallback_models = []
24
24
25
25
@classmethod
26
def get_models(cls, api_key: str = None):
26
def get_models(cls, api_key: str = None, api_base: str = api_base) -> list[str]:
27
27
if not cls.models:
28
28
try:
29
29
headers = {}
30
30
if api_key is not None:
31
31
headers["authorization"] = f"Bearer {api_key}"
32
response = requests.get(f"{cls.api_base}/models", headers=headers)
32
response = requests.get(f"{api_base}/models", headers=headers)
33
33
raise_for_status(response)
34
34
data = response.json()
35
35
cls.models = [model.get("id") for model in data.get("data")]
@@ -82,7 +82,7 @@ class OpenaiAPI(AsyncGeneratorProvider, ProviderModelMixin):
82
82
) as session:
83
83
data = filter_none(
84
84
messages=messages,
85
model=cls.get_model(model),
85
model=cls.get_model(model, api_key=api_key, api_base=api_base),
86
86
temperature=temperature,
87
87
max_tokens=max_tokens,
88
88
top_p=top_p,
@@ -147,4 +147,4 @@ class OpenaiAPI(AsyncGeneratorProvider, ProviderModelMixin):
147
147
if api_key is not None else {}
148
148
),
149
149
**({} if headers is None else headers)
150
}
150
}
@@ -243,19 +243,20 @@ class ProviderModelMixin:
243
243
last_model: str = None
244
244
245
245
@classmethod
246
def get_models(cls) -> list[str]:
246
def get_models(cls, **kwargs) -> list[str]:
247
247
if not cls.models and cls.default_model is not None:
248
248
return [cls.default_model]
249
249
return cls.models
250
250
251
251
@classmethod
252
def get_model(cls, model: str) -> str:
252
def get_model(cls, model: str, **kwargs) -> str:
253
253
if not model and cls.default_model is not None:
254
254
model = cls.default_model
255
255
elif model in cls.model_aliases:
256
256
model = cls.model_aliases[model]
257
elif model not in cls.get_models() and cls.models:
258
raise ModelNotSupportedError(f"Model is not supported: {model} in: {cls.__name__}")
257
else:
258
if model not in cls.get_models(**kwargs) and cls.models:
259
raise ModelNotSupportedError(f"Model is not supported: {model} in: {cls.__name__}")
259
260
cls.last_model = model
260
261
debug.last_model = model
261
262
return model