summaryrefslogtreecommitdiffstats
path: root/src/client/cohere.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-12-06 18:14:43 +0800
committerGitHub <noreply@github.com>2024-12-06 18:14:43 +0800
commitb164a1587deff738d04a75665bdb863d82b3264c (patch)
tree556e94d58e378246a736c638705abd84e9b40e81 /src/client/cohere.rs
parent5f7edc0ded187d86712639f3f7c48e6e4d72927b (diff)
downloadaichat-b164a1587deff738d04a75665bdb863d82b3264c.tar.gz
feat: update cohere api to v2 (#1041)
Diffstat (limited to 'src/client/cohere.rs')
-rw-r--r--src/client/cohere.rs274
1 files changed, 95 insertions, 179 deletions
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)
}