summaryrefslogtreecommitdiffstats
path: root/src/client
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
parent68882ecd4dced38e92b6ef581a30e45358fc61e0 (diff)
downloadaichat-34041a976c0977d51123e6d5fefd944e222e2918.tar.gz
feat: support cloudflare client (#459)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/cloudflare.rs115
-rw-r--r--src/client/common.rs4
-rw-r--r--src/client/mod.rs13
3 files changed, 126 insertions, 6 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()))
+}
diff --git a/src/client/common.rs b/src/client/common.rs
index e35e956..842741c 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -506,6 +506,10 @@ pub fn catch_error(data: &Value, status: u16) -> Result<()> {
if let (Some(typ), Some(message)) = (error["type"].as_str(), error["message"].as_str()) {
bail!("{message} (type: {typ})");
}
+ } else if let Some(error) = data["errors"][0].as_object() {
+ if let (Some(code), Some(message)) = (error["code"].as_u64(), error["message"].as_str()) {
+ bail!("{message} (status: {code})")
+ }
} else if let Some(error) = data[0]["error"].as_object() {
if let (Some(status), Some(message)) = (error["status"].as_str(), error["message"].as_str())
{
diff --git a/src/client/mod.rs b/src/client/mod.rs
index 0916801..a311efe 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -19,12 +19,6 @@ register_client!(
(cohere, "cohere", CohereConfig, CohereClient),
(perplexity, "perplexity", PerplexityConfig, PerplexityClient),
(groq, "groq", GroqConfig, GroqClient),
- (
- openai_compatible,
- "openai-compatible",
- OpenAICompatibleConfig,
- OpenAICompatibleClient
- ),
(ollama, "ollama", OllamaConfig, OllamaClient),
(
azure_openai,
@@ -34,7 +28,14 @@ register_client!(
),
(vertexai, "vertexai", VertexAIConfig, VertexAIClient),
(bedrock, "bedrock", BedrockConfig, BedrockClient),
+ (cloudflare, "cloudflare", CloudflareConfig, CloudflareClient),
(ernie, "ernie", ErnieConfig, ErnieClient),
(qianwen, "qianwen", QianwenConfig, QianwenClient),
(moonshot, "moonshot", MoonshotConfig, MoonshotClient),
+ (
+ openai_compatible,
+ "openai-compatible",
+ OpenAICompatibleConfig,
+ OpenAICompatibleClient
+ ),
);