diff options
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/cloudflare.rs | 115 | ||||
| -rw-r--r-- | src/client/common.rs | 4 | ||||
| -rw-r--r-- | src/client/mod.rs | 13 |
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 + ), ); |
