返回提交历史
Modified
g4f/Provider/needs_auth/Cohere.py
+42
-84
XFEstudio/gpt4free
Refactor Cohere provider to update API endpoint and improve model retrieval logic
a1c3ed72
代码差异
1 个文件
+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: