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

XFEstudio/gpt4free

Enhance Cohere and OpenRouter providers with API key handling and model management improvements

e1e34ecc
hlohaus <983577+hlohaus@users.noreply.github.com>
提交于

代码差异

3 个文件 +15 -10
Modified g4f/Provider/needs_auth/Cohere.py +6 -2
@@ -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
Modified g4f/Provider/needs_auth/OpenRouter.py +7 -7
@@ -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
Modified g4f/Provider/template/OpenaiTemplate.py +2 -1
@@ -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,