返回提交历史
Modified
g4f/Provider/needs_auth/Cohere.py
+6
-2
Modified
g4f/Provider/needs_auth/OpenRouter.py
+7
-7
Modified
g4f/Provider/template/OpenaiTemplate.py
+2
-1
XFEstudio/gpt4free
Enhance Cohere and OpenRouter providers with API key handling and model management improvements
e1e34ecc
代码差异
3 个文件
+15
-10
@@ -7,6 +7,7 @@ from ...typing import AsyncResult, Messages
7
7
from ...requests import StreamSession, raise_for_status, sse_stream
8
8
from ...providers.response import FinishReason, Usage
9
9
from ...errors import MissingAuthError
10
from ...tools.run_tools import AuthManager
10
11
from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin
11
12
from ... import debug
12
13
@@ -18,6 +19,7 @@ class Cohere(AsyncGeneratorProvider, ProviderModelMixin):
18
19
working = True
19
20
active_by_default = True
20
21
needs_auth = True
22
models_needs_auth = True
21
23
supports_stream = True
22
24
supports_system_message = True
23
25
supports_message_history = True
@@ -25,10 +27,12 @@ class Cohere(AsyncGeneratorProvider, ProviderModelMixin):
25
27
default_model = "command-r-plus"
26
28
27
29
@classmethod
28
def get_models(cls, **kwargs):
30
def get_models(cls, api_key: str = None, **kwargs):
29
31
if not cls.models:
32
if not api_key:
33
api_key = AuthManager.load_api_key(cls)
30
34
url = "https://api.cohere.com/v1/models?page_size=500&endpoint=chat"
31
models = requests.get(url).json().get("models", [])
35
models = requests.get(url, headers={"Authorization": f"Bearer {api_key}" }).json().get("models", [])
32
36
cls.models = [model.get("name") for model in models if "chat" in model.get("endpoints")]
33
37
cls.vision_models = {model.get("name") for model in models if model.get("supports_vision")}
34
38
return cls.models
@@ -15,17 +15,17 @@ class OpenRouter(OpenaiTemplate):
15
15
class OpenRouterFree(OpenRouter):
16
16
parent = "OpenRouter"
17
17
label = "OpenRouter (free)"
18
max_tokens = 5012
18
19
19
20
@classmethod
20
21
def get_models(cls, api_key: str = None, **kwargs):
21
if not cls.models:
22
models = super().get_models(api_key=api_key, **kwargs)
23
models = [model for model in models if model.endswith(":free")]
24
cls.model_aliases = {model.replace(":free", ""): model for model in models}
25
cls.models = [model.replace(":free", "") for model in models]
26
cls.default_model = models[0] if models else cls.default_model
22
models = super().get_models(api_key=api_key, **kwargs)
23
models = [model for model in models if model.endswith(":free")]
24
cls.model_aliases = {model.replace(":free", ""): model for model in models}
25
cls.models = [model.replace(":free", "") for model in models]
26
cls.default_model = models[0] if models else cls.default_model
27
27
return cls.models
28
28
29
29
@classmethod
30
30
def get_model(cls, model: str, **kwargs) -> str:
31
31
# Load model aliases if not already done
@@ -28,6 +28,7 @@ class OpenaiTemplate(AsyncGeneratorProvider, ProviderModelMixin, RaiseErrorMixin
28
28
ssl = None
29
29
add_user = True
30
30
use_image_size = False
31
max_tokens: int = None
31
32
32
33
@classmethod
33
34
def get_models(cls, api_key: str = None, api_base: str = None) -> list[str]:
@@ -126,7 +127,7 @@ class OpenaiTemplate(AsyncGeneratorProvider, ProviderModelMixin, RaiseErrorMixin
126
127
messages=list(render_messages(messages, media)),
127
128
model=model,
128
129
temperature=temperature,
129
max_tokens=max_tokens,
130
max_tokens=max_tokens if max_tokens is not None else cls.max_tokens,
130
131
top_p=top_p,
131
132
stop=stop,
132
133
stream="audio" not in extra_parameters if stream is None else stream,