From 34041a976c0977d51123e6d5fefd944e222e2918 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 29 Apr 2024 09:27:11 +0800 Subject: feat: support cloudflare client (#459) --- src/client/cloudflare.rs | 115 +++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 115 insertions(+) create mode 100644 src/client/cloudflare.rs (limited to 'src/client/cloudflare.rs') diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs new file mode 100644 index 0000000..09f020a --- /dev/null +++ b/src/client/cloudflare.rs @@ -0,0 +1,115 @@ +use super::{ + catch_error, sse_stream, CloudflareClient, CompletionDetails, ExtraConfig, Model, ModelConfig, + PromptType, SendData, SseHandler, +}; + +use crate::utils::PromptKind; + +use anyhow::{anyhow, Result}; +use reqwest::{Client as ReqwestClient, 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_key: Option, + #[serde(default)] + pub models: Vec, + pub extra: Option, +} + +impl CloudflareClient { + config_get_fn!(account_id, get_account_id); + config_get_fn!(api_key, get_api_key); + + pub const PROMPTS: [PromptType<'static>; 2] = [ + ("account_id", "Account ID:", false, PromptKind::String), + ("api_key", "API Key:", false, PromptKind::String), + ]; + + fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { + let account_id = self.get_account_id()?; + let api_key = self.get_api_key()?; + + let body = build_body(data, &self.model)?; + + let url = format!( + "{API_BASE}/accounts/{account_id}/ai/run/{}", + self.model.name + ); + + debug!("Cloudflare Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } +} + +impl_client_trait!(CloudflareClient, send_message, send_message_streaming); + +async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionDetails)> { + let res = builder.send().await?; + let status = res.status(); + let data: Value = res.json().await?; + if status != 200 { + catch_error(&data, status.as_u16())?; + } + + extract_completion(&data) +} + +async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandler) -> Result<()> { + let handle = |data: &str| -> Result { + if data == "[DONE]" { + return Ok(true); + } + let data: Value = serde_json::from_str(data)?; + if let Some(text) = data["response"].as_str() { + handler.text(text)?; + } + Ok(false) + }; + sse_stream(builder, handle).await +} + +fn build_body(data: SendData, model: &Model) -> Result { + let SendData { + messages, + temperature, + top_p, + stream, + } = data; + + let mut body = json!({ + "model": &model.name, + "messages": messages, + }); + + if let Some(v) = model.max_output_tokens { + 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_completion(data: &Value) -> Result<(String, CompletionDetails)> { + let text = data["result"]["response"] + .as_str() + .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; + + Ok((text.to_string(), CompletionDetails::default())) +} -- cgit v1.2.3