summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/common.rs20
-rw-r--r--src/client/mod.rs1
-rw-r--r--src/client/palm.rs135
3 files changed, 154 insertions, 2 deletions
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<Value> {
Ok(clients)
}
+pub async fn send_message_as_streaming<F, Fut>(
+ builder: RequestBuilder,
+ handler: &mut ReplyHandler,
+ f: F,
+) -> Result<()>
+where
+ F: FnOnce(RequestBuilder) -> Fut,
+ Fut: Future<Output = Result<String>>,
+{
+ 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<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, 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))
+ .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.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<String> {
+ 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<Value> = 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
+}