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

XFEstudio/gpt4free

Refactor Cohere provider to update API endpoint and improve model retrieval logic

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

代码差异

1 个文件 +42 -84
Modified g4f/Provider/needs_auth/Cohere.py +42 -84
@@ -1,42 +1,37 @@
1 1 from __future__ import annotations
2 2
3 import json
4 from typing import Optional
3 import requests
5 4
6 5 from ..helper import filter_none
7 6 from ...typing import AsyncResult, Messages
8 from ...requests import StreamSession, raise_for_status
7 from ...requests import StreamSession, raise_for_status, sse_stream
9 8 from ...providers.response import FinishReason, Usage
10 9 from ...errors import MissingAuthError
11 10 from ..base_provider import AsyncGeneratorProvider, ProviderModelMixin
12 from ...tools.run_tools import AuthManager
13 11 from ... import debug
14 12
15 13 class Cohere(AsyncGeneratorProvider, ProviderModelMixin):
16 14 label = "Cohere API"
17 15 url = "https://cohere.com"
18 16 login_url = "https://dashboard.cohere.com/api-keys"
19 api_base = "https://api.cohere.ai/v1"
17 api_endpoint = "https://api.cohere.ai/v2/chat"
20 18 working = True
19 active_by_default = True
21 20 needs_auth = True
22 21 supports_stream = True
23 22 supports_system_message = True
24 23 supports_message_history = True
25 24
26 25 default_model = "command-r-plus"
27 models = [
28 default_model,
29 "command-r",
30 "command",
31 "command-nightly",
32 "command-light",
33 "command-light-nightly",
34 ]
35
36 model_aliases = {
37 "command-r-plus-08-2024": "command-r-plus",
38 "command-r-08-2024": "command-r",
39 }
26
27 @classmethod
28 def get_models(cls, **kwargs):
29 if not cls.models:
30 url = "https://api.cohere.com/v1/models?page_size=500&endpoint=chat"
31 models = requests.get(url).json().get("models", [])
32 cls.models = [model.get("name") for model in models if "chat" in model.get("endpoints")]
33 cls.vision_models = {model.get("name") for model in models if model.get("supports_vision")}
34 return cls.models
40 35
41 36 @classmethod
42 37 async def create_async_generator(
@@ -51,43 +46,14 @@ class Cohere(AsyncGeneratorProvider, ProviderModelMixin):
51 46 top_k: int = None,
52 47 top_p: float = None,
53 48 stop: list[str] = None,
54 stream: bool = False,
49 stream: bool = True,
55 50 headers: dict = None,
56 51 impersonate: str = None,
57 52 **kwargs
58 53 ) -> AsyncResult:
59 if api_key is None:
60 api_key = AuthManager.load_api_key(cls)
61 54 if api_key is None:
62 55 raise MissingAuthError('Add a "api_key"')
63 56
64 # Convert messages to Cohere format
65 system_message = None
66 chat_history = []
67 user_message = None
68
69 # Filter out system messages first
70 system_messages = [msg for msg in messages if msg.get("role") == "system"]
71 if system_messages:
72 system_message = "\n".join([msg.get("content", "") for msg in system_messages])
73
74 # Process conversation messages (non-system)
75 conversation_messages = [msg for msg in messages if msg.get("role") != "system"]
76
77 # The last message should be from user
78 if conversation_messages and conversation_messages[-1].get("role") == "user":
79 user_message = conversation_messages[-1].get("content", "")
80 # All previous messages become chat history
81 for msg in conversation_messages[:-1]:
82 role = msg.get("role")
83 content = msg.get("content", "")
84 if role == "user":
85 chat_history.append({"role": "USER", "message": content})
86 elif role == "assistant":
87 chat_history.append({"role": "CHATBOT", "message": content})
88 else:
89 raise ValueError("The last message must be from the user")
90
91 57 async with StreamSession(
92 58 proxy=proxy,
93 59 headers=cls.get_headers(stream, api_key, headers),
@@ -95,19 +61,16 @@ class Cohere(AsyncGeneratorProvider, ProviderModelMixin):
95 61 impersonate=impersonate,
96 62 ) as session:
97 63 data = filter_none(
98 message=user_message,
64 messages=messages,
99 65 model=cls.get_model(model, api_key=api_key),
100 66 temperature=temperature,
101 67 max_tokens=max_tokens,
102 68 k=top_k,
103 69 p=top_p,
104 70 stop_sequences=stop,
105 preamble=system_message,
106 chat_history=chat_history if chat_history else None,
107 71 stream=stream,
108 72 )
109
110 async with session.post(f"{cls.api_base}/chat", json=data) as response:
73 async with session.post(cls.api_endpoint, json=data) as response:
111 74 await raise_for_status(response)
112 75
113 76 if not stream:
@@ -120,40 +83,35 @@ class Cohere(AsyncGeneratorProvider, ProviderModelMixin):
120 83 yield FinishReason("stop")
121 84 elif data["finish_reason"] == "MAX_TOKENS":
122 85 yield FinishReason("length")
123 if "meta" in data and "tokens" in data["meta"]:
86 if "usage" in data:
87 tokens = data.get("usage", {}).get("tokens", {})
124 88 yield Usage(
125 prompt_tokens=data["meta"]["tokens"]["input_tokens"],
126 completion_tokens=data["meta"]["tokens"]["output_tokens"],
127 total_tokens=data["meta"]["tokens"]["input_tokens"] + data["meta"]["tokens"]["output_tokens"]
89 prompt_tokens=tokens.get("input_tokens"),
90 completion_tokens=tokens.get("output_tokens"),
91 total_tokens=tokens.get("input_tokens", 0) + tokens.get("output_tokens", 0),
92 billed_units=data.get("usage", {}).get("billed_units")
128 93 )
129 94 else:
130 async for line in response.iter_lines():
131 if line.startswith(b"data: "):
132 chunk = line[6:]
133 if chunk == b"[DONE]":
134 break
135 try:
136 data = json.loads(chunk)
137 cls.raise_error(data)
138
139 if "event_type" in data:
140 if data["event_type"] == "text-generation":
141 if "text" in data:
142 yield data["text"]
143 elif data["event_type"] == "stream-end":
144 if "finish_reason" in data:
145 if data["finish_reason"] == "COMPLETE":
146 yield FinishReason("stop")
147 elif data["finish_reason"] == "MAX_TOKENS":
148 yield FinishReason("length")
149 if "meta" in data and "tokens" in data["meta"]:
150 yield Usage(
151 prompt_tokens=data["meta"]["tokens"]["input_tokens"],
152 completion_tokens=data["meta"]["tokens"]["output_tokens"],
153 total_tokens=data["meta"]["tokens"]["input_tokens"] + data["meta"]["tokens"]["output_tokens"]
154 )
155 except json.JSONDecodeError:
156 continue
95 async for data in sse_stream(response):
96 cls.raise_error(data)
97 if "type" in data:
98 if data["type"] == "content-delta":
99 yield data.get("delta", {}).get("message", {}).get("content", {}).get("text")
100 elif data["type"] == "message-end":
101 delta = data.get("delta", {})
102 if "finish_reason" in delta:
103 if delta["finish_reason"] == "COMPLETE":
104 yield FinishReason("stop")
105 elif delta["finish_reason"] == "MAX_TOKENS":
106 yield FinishReason("length")
107 if "usage" in delta:
108 tokens = delta.get("usage", {}).get("tokens", {})
109 yield Usage(
110 prompt_tokens=tokens.get("input_tokens"),
111 completion_tokens=tokens.get("output_tokens"),
112 total_tokens=tokens.get("input_tokens", 0) + tokens.get("output_tokens", 0),
113 billed_units=delta.get("usage", {}).get("billed_units")
114 )
157 115
158 116 @classmethod
159 117 def get_headers(cls, stream: bool, api_key: str = None, headers: dict = None) -> dict: