From 573e0d58b44cd0686c9e7405723e5cc3a5c5126f Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 1 Sep 2024 08:27:08 +0800 Subject: feat: migrate `ollama`/`qianwen` clients to `openai-compatible` (#816) --- src/client/qianwen.rs | 509 -------------------------------------------------- 1 file changed, 509 deletions(-) delete mode 100644 src/client/qianwen.rs (limited to 'src/client/qianwen.rs') diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs deleted file mode 100644 index 38534d8..0000000 --- a/src/client/qianwen.rs +++ /dev/null @@ -1,509 +0,0 @@ -use super::*; - -use crate::utils::{base64_decode, sha256}; - -use anyhow::{anyhow, bail, Context, Result}; -use reqwest::{ - multipart::{Form, Part}, - Client as ReqwestClient, RequestBuilder, -}; -use serde::Deserialize; -use serde_json::{json, Value}; -use std::borrow::BorrowMut; - -const API_BASE: &str = "https://dashscope.aliyuncs.com/api/v1"; - -const CHAT_COMPLETIONS_ENDPOINT: &str = "/services/aigc/text-generation/generation"; - -const CHAT_COMPLETIONS_VL_ENDPOINT: &str = "/services/aigc/multimodal-generation/generation"; - -const EMBEDDINGS_ENDPOINT: &str = "/services/embeddings/text-embedding/text-embedding"; - -#[derive(Debug, Clone, Deserialize, Default)] -pub struct QianwenConfig { - pub name: Option, - pub api_key: Option, - pub api_base: Option, - #[serde(default)] - pub models: Vec, - pub patch: Option, - pub extra: Option, -} - -impl QianwenClient { - config_get_fn!(api_key, get_api_key); - config_get_fn!(api_base, get_api_base); - - pub const PROMPTS: [PromptAction<'static>; 1] = - [("api_key", "API Key:", true, PromptKind::String)]; -} - -#[async_trait::async_trait] -impl Client for QianwenClient { - client_common_fns!(); - - async fn chat_completions_inner( - &self, - client: &ReqwestClient, - mut data: ChatCompletionsData, - ) -> Result { - let api_key = self.get_api_key()?; - patch_messages(self.model.name(), &api_key, &mut data.messages).await?; - let request_data = prepare_chat_completions(self, data)?; - let builder = self.request_builder(client, request_data, ApiType::ChatCompletions); - chat_completions(builder, &self.model).await - } - - async fn chat_completions_streaming_inner( - &self, - client: &ReqwestClient, - handler: &mut SseHandler, - mut data: ChatCompletionsData, - ) -> Result<()> { - let api_key = self.get_api_key()?; - patch_messages(self.model.name(), &api_key, &mut data.messages).await?; - let request_data = prepare_chat_completions(self, data)?; - let builder = self.request_builder(client, request_data, ApiType::ChatCompletions); - chat_completions_streaming(builder, handler, &self.model).await - } - - async fn embeddings_inner( - &self, - client: &ReqwestClient, - data: EmbeddingsData, - ) -> Result>> { - let request_data = prepare_embeddings(self, data)?; - let builder = self.request_builder(client, request_data, ApiType::Embeddings); - embeddings(builder, &self.model).await - } -} - -fn prepare_chat_completions( - self_: &QianwenClient, - data: ChatCompletionsData, -) -> Result { - let api_key = self_.get_api_key()?; - let api_base = self_ - .get_api_base() - .unwrap_or_else(|_| API_BASE.to_string()); - - let stream = data.stream; - - let url = match self_.model().supports_vision() { - true => format!( - "{}{CHAT_COMPLETIONS_VL_ENDPOINT}", - api_base.trim_end_matches('/'), - ), - false => format!( - "{}{CHAT_COMPLETIONS_ENDPOINT}", - api_base.trim_end_matches('/'), - ), - }; - - let (body, has_upload) = build_chat_completions_body(data, &self_.model)?; - - let mut request_data = RequestData::new(url, body); - - request_data.bearer_auth(api_key); - - if stream { - request_data.header("X-DashScope-SSE", "enable"); - } - if has_upload { - request_data.header("X-DashScope-OssResourceResolve", "enable"); - } - - Ok(request_data) -} - -fn prepare_embeddings(self_: &QianwenClient, data: EmbeddingsData) -> Result { - let api_key = self_.get_api_key()?; - let api_base = self_ - .get_api_base() - .unwrap_or_else(|_| API_BASE.to_string()); - - let url = format!("{}{EMBEDDINGS_ENDPOINT}", api_base.trim_end_matches('/'),); - - let text_type = match data.query { - true => "query", - false => "document", - }; - - let body = json!({ - "model": self_.model.name(), - "input": { - "texts": data.texts, - }, - "parameters": { - "text_type": text_type, - } - }); - - let mut request_data = RequestData::new(url, body); - - request_data.bearer_auth(api_key); - - Ok(request_data) -} - -async fn chat_completions(builder: RequestBuilder, model: &Model) -> Result { - let data: Value = builder.send().await?.json().await?; - maybe_catch_error(&data)?; - - debug!("non-stream-data: {data}"); - extract_chat_completions_text(&data, model) -} - -async fn chat_completions_streaming( - builder: RequestBuilder, - handler: &mut SseHandler, - model: &Model, -) -> Result<()> { - let model_name = model.name(); - let mut prev_text = String::new(); - let handle = |message: SseMmessage| -> Result { - let data: Value = serde_json::from_str(&message.data)?; - maybe_catch_error(&data)?; - debug!("stream-data: {data}"); - if model_name == "qwen-long" { - if let Some(text) = data["output"]["choices"][0]["message"]["content"].as_str() { - handler.text(text)?; - } - } else if model.supports_vision() { - if let Some(text) = - data["output"]["choices"][0]["message"]["content"][0]["text"].as_str() - { - handler.text(text)?; - } - } else if let Some(text) = data["output"]["text"].as_str() { - if let Some(pos) = text.rfind("✿FUNCTION") { - if pos > prev_text.len() { - let delta_text = &text[prev_text.len()..pos]; - if delta_text != ": \n" { - handler.text(delta_text)?; - } - } - prev_text = text.to_string(); - if let Some((name, arguments)) = parse_tool_call(&text[pos..]) { - let arguments: Value = arguments - .parse() - .with_context(|| format!("Invalid function call {name} {arguments}"))?; - handler.tool_call(ToolCall::new(name.to_string(), arguments, None))?; - } - } else { - let mut delta_text = &text[prev_text.len()..]; - if prev_text.is_empty() && delta_text.starts_with(": ") { - delta_text = &delta_text[2..]; - } - prev_text = text.to_string(); - handler.text(delta_text)?; - } - } - Ok(false) - }; - - sse_stream(builder, handle).await -} - -fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<(Value, bool)> { - let ChatCompletionsData { - messages, - temperature, - top_p, - functions, - stream, - } = data; - - let mut has_upload = false; - let input = if model.supports_vision() { - let messages: Vec = messages - .into_iter() - .map(|message| { - let role = message.role; - let content = match message.content { - MessageContent::Text(text) => vec![json!({"text": text})], - MessageContent::Array(list) => list - .into_iter() - .map(|item| match item { - MessageContentPart::Text { text } => json!({"text": text}), - MessageContentPart::ImageUrl { - image_url: ImageUrl { url }, - } => { - if url.starts_with("oss:") { - has_upload = true; - } - json!({"image": url}) - } - }) - .collect(), - MessageContent::ToolResults(_) => { - vec![] - } - }; - json!({ "role": role, "content": content }) - }) - .collect(); - - json!({ - "messages": messages, - }) - } else { - let messages: Vec = - messages - .into_iter() - .flat_map(|message| { - let role = message.role; - match message.content { - MessageContent::Text(text) => vec![json!({ "role": role, "content": text })], - MessageContent::Array(list) => { - let parts: Vec<_> = list - .into_iter() - .map(|item| match item { - MessageContentPart::Text { text } => json!({"text": text}), - MessageContentPart::ImageUrl { - image_url: ImageUrl { url }, - } => { - if url.starts_with("oss:") { - has_upload = true; - } - json!({"image": url}) - } - }) - .collect(); - vec![json!({ "role": role, "content": parts })] - } - MessageContent::ToolResults((tool_results, _)) => { - tool_results.into_iter().flat_map(|tool_result| vec![ - json!({ - "role": MessageRole::Assistant, - "content": "", - "tool_calls": vec![ - json!({ - "type": "function", - "function": { - "name": tool_result.call.name, - "arguments": tool_result.call.arguments.to_string(), - }, - }) - ], - }), - json!({ - "role": "tool", - "content": tool_result.output.to_string(), - "name": tool_result.call.name, - }), - ]).collect() - } - } - }) - .collect(); - json!({ - "messages": messages, - }) - }; - - let mut parameters = json!({}); - - if stream && (model.name() == "qwen-long" || model.supports_vision()) { - parameters["incremental_output"] = true.into(); - } - - if let Some(v) = model.max_tokens_param() { - parameters["max_tokens"] = v.into(); - } - if let Some(v) = temperature { - parameters["temperature"] = v.into(); - } - if let Some(v) = top_p { - parameters["top_p"] = v.into(); - } - - if let Some(functions) = functions { - parameters["tools"] = functions - .iter() - .map(|v| { - json!({ - "type": "function", - "function": v, - }) - }) - .collect(); - } - - let body = json!({ - "model": &model.name(), - "input": input, - "parameters": parameters - }); - - Ok((body, has_upload)) -} - -async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result { - let data: Value = builder.send().await?.json().await?; - maybe_catch_error(&data)?; - let res_body: EmbeddingsResBody = - serde_json::from_value(data).context("Invalid embeddings data")?; - let output = res_body - .output - .embeddings - .into_iter() - .map(|v| v.embedding) - .collect(); - Ok(output) -} - -#[derive(Deserialize)] -struct EmbeddingsResBody { - output: EmbeddingsResBodyOutput, -} - -#[derive(Deserialize)] -struct EmbeddingsResBodyOutput { - embeddings: Vec, -} - -#[derive(Deserialize)] -struct EmbeddingsResBodyOutputEmbedding { - embedding: Vec, -} - -fn extract_chat_completions_text(data: &Value, model: &Model) -> Result { - let err = || anyhow!("Invalid response data: {data}"); - let mut tool_calls = vec![]; - let text = if model.name() == "qwen-long" { - data["output"]["choices"][0]["message"]["content"] - .as_str() - .ok_or_else(err)? - } else if model.supports_vision() { - data["output"]["choices"][0]["message"]["content"][0]["text"] - .as_str() - .ok_or_else(err)? - } else { - let text = data["output"]["text"].as_str().ok_or_else(err)?; - match parse_tool_call(text) { - Some((name, arguments)) => { - let arguments: Value = arguments - .parse() - .with_context(|| format!("Invalid function call {name} {arguments}"))?; - tool_calls.push(ToolCall::new(name.to_string(), arguments, None)); - "" - } - None => text, - } - }; - let output = ChatCompletionsOutput { - text: text.to_string(), - tool_calls, - id: data["request_id"].as_str().map(|v| v.to_string()), - input_tokens: data["usage"]["input_tokens"].as_u64(), - output_tokens: data["usage"]["output_tokens"].as_u64(), - }; - - Ok(output) -} - -/// Patch messages, upload embedded images to oss -async fn patch_messages(model: &str, api_key: &str, messages: &mut Vec) -> Result<()> { - for message in messages { - if let MessageContent::Array(list) = message.content.borrow_mut() { - for item in list { - if let MessageContentPart::ImageUrl { - image_url: ImageUrl { url }, - } = item - { - if url.starts_with("data:") { - *url = upload(model, api_key, url) - .await - .with_context(|| "Failed to upload embedded image to oss")?; - } - } - } - } - } - Ok(()) -} - -#[derive(Debug, Deserialize)] -struct Policy { - data: PolicyData, -} - -#[derive(Debug, Deserialize)] -struct PolicyData { - policy: String, - signature: String, - upload_dir: String, - upload_host: String, - oss_access_key_id: String, - x_oss_object_acl: String, - x_oss_forbid_overwrite: String, -} - -/// Upload image to dashscope -async fn upload(model: &str, api_key: &str, url: &str) -> Result { - let (mime_type, data) = url - .strip_prefix("data:") - .and_then(|v| v.split_once(";base64,")) - .ok_or_else(|| anyhow!("Invalid image url"))?; - let mut name = sha256(data); - if let Some(ext) = mime_type.strip_prefix("image/") { - name.push('.'); - name.push_str(ext); - } - let data = base64_decode(data)?; - - let client = reqwest::Client::new(); - let policy: Policy = client - .get(format!( - "https://dashscope.aliyuncs.com/api/v1/uploads?action=getPolicy&model={model}" - )) - .header("Authorization", format!("Bearer {api_key}")) - .send() - .await? - .json() - .await?; - let PolicyData { - policy, - signature, - upload_dir, - upload_host, - oss_access_key_id, - x_oss_object_acl, - x_oss_forbid_overwrite, - .. - } = policy.data; - - let key = format!("{upload_dir}/{name}"); - let file = Part::bytes(data).file_name(name).mime_str(mime_type)?; - let form = Form::new() - .text("OSSAccessKeyId", oss_access_key_id) - .text("Signature", signature) - .text("policy", policy) - .text("key", key.clone()) - .text("x-oss-object-acl", x_oss_object_acl) - .text("x-oss-forbid-overwrite", x_oss_forbid_overwrite) - .text("success_action_status", "200") - .text("x-oss-content-type", mime_type.to_string()) - .part("file", file); - - let res = client.post(upload_host).multipart(form).send().await?; - - let status = res.status(); - if !status.is_success() { - let text = res.text().await?; - bail!("Invalid response data: {text} (status: {status})") - } - Ok(format!("oss://{key}")) -} - -fn parse_tool_call(text: &str) -> Option<(&str, &str)> { - let function_symbol = "✿FUNCTION✿: "; - let result_symbol = "\n✿RESULT✿: "; - let args_symbol = "\n✿ARGS✿: "; - let start = text.find(function_symbol)? + function_symbol.len(); - let text = &text[start..]; - let end = text.find(result_symbol)?; - let text = &text[..end]; - text.split_once(args_symbol) -} -- cgit v1.2.3