From bba5028615a1c81ee513096eff3e5f2e71eed5a5 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 4 Nov 2023 08:11:51 +0800 Subject: feat: support PaLM (#209) --- src/client/common.rs | 20 +++++++- src/client/mod.rs | 1 + src/client/palm.rs | 135 +++++++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 154 insertions(+), 2 deletions(-) create mode 100644 src/client/palm.rs (limited to 'src/client') diff --git a/src/client/common.rs b/src/client/common.rs index d43f1b6..7959b46 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -11,10 +11,10 @@ use crate::{ use anyhow::{Context, Result}; use async_trait::async_trait; -use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy}; +use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; -use std::{env, time::Duration}; +use std::{env, future::Future, time::Duration}; use tokio::time::sleep; #[macro_export] @@ -298,6 +298,22 @@ pub fn create_config(list: &[PromptType], client: &str) -> Result { Ok(clients) } +pub async fn send_message_as_streaming( + builder: RequestBuilder, + handler: &mut ReplyHandler, + f: F, +) -> Result<()> +where + F: FnOnce(RequestBuilder) -> Fut, + Fut: Future>, +{ + let text = f(builder).await?; + handler.text(&text)?; + handler.done()?; + + Ok(()) +} + fn set_config_value(json: &mut Value, path: &str, kind: &PromptKind, value: &str) { let segs: Vec<&str> = path.split('.').collect(); match segs.as_slice() { diff --git a/src/client/mod.rs b/src/client/mod.rs index f124b62..55c0739 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -16,4 +16,5 @@ register_client!( AzureOpenAIConfig, AzureOpenAIClient ), + (palm, "palm", PaLMConfig, PaLMClient), ); diff --git a/src/client/palm.rs b/src/client/palm.rs new file mode 100644 index 0000000..87b60d1 --- /dev/null +++ b/src/client/palm.rs @@ -0,0 +1,135 @@ +use super::{PaLMClient, Client, ExtraConfig, Model, PromptType, SendData, TokensCountFactors, send_message_as_streaming}; + +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"; + +const MODELS: [(&str, usize, &str); 1] = [("chat-bison-001", 4096, "/models/chat-bison-001")]; + +const TOKENS_COUNT_FACTORS: TokensCountFactors = (3, 8); + +#[derive(Debug, Clone, Deserialize, Default)] +pub struct PaLMConfig { + pub name: Option, + pub api_key: Option, + pub extra: Option, +} + +#[async_trait] +impl Client for PaLMClient { + fn config(&self) -> (&GlobalConfig, &Option) { + (&self.global_config, &self.config.extra) + } + + async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result { + 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, client_index: usize) -> Vec { + 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)) + .set_tokens_count_factors(TOKENS_COUNT_FACTORS) + }) + .collect() + } + + fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { + let api_key = self.get_api_key()?; + + let body = build_body(data, self.model.llm_name.clone()); + + let model = self.model.llm_name.clone(); + let (_, _, endpoint) = MODELS + .iter() + .find(|(v, _, _)| v == &model) + .ok_or_else(|| anyhow!("Miss Model '{}' in {}", model, self.model.client_name))?; + + let url = format!("{API_BASE}{endpoint}:generateMessage?key={}", api_key); + + let builder = client.post(url).json(&body); + + Ok(builder) + } +} + +async fn send_message(builder: RequestBuilder) -> Result { + let data: Value = builder.send().await?.json().await?; + check_error(&data)?; + + let output = data["candidates"][0]["content"] + .as_str() + .ok_or_else(|| anyhow!("Unexpected response {data}"))?; + + 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!("Request failed. {}", data); + } + } + Ok(()) +} + +fn build_body(data: SendData, _model: String) -> Value { + let SendData { + mut messages, + temperature, + .. + } = data; + + let mut context = None; + if messages[0].role.is_system() { + let message = messages.remove(0); + context = Some(message.content); + } + + let messages: Vec = messages.into_iter().map(|v| json!({ "content": v.content })).collect(); + + let mut prompt = json!({ "messages": messages }); + + if let Some(context) = context { + prompt["context"] = context.into(); + }; + + let mut body = json!({ + "prompt": prompt, + }); + + if let Some(temperature) = temperature { + body["temperature"] = (temperature / 2.0).into(); + } + + body +} -- cgit v1.2.3