From 2cba09c064b366ad91c7eda3acbc2030215a1842 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 2 Sep 2024 08:23:21 +0800 Subject: feat: migrate `cloudflare` client to `openai-compatible` (#821) --- src/client/cloudflare.rs | 181 ----------------------------------------------- src/client/common.rs | 17 +++-- src/client/gemini.rs | 4 +- src/client/mod.rs | 4 +- 4 files changed, 14 insertions(+), 192 deletions(-) delete mode 100644 src/client/cloudflare.rs (limited to 'src') diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs deleted file mode 100644 index 3626c73..0000000 --- a/src/client/cloudflare.rs +++ /dev/null @@ -1,181 +0,0 @@ -use super::*; - -use anyhow::{anyhow, Context, Result}; -use reqwest::RequestBuilder; -use serde::Deserialize; -use serde_json::{json, Value}; - -const API_BASE: &str = "https://api.cloudflare.com/client/v4"; - -#[derive(Debug, Clone, Deserialize, Default)] -pub struct CloudflareConfig { - pub name: Option, - pub account_id: Option, - pub api_base: Option, - pub api_key: Option, - #[serde(default)] - pub models: Vec, - pub patch: Option, - pub extra: Option, -} - -impl CloudflareClient { - config_get_fn!(account_id, get_account_id); - config_get_fn!(api_key, get_api_key); - config_get_fn!(api_base, get_api_base); - - pub const PROMPTS: [PromptAction<'static>; 2] = [ - ("account_id", "Account ID:", true, PromptKind::String), - ("api_key", "API Key:", true, PromptKind::String), - ]; -} - -impl_client_trait!( - CloudflareClient, - ( - prepare_chat_completions, - chat_completions, - chat_completions_streaming - ), - (prepare_embeddings, embeddings), - (noop_prepare_rerank, noop_rerank), -); - -fn prepare_chat_completions( - self_: &CloudflareClient, - data: ChatCompletionsData, -) -> Result { - let account_id = self_.get_account_id()?; - let api_key = self_.get_api_key()?; - let api_base = self_ - .get_api_base() - .unwrap_or_else(|_| API_BASE.to_string()); - - let url = format!( - "{}/accounts/{account_id}/ai/run/{}", - api_base.trim_end_matches('/'), - self_.model.name() - ); - - let body = build_chat_completions_body(data, &self_.model)?; - - let mut request_data = RequestData::new(url, body); - - request_data.bearer_auth(api_key); - - Ok(request_data) -} - -fn prepare_embeddings(self_: &CloudflareClient, data: EmbeddingsData) -> Result { - let account_id = self_.get_account_id()?; - let api_key = self_.get_api_key()?; - - let url = format!( - "{API_BASE}/accounts/{account_id}/ai/run/{}", - self_.model.name() - ); - - let body = json!({ - "text": data.texts, - }); - - let mut request_data = RequestData::new(url, body); - - request_data.bearer_auth(api_key); - - Ok(request_data) -} - -async fn chat_completions( - builder: RequestBuilder, - _model: &Model, -) -> Result { - let res = builder.send().await?; - let status = res.status(); - let data: Value = res.json().await?; - if !status.is_success() { - catch_error(&data, status.as_u16())?; - } - - debug!("non-stream-data: {data}"); - extract_chat_completions(&data) -} - -async fn chat_completions_streaming( - builder: RequestBuilder, - handler: &mut SseHandler, - _model: &Model, -) -> Result<()> { - let handle = |message: SseMmessage| -> Result { - if message.data == "[DONE]" { - return Ok(true); - } - let data: Value = serde_json::from_str(&message.data)?; - debug!("stream-data: {data}"); - if let Some(text) = data["response"].as_str() { - handler.text(text)?; - } - Ok(false) - }; - sse_stream(builder, handle).await -} - -async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result { - let res = builder.send().await?; - let status = res.status(); - let data: Value = res.json().await?; - if !status.is_success() { - catch_error(&data, status.as_u16())?; - } - let res_body: EmbeddingsResBody = - serde_json::from_value(data).context("Invalid embeddings data")?; - Ok(res_body.result.data) -} - -#[derive(Deserialize)] -struct EmbeddingsResBody { - result: EmbeddingsResBodyResult, -} - -#[derive(Deserialize)] -struct EmbeddingsResBodyResult { - data: Vec>, -} - -fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result { - let ChatCompletionsData { - messages, - temperature, - top_p, - functions: _, - stream, - } = data; - - let mut body = json!({ - "model": &model.name(), - "messages": messages, - }); - - if let Some(v) = model.max_tokens_param() { - body["max_tokens"] = v.into(); - } - if let Some(v) = temperature { - body["temperature"] = v.into(); - } - if let Some(v) = top_p { - body["top_p"] = v.into(); - } - if stream { - body["stream"] = true.into(); - } - - Ok(body) -} - -fn extract_chat_completions(data: &Value) -> Result { - let text = data["result"]["response"] - .as_str() - .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - - Ok(ChatCompletionsOutput::new(text)) -} diff --git a/src/client/common.rs b/src/client/common.rs index 684d10e..9d6480a 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -364,7 +364,7 @@ pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, pub fn create_openai_compatible_client_config(client: &str) -> Result> { match super::OPENAI_COMPATIBLE_PLATFORMS - .iter() + .into_iter() .find(|(name, _)| client == *name) { None => Ok(None), @@ -372,13 +372,16 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result Result Result Result { +async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; diff --git a/src/client/mod.rs b/src/client/mod.rs index 8ab7c78..ecafd3c 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -33,13 +33,13 @@ register_client!( ), (vertexai, "vertexai", VertexAIConfig, VertexAIClient), (bedrock, "bedrock", BedrockConfig, BedrockClient), - (cloudflare, "cloudflare", CloudflareConfig, CloudflareClient), (replicate, "replicate", ReplicateConfig, ReplicateClient), (ernie, "ernie", ErnieConfig, ErnieClient), ); -pub const OPENAI_COMPATIBLE_PLATFORMS: [(&str, &str); 18] = [ +pub const OPENAI_COMPATIBLE_PLATFORMS: [(&str, &str); 19] = [ ("ai21", "https://api.ai21.com/studio/v1"), + ("cloudflare", ""), ("deepinfra", "https://api.deepinfra.com/v1/openai"), ("deepseek", "https://api.deepseek.com"), ("fireworks", "https://api.fireworks.ai/inference/v1"), -- cgit v1.2.3