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

XFEstudio/gpt4free

Add get_models to GeminiPro provider

ec9df598
Heiner Lohaus <hlohaus@users.noreply.github.com>
提交于

代码差异

6 个文件 +44 -24
Modified g4f/Provider/Cloudflare.py +4 -7
@@ -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,
Modified g4f/Provider/PollinationsAI.py +1 -1
@@ -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:
Modified g4f/Provider/needs_auth/DeepInfra.py +1 -1
@@ -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()
Modified g4f/Provider/needs_auth/GeminiPro.py +29 -7
@@ -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}"}
Modified g4f/Provider/needs_auth/OpenaiAPI.py +4 -4
@@ -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 }
Modified g4f/providers/base_provider.py +5 -4
@@ -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