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 ++++++++++++++++++-- 1 file changed, 18 insertions(+), 2 deletions(-) (limited to 'src/client/common.rs') 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() { -- cgit v1.2.3