summaryrefslogtreecommitdiffstats
path: root/src/client/cloudflare.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-29 09:27:11 +0800
committerGitHub <noreply@github.com>2024-04-29 09:27:11 +0800
commit34041a976c0977d51123e6d5fefd944e222e2918 (patch)
tree1186205314499e911a702289c1dafb898479ed3d /src/client/cloudflare.rs
parent68882ecd4dced38e92b6ef581a30e45358fc61e0 (diff)
downloadaichat-34041a976c0977d51123e6d5fefd944e222e2918.tar.gz
feat: support cloudflare client (#459)
Diffstat (limited to 'src/client/cloudflare.rs')
-rw-r--r--src/client/cloudflare.rs115
1 files changed, 115 insertions, 0 deletions
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<String>,
+ pub account_id: Option<String>,
+ pub api_key: Option<String>,
+ #[serde(default)]
+ pub models: Vec<ModelConfig>,
+ pub extra: Option<ExtraConfig>,
+}
+
+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<RequestBuilder> {
+ 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<bool> {
+ 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<Value> {
+ 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()))
+}