diff options
| author | sigoden <sigoden@gmail.com> | 2023-11-04 09:51:01 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-11-04 09:51:01 +0800 |
| commit | 7627771e2ddc3f3c86f691ebfa84e31d8518a22d (patch) | |
| tree | f4a7e86635d8ed4019027fa39a013385b1cf5b3d /src | |
| parent | ec1f99fbf5c5c4f49a204a368daf04a68f69f088 (diff) | |
| download | aichat-7627771e2ddc3f3c86f691ebfa84e31d8518a22d.tar.gz | |
feat: support Qianwen (#211)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/mod.rs | 1 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 160 |
2 files changed, 161 insertions, 0 deletions
diff --git a/src/client/mod.rs b/src/client/mod.rs index ab58056..be29d79 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -18,4 +18,5 @@ register_client!( ), (palm, "palm", PaLMConfig, PaLMClient), (ernie, "ernie", ErnieConfig, ErnieClient), + (qianwen, "qianwen", QianwenConfig, QianwenClient), ); diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs new file mode 100644 index 0000000..d5145c2 --- /dev/null +++ b/src/client/qianwen.rs @@ -0,0 +1,160 @@ +use super::{QianwenClient, Client, ExtraConfig, PromptType, SendData, Model}; + +use crate::{ + config::GlobalConfig, + render::ReplyHandler, + utils::PromptKind, +}; + +use anyhow::{anyhow, bail, Result}; +use async_trait::async_trait; +use futures_util::StreamExt; +use reqwest::{Client as ReqwestClient, RequestBuilder}; +use reqwest_eventsource::{Error as EventSourceError, Event, RequestBuilderExt}; +use serde::Deserialize; +use serde_json::{json, Value}; + +const API_URL: &str = + "https://dashscope.aliyuncs.com/api/v1/services/aigc/text-generation/generation"; + +const MODELS: [(&str, usize); 2] = [("qwen-turbo", 6144), ("qwen-plus", 6144)]; + +#[derive(Debug, Clone, Deserialize, Default)] +pub struct QianwenConfig { + pub name: Option<String>, + pub api_key: Option<String>, + pub extra: Option<ExtraConfig>, +} + +#[async_trait] +impl Client for QianwenClient { + fn config(&self) -> (&GlobalConfig, &Option<ExtraConfig>) { + (&self.global_config, &self.config.extra) + } + + async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> { + let builder = self.request_builder(client, data)?; + send_message(builder).await + } + + async fn send_message_streaming_inner( + &self, + client: &ReqwestClient, + handler: &mut ReplyHandler, + data: SendData, + ) -> Result<()> { + let builder = self.request_builder(client, data)?; + send_message_streaming(builder, handler).await + } +} + +impl QianwenClient { + config_get_fn!(api_key, get_api_key); + + pub const PROMPTS: [PromptType<'static>; 1] = + [("api_key", "API Key:", true, PromptKind::String)]; + + pub fn list_models(local_config: &QianwenConfig, client_index: usize) -> Vec<Model> { + let client_name = Self::name(local_config); + MODELS + .into_iter() + .map(|(name, max_tokens)| Model::new(client_index, client_name, name).set_max_tokens(Some(max_tokens))) + .collect() + } + + fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { + let api_key = self.get_api_key()?; + + let stream = data.stream; + let body = build_body(data, self.model.llm_name.clone()); + + let mut builder = client.post(API_URL).bearer_auth(api_key).json(&body); + if stream { + builder = builder.header("X-DashScope-SSE", "enable"); + } + + Ok(builder) + } +} + +async fn send_message(builder: RequestBuilder) -> Result<String> { + let data: Value = builder.send().await?.json().await?; + check_error(&data)?; + + let output = data["output"]["text"].as_str() + .ok_or_else(|| anyhow!("Unexpected response {data}"))?; + + Ok(output.to_string()) +} + +async fn send_message_streaming( + builder: RequestBuilder, + handler: &mut ReplyHandler, +) -> Result<()> { + let mut es = builder.eventsource()?; + let mut offset = 0; + + while let Some(event) = es.next().await { + match event { + Ok(Event::Open) => {} + Ok(Event::Message(message)) => { + let data: Value = serde_json::from_str(&message.data)?; + if let Some(text) = data["output"]["text"].as_str() { + + let text = &text[offset..]; + handler.text(text)?; + offset += text.len(); + } + } + Err(err) => { + match err { + EventSourceError::InvalidStatusCode(_, res) => { + let data: Value = res.json().await?; + check_error(&data)?; + bail!("Request failed"); + } + EventSourceError::StreamEnded => {} + _ => { + bail!("{}", err); + } + } + es.close(); + } + } + } + + Ok(()) +} + +fn check_error(data: &Value) -> Result<()> { + if let Some(code) = data["code"].as_str() { + if let Some(message) = data["message"].as_str() { + bail!("{message}"); + } else { + bail!("{code}"); + } + } + Ok(()) +} + +fn build_body(data: SendData, model: String) -> Value { + let SendData { + messages, + temperature, + stream: _, + } = data; + + let mut parameters = json!({}); + + if let Some(v) = temperature { + parameters["temperature"] = v.into(); + } + + json!({ + "model": model, + "input": json!({ + "messages": messages, + }), + "parameters": parameters + }) +} |
