summaryrefslogtreecommitdiffstats
path: root/src/client/qianwen.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-09-01 08:27:08 +0800
committerGitHub <noreply@github.com>2024-09-01 08:27:08 +0800
commit573e0d58b44cd0686c9e7405723e5cc3a5c5126f (patch)
tree238266a756e26a0d3543a44bf901fc5111ebfd88 /src/client/qianwen.rs
parent55e36c7e9da2e1c93ebeaabdc8355d0a22361f03 (diff)
downloadaichat-573e0d58b44cd0686c9e7405723e5cc3a5c5126f.tar.gz
feat: migrate `ollama`/`qianwen` clients to `openai-compatible` (#816)
Diffstat (limited to 'src/client/qianwen.rs')
-rw-r--r--src/client/qianwen.rs509
1 files changed, 0 insertions, 509 deletions
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<String>,
- pub api_key: Option<String>,
- pub api_base: Option<String>,
- #[serde(default)]
- pub models: Vec<ModelData>,
- pub patch: Option<RequestPatch>,
- pub extra: Option<ExtraConfig>,
-}
-
-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<ChatCompletionsOutput> {
- 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<Vec<Vec<f32>>> {
- 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<RequestData> {
- 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<RequestData> {
- 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<ChatCompletionsOutput> {
- 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<bool> {
- 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<Value> = 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<Value> =
- 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<EmbeddingsOutput> {
- 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<EmbeddingsResBodyOutputEmbedding>,
-}
-
-#[derive(Deserialize)]
-struct EmbeddingsResBodyOutputEmbedding {
- embedding: Vec<f32>,
-}
-
-fn extract_chat_completions_text(data: &Value, model: &Model) -> Result<ChatCompletionsOutput> {
- 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<Message>) -> 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<String> {
- 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)
-}