diff options
| author | sigoden <sigoden@gmail.com> | 2024-12-06 18:14:43 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-12-06 18:14:43 +0800 |
| commit | b164a1587deff738d04a75665bdb863d82b3264c (patch) | |
| tree | 556e94d58e378246a736c638705abd84e9b40e81 | |
| parent | 5f7edc0ded187d86712639f3f7c48e6e4d72927b (diff) | |
| download | aichat-b164a1587deff738d04a75665bdb863d82b3264c.tar.gz | |
feat: update cohere api to v2 (#1041)
| -rwxr-xr-x | Argcfile.sh | 11 | ||||
| -rw-r--r-- | config.example.yaml | 2 | ||||
| -rw-r--r-- | models.yaml | 15 | ||||
| -rw-r--r-- | src/client/cohere.rs | 274 |
4 files changed, 113 insertions, 189 deletions
diff --git a/Argcfile.sh b/Argcfile.sh index c192232..c9fb34a 100755 --- a/Argcfile.sh +++ b/Argcfile.sh @@ -250,7 +250,7 @@ chat-claude() { # @flag -S --no-stream # @arg text~ chat-cohere() { - _wrapper curl -i https://api.cohere.ai/v1/chat \ + _wrapper curl -i https://api.cohere.ai/v2/chat \ -X POST \ -H 'Content-Type: application/json' \ -H "Authorization: Bearer $COHERE_API_KEY" \ @@ -398,7 +398,7 @@ _build_body() { else shift case "$kind" in - openai) + openai|cohere) echo '{ "model": "'$argc_model'", "messages": [ @@ -410,13 +410,6 @@ _build_body() { "stream": '$stream' }' ;; - cohere) - echo '{ - "model": "'$argc_model'", - "message": "'"$*"'", - "stream": '$stream' -}' - ;; claude) echo '{ "model": "'$argc_model'", diff --git a/config.example.yaml b/config.example.yaml index e50b7ed..25edcfc 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -181,7 +181,7 @@ clients: # See https://docs.cohere.com/docs/the-cohere-platform - type: cohere - api_base: https://api.cohere.ai/v1 # Optional + api_base: https://api.cohere.ai/v2 # Optional api_key: xxx # See https://docs.perplexity.ai/docs/getting-started diff --git a/models.yaml b/models.yaml index 20b6626..9d353e2 100644 --- a/models.yaml +++ b/models.yaml @@ -283,12 +283,27 @@ max_tokens_per_chunk: 512 default_chunk_size: 1000 max_batch_size: 96 + - name: embed-english-light-v3.0 + type: embedding + input_price: 0.1 + max_tokens_per_chunk: 512 + default_chunk_size: 700 + max_batch_size: 96 - name: embed-multilingual-v3.0 type: embedding input_price: 0.1 max_tokens_per_chunk: 512 default_chunk_size: 1000 max_batch_size: 96 + - name: embed-multilingual-light-v3.0 + type: embedding + input_price: 0.1 + max_tokens_per_chunk: 512 + default_chunk_size: 700 + max_batch_size: 96 + - name: rerank-v3.5 + type: reranker + max_input_tokens: 4096 - name: rerank-english-v3.0 type: reranker max_input_tokens: 4096 diff --git a/src/client/cohere.rs b/src/client/cohere.rs index 9726300..457ea66 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,3 +1,4 @@ +use super::openai::*; use super::openai_compatible::*; use super::*; @@ -6,7 +7,7 @@ use reqwest::RequestBuilder; use serde::Deserialize; use serde_json::{json, Value}; -const API_BASE: &str = "https://api.cohere.ai/v1"; +const API_BASE: &str = "https://api.cohere.ai/v2"; #[derive(Debug, Clone, Deserialize, Default)] pub struct CohereConfig { @@ -48,7 +49,7 @@ fn prepare_chat_completions( .unwrap_or_else(|_| API_BASE.to_string()); let url = format!("{}/chat", api_base.trim_end_matches('/')); - let body = build_chat_completions_body(data, &self_.model)?; + let body = openai_build_chat_completions_body(data, &self_.model); let mut request_data = RequestData::new(url, body); @@ -74,6 +75,7 @@ fn prepare_embeddings(self_: &CohereClient, data: &EmbeddingsData) -> Result<Req "model": self_.model.name(), "texts": data.texts, "input_type": input_type, + "embedding_types": ["float"], }); let mut request_data = RequestData::new(url, body); @@ -119,39 +121,67 @@ async fn chat_completions_streaming( handler: &mut SseHandler, _model: &Model, ) -> Result<()> { - let res = builder.send().await?; - let status = res.status(); - if !status.is_success() { - let data: Value = res.json().await?; - catch_error(&data, status.as_u16())?; - } else { - let handle = |data: &str| -> Result<()> { - let data: Value = serde_json::from_str(data)?; - debug!("stream-data: {data}"); - if let Some("text-generation") = data["event_type"].as_str() { - if let Some(text) = data["text"].as_str() { - handler.text(text)?; + let mut function_name = String::new(); + let mut function_arguments = String::new(); + let mut function_id = String::new(); + let handle = |message: SseMmessage| -> Result<bool> { + if message.data == "[DONE]" { + return Ok(true); + } + let data: Value = serde_json::from_str(&message.data)?; + debug!("stream-data: {data}"); + if let Some(typ) = data["type"].as_str() { + match typ { + "content-delta" => { + if let Some(text) = data["delta"]["message"]["content"]["text"].as_str() { + handler.text(text)?; + } } - } else if let Some("tool-calls-generation") = data["event_type"].as_str() { - if let Some(tool_calls) = data["tool_calls"].as_array() { - for call in tool_calls { - if let (Some(name), Some(args)) = - (call["name"].as_str(), call["parameters"].as_object()) - { - handler.tool_call(ToolCall::new( - name.to_string(), - json!(args), - None, - ))?; + "tool-plan-delta" => { + if let Some(text) = data["delta"]["message"]["tool_plan"].as_str() { + handler.text(text)?; + } + } + "tool-call-start" => { + if let (Some(function), Some(id)) = ( + data["delta"]["message"]["tool_calls"]["function"].as_object(), + data["delta"]["message"]["tool_calls"]["id"].as_str(), + ) { + if let Some(name) = function.get("name").and_then(|v| v.as_str()) { + function_name = name.to_string(); } + function_id = id.to_string(); + } + } + "tool-call-delta" => { + if let Some(text) = + data["delta"]["message"]["tool_calls"]["function"]["arguments"].as_str() + { + function_arguments.push_str(text); + } + } + "tool-call-end" => { + if !function_name.is_empty() { + let arguments: Value = function_arguments.parse().with_context(|| { + format!("Tool call '{function_name}' have non-JSON arguments '{function_arguments}'") + })?; + handler.tool_call(ToolCall::new( + function_name.clone(), + arguments, + Some(function_id.clone()), + ))?; } + function_name.clear(); + function_arguments.clear(); + function_id.clear(); } + _ => {} } - Ok(()) - }; - json_stream(res.bytes_stream(), handle).await?; - } - Ok(()) + } + Ok(false) + }; + + sse_stream(builder, handle).await } async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result<EmbeddingsOutput> { @@ -163,173 +193,59 @@ async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result<Embedding } let res_body: EmbeddingsResBody = serde_json::from_value(data).context("Invalid embeddings data")?; - Ok(res_body.embeddings) + Ok(res_body.embeddings.float) } #[derive(Deserialize)] struct EmbeddingsResBody { - embeddings: Vec<Vec<f32>>, + embeddings: EmbeddingsResBodyEmbeddings, } -fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> { - let ChatCompletionsData { - mut messages, - temperature, - top_p, - functions, - stream, - } = data; - - let system_message = extract_system_message(&mut messages); - - let mut image_urls = vec![]; - let mut tool_results = None; - - let mut messages: Vec<Value> = messages - .into_iter() - .filter_map(|message| { - let Message { role, content } = message; - let role = match role { - MessageRole::User => "USER", - _ => "CHATBOT", - }; - match content { - MessageContent::Text(text) => Some(json!({ - "role": role, - "message": text, - })), - MessageContent::Array(list) => { - let list: Vec<String> = list - .into_iter() - .filter_map(|item| match item { - MessageContentPart::Text { text } => Some(text), - MessageContentPart::ImageUrl { - image_url: ImageUrl { url }, - } => { - image_urls.push(url.clone()); - None - } - }) - .collect(); - Some(json!({ "role": role, "message": list.join("\n\n") })) - } - MessageContent::ToolCalls(tool_calls) => { - tool_results = Some(tool_calls.tool_results); - None - } - } - }) - .collect(); - - if !image_urls.is_empty() { - bail!("The model does not support images: {:?}", image_urls); - } - let message = messages.pop().unwrap(); - let message = message["message"].as_str().unwrap_or_default(); - - let mut body = json!({ - "model": &model.name(), - "message": message, - }); - - if let Some(v) = system_message { - body["preamble"] = v.into(); - } - - if !messages.is_empty() { - body["chat_history"] = messages.into(); - } - - if let Some(v) = model.max_tokens_param() { - body["max_tokens"] = v.into(); - } - if let Some(v) = temperature { - body["temperature"] = v.into(); - } - if let Some(v) = top_p { - body["p"] = v.into(); - } - if stream { - body["stream"] = true.into(); - } - - if let Some(tool_results) = tool_results { - let tool_results: Vec<_> = tool_results - .into_iter() - .map(|tool_result| { - json!({ - "call": { - "name": tool_result.call.name, - "parameters": tool_result.call.arguments, - }, - "outputs": [ - tool_result.output, - ] - - }) - }) - .collect(); - body["tool_results"] = json!(tool_results); - if let Some(object) = body.as_object_mut() { - object.remove("chat_history"); - object.remove("message"); - } - } - - if let Some(functions) = functions { - body["tools"] = functions - .iter() - .map(|v| { - let required = v.parameters.required.clone().unwrap_or_default(); - let mut parameter_definitions = json!({}); - if let Some(properties) = &v.parameters.properties { - for (key, value) in properties { - let mut value: Value = json!(value); - if value.is_object() && required.iter().any(|x| x == key) { - value["required"] = true.into(); - } - parameter_definitions[key] = value; - } - } - json!({ - "name": v.name, - "description": v.description, - "parameter_definitions": parameter_definitions, - }) - }) - .collect(); - } - Ok(body) +#[derive(Deserialize)] +struct EmbeddingsResBodyEmbeddings { + float: Vec<Vec<f32>>, } fn extract_chat_completions(data: &Value) -> Result<ChatCompletionsOutput> { - let text = data["text"].as_str().unwrap_or_default(); + let mut text = data["message"]["content"][0]["text"] + .as_str() + .unwrap_or_default() + .to_string(); let mut tool_calls = vec![]; - if let Some(calls) = data["tool_calls"].as_array() { - tool_calls = calls - .iter() - .filter_map(|call| { - if let (Some(name), Some(parameters)) = - (call["name"].as_str(), call["parameters"].as_object()) - { - Some(ToolCall::new(name.to_string(), json!(parameters), None)) - } else { - None - } - }) - .collect() + if let Some(calls) = data["message"]["tool_calls"].as_array() { + if text.is_empty() { + if let Some(tool_plain) = data["message"]["tool_plan"].as_str() { + text = tool_plain.to_string(); + } + } + for call in calls { + if let (Some(name), Some(arguments), Some(id)) = ( + call["function"]["name"].as_str(), + call["function"]["arguments"].as_str(), + call["id"].as_str(), + ) { + let arguments: Value = arguments.parse().with_context(|| { + format!("Tool call '{name}' have non-JSON arguments '{arguments}'") + })?; + tool_calls.push(ToolCall::new( + name.to_string(), + arguments, + Some(id.to_string()), + )); + } + } } if text.is_empty() && tool_calls.is_empty() { bail!("Invalid response data: {data}"); } let output = ChatCompletionsOutput { - text: text.to_string(), + text, tool_calls, - id: data["generation_id"].as_str().map(|v| v.to_string()), - input_tokens: data["meta"]["billed_units"]["input_tokens"].as_u64(), - output_tokens: data["meta"]["billed_units"]["output_tokens"].as_u64(), + id: data["id"].as_str().map(|v| v.to_string()), + input_tokens: data["usage"]["billed_units"]["input_tokens"].as_u64(), + output_tokens: data["usage"]["billed_units"]["output_tokens"].as_u64(), }; Ok(output) } |
