diff options
| author | sigoden <sigoden@gmail.com> | 2024-09-02 08:23:21 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-09-02 08:23:21 +0800 |
| commit | 2cba09c064b366ad91c7eda3acbc2030215a1842 (patch) | |
| tree | 467d550c12b2cab8172da7cc3af01903a5b5477a /src/client/cloudflare.rs | |
| parent | 6462a587428c38c17cbb6bce440cbf10c966eecc (diff) | |
| download | aichat-2cba09c064b366ad91c7eda3acbc2030215a1842.tar.gz | |
feat: migrate `cloudflare` client to `openai-compatible` (#821)
Diffstat (limited to 'src/client/cloudflare.rs')
| -rw-r--r-- | src/client/cloudflare.rs | 181 |
1 files changed, 0 insertions, 181 deletions
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<String>, - pub account_id: Option<String>, - pub api_base: Option<String>, - pub api_key: Option<String>, - #[serde(default)] - pub models: Vec<ModelData>, - pub patch: Option<RequestPatch>, - pub extra: Option<ExtraConfig>, -} - -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<RequestData> { - 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<RequestData> { - 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<ChatCompletionsOutput> { - 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<bool> { - 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<EmbeddingsOutput> { - 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<Vec<f32>>, -} - -fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> { - 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<ChatCompletionsOutput> { - let text = data["result"]["response"] - .as_str() - .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - - Ok(ChatCompletionsOutput::new(text)) -} |
