summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-27 16:06:35 +0800
committerGitHub <noreply@github.com>2023-11-27 16:06:35 +0800
commit18f16c6511e20d11ea1ae4b4d81b43d4d35486b3 (patch)
treeb2e3022c39c7db4927f1a35eecc1193f46c953e8 /src
parent2508d56598a37844e369ab623ceb5bc7b2c78d38 (diff)
downloadaichat-18f16c6511e20d11ea1ae4b4d81b43d4d35486b3.tar.gz
feat: add ernie:ernie-bot-8k qianwen:qwen-max (#252)
Diffstat (limited to 'src')
-rw-r--r--src/client/ernie.rs5
-rw-r--r--src/client/qianwen.rs26
2 files changed, 16 insertions, 15 deletions
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 4bb3435..75a362a 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -18,9 +18,10 @@ use std::env;
const API_BASE: &str = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1";
const ACCESS_TOKEN_URL: &str = "https://aip.baidubce.com/oauth/2.0/token";
-const MODELS: [(&str, &str); 3] = [
- ("eb-instant", "/wenxinworkshop/chat/eb-instant"),
+const MODELS: [(&str, &str); 4] = [
+ ("ernie-bot-turbo", "/wenxinworkshop/chat/eb-instant"),
("ernie-bot", "/wenxinworkshop/chat/completions"),
+ ("ernie-bot-8k", "/wenxinworkshop/chat/ernie_bot_8k"),
("ernie-bot-4", "/wenxinworkshop/chat/completions_pro"),
];
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 1ee5f3d..619ed76 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -1,10 +1,6 @@
-use super::{QianwenClient, Client, ExtraConfig, PromptType, SendData, Model};
+use super::{Client, ExtraConfig, Model, PromptType, QianwenClient, SendData};
-use crate::{
- config::GlobalConfig,
- render::ReplyHandler,
- utils::PromptKind,
-};
+use crate::{config::GlobalConfig, render::ReplyHandler, utils::PromptKind};
use anyhow::{anyhow, bail, Result};
use async_trait::async_trait;
@@ -17,7 +13,11 @@ 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)];
+const MODELS: [(&str, usize); 3] = [
+ ("qwen-turbo", 6144),
+ ("qwen-plus", 6144),
+ ("qwen-max", 6144),
+];
#[derive(Debug, Clone, Deserialize, Default)]
pub struct QianwenConfig {
@@ -58,7 +58,9 @@ impl QianwenClient {
let client_name = Self::name(local_config);
MODELS
.into_iter()
- .map(|(name, max_tokens)| Model::new(client_name, name).set_max_tokens(Some(max_tokens)))
+ .map(|(name, max_tokens)| {
+ Model::new(client_name, name).set_max_tokens(Some(max_tokens))
+ })
.collect()
}
@@ -83,16 +85,14 @@ 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()
+ 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<()> {
+async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> {
let mut es = builder.eventsource()?;
while let Some(event) = es.next().await {