summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-12-19 20:47:15 +0800
committerGitHub <noreply@github.com>2023-12-19 20:47:15 +0800
commit6fb13359f474f575fa87924884d9f46a312b6d21 (patch)
tree5d8147badd84c9be4b3e025a6bd6b31b778eb3de
parent6286251d320a7f512e49534c30b42452d30d93cc (diff)
downloadaichat-6fb13359f474f575fa87924884d9f46a312b6d21.tar.gz
feat: abandon PaLM2 (#274)
-rw-r--r--README.md1
-rw-r--r--config.example.yaml4
-rw-r--r--src/client/common.rs1
-rw-r--r--src/client/mod.rs1
-rw-r--r--src/client/palm.rs131
5 files changed, 1 insertions, 137 deletions
diff --git a/README.md b/README.md
index f896c13..cb64d4a 100644
--- a/README.md
+++ b/README.md
@@ -52,7 +52,6 @@ Download it from [GitHub Releases](https://github.com/sigoden/aichat/releases),
- [x] LocalAI: user deployed opensource LLMs
- [x] Azure-OpenAI: user created gpt3.5/gpt4
- [x] Gemini: gemini-pro/gemini-pro-vision/gemini-ultra
-- [x] PaLM: chat-bison-001
- [x] Ernie: ernie-bot-turbo/ernie-bot/ernie-bot-8k/ernie-bot-4
- [x] Qianwen: qwen-turbo/qwen-plus/qwen-max
diff --git a/config.example.yaml b/config.example.yaml
index 847beca..502969b 100644
--- a/config.example.yaml
+++ b/config.example.yaml
@@ -43,10 +43,6 @@ clients:
- type: gemini
api_key: AIxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx
- # See https://developers.generativeai.google/guide
- - type: palm
- api_key: AIxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx
-
# See https://cloud.baidu.com/doc/WENXINWORKSHOP/index.html
- type: ernie
api_key: xxxxxxxxxxxxxxxxxxxxxxxx
diff --git a/src/client/common.rs b/src/client/common.rs
index 98cffa2..9ff02ba 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -292,6 +292,7 @@ pub fn create_config(list: &[PromptType], client: &str) -> Result<Value> {
Ok(clients)
}
+#[allow(unused)]
pub async fn send_message_as_streaming<F, Fut>(
builder: RequestBuilder,
handler: &mut ReplyHandler,
diff --git a/src/client/mod.rs b/src/client/mod.rs
index dc6b093..d2ada44 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -17,7 +17,6 @@ register_client!(
AzureOpenAIClient
),
(gemini, "gemini", GeminiConfig, GeminiClient),
- (palm, "palm", PaLMConfig, PaLMClient),
(ernie, "ernie", ErnieConfig, ErnieClient),
(qianwen, "qianwen", QianwenConfig, QianwenClient),
);
diff --git a/src/client/palm.rs b/src/client/palm.rs
deleted file mode 100644
index a519768..0000000
--- a/src/client/palm.rs
+++ /dev/null
@@ -1,131 +0,0 @@
-use super::{PaLMClient, Client, ExtraConfig, Model, PromptType, SendData, TokensCountFactors, send_message_as_streaming, patch_system_message};
-
-use crate::{config::GlobalConfig, render::ReplyHandler, utils::PromptKind};
-
-use anyhow::{anyhow, bail, Result};
-use async_trait::async_trait;
-use reqwest::{Client as ReqwestClient, RequestBuilder};
-use serde::Deserialize;
-use serde_json::{json, Value};
-
-const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta2/models/";
-
-const MODELS: [(&str, usize); 1] = [("chat-bison-001", 4096)];
-
-const TOKENS_COUNT_FACTORS: TokensCountFactors = (3, 8);
-
-#[derive(Debug, Clone, Deserialize, Default)]
-pub struct PaLMConfig {
- pub name: Option<String>,
- pub api_key: Option<String>,
- pub extra: Option<ExtraConfig>,
-}
-
-#[async_trait]
-impl Client for PaLMClient {
- 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_as_streaming(builder, handler, send_message).await
- }
-}
-
-impl PaLMClient {
- 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: &PaLMConfig) -> Vec<Model> {
- 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))
- .set_tokens_count_factors(TOKENS_COUNT_FACTORS)
- })
- .collect()
- }
-
- fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
- let api_key = self.get_api_key()?;
-
- let body = build_body(data, self.model.name.clone());
-
- let model = self.model.name.clone();
-
- let url = format!("{API_BASE}{}:generateMessage?key={}", model, api_key);
-
- debug!("PaLM Request: {url} {body}");
-
- let builder = client.post(url).json(&body);
-
- Ok(builder)
- }
-}
-
-async fn send_message(builder: RequestBuilder) -> Result<String> {
- let data: Value = builder.send().await?.json().await?;
- check_error(&data)?;
-
- let output = data["candidates"][0]["content"]
- .as_str()
- .ok_or_else(|| {
- if let Some(reason) = data["filters"][0]["reason"].as_str() {
- anyhow!("Content Filtering: {reason}")
- } else {
- anyhow!("Unexpected response")
- }
- })?;
-
- Ok(output.to_string())
-}
-
-fn check_error(data: &Value) -> Result<()> {
- if let Some(error) = data["error"].as_object() {
- if let Some(message) = error["message"].as_str() {
- bail!("{message}");
- } else {
- bail!("Error {}", Value::Object(error.clone()));
- }
- }
- Ok(())
-}
-
-fn build_body(data: SendData, _model: String) -> Value {
- let SendData {
- mut messages,
- temperature,
- ..
- } = data;
-
- patch_system_message(&mut messages);
-
- let messages: Vec<Value> = messages.into_iter().map(|v| json!({ "content": v.content })).collect();
-
- let prompt = json!({ "messages": messages });
-
- let mut body = json!({
- "prompt": prompt,
- });
-
- if let Some(temperature) = temperature {
- body["temperature"] = (temperature / 2.0).into();
- }
-
- body
-}