summaryrefslogtreecommitdiffstats
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
parent68882ecd4dced38e92b6ef581a30e45358fc61e0 (diff)
downloadaichat-34041a976c0977d51123e6d5fefd944e222e2918.tar.gz
feat: support cloudflare client (#459)
-rwxr-xr-xArgcfile.sh27
-rw-r--r--config.example.yaml35
-rw-r--r--models.yaml23
-rw-r--r--src/client/cloudflare.rs115
-rw-r--r--src/client/common.rs4
-rw-r--r--src/client/mod.rs13
6 files changed, 190 insertions, 27 deletions
diff --git a/Argcfile.sh b/Argcfile.sh
index b830e28..7ba83d3 100755
--- a/Argcfile.sh
+++ b/Argcfile.sh
@@ -309,6 +309,22 @@ chat-bedrock() {
cat "$file"
}
+# @cmd Chat with cloudflare api
+# @env CLOUDFLARE_API_KEY!
+# @option -m --model=@cf/meta/llama-3-8b-instruct $CLOUDFLARE_MODEL
+# @flag -S --no-stream
+# @arg text~
+chat-cloudflare() {
+ url="https://api.cloudflare.com/client/v4/accounts/$CLOUDFLARE_ACCOUNT_ID/ai/run/$argc_model"
+ _wrapper curl -i $CLOUDFLARE_CURL_ARGS "$url" \
+-X POST \
+-H "Authorization: Bearer $CLOUDFLARE_API_KEY" \
+-d '{
+ "messages": '"$(_build_msg $*)"',
+ "stream": '$stream'
+}'
+}
+
# @cmd Chat with ernie api
# @meta require-tools jq
# @env ERNIE_API_KEY!
@@ -316,8 +332,10 @@ chat-bedrock() {
# @flag -S --no-stream
# @arg text~
chat-ernie() {
- ACCESS_TOKEN="$(curl -fsSL "https://aip.baidubce.com/oauth/2.0/token?grant_type=client_credentials&client_id=$ERNIE_API_KEY&client_secret=$ERNIE_SECRET_KEY" | jq -r '.access_token')"
- _wrapper curl -i $ERNIE_CURL_ARGS "https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/$argc_model?access_token=$ACCESS_TOKEN" \
+ auth_url="https://aip.baidubce.com/oauth/2.0/token?grant_type=client_credentials&client_id=$ERNIE_API_KEY&client_secret=$ERNIE_SECRET_KEY"
+ ACCESS_TOKEN="$(curl -fsSL "$auth_url" | jq -r '.access_token')"
+ url="https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/$argc_model?access_token=$ACCESS_TOKEN"
+ _wrapper curl -i $ERNIE_CURL_ARGS "$url" \
-X POST \
-d '{
"messages": '"$(_build_msg $*)"',
@@ -338,9 +356,8 @@ chat-qianwen() {
stream_args=""
parameters_args='{}'
fi
- parameters_args='{ "temperature": 0.5 }'
-
- _wrapper curl -i $QIANWEN_CURL_ARGS 'https://dashscope.aliyuncs.com/api/v1/services/aigc/text-generation/generation' \
+ url=https://dashscope.aliyuncs.com/api/v1/services/aigc/text-generation/generation
+ _wrapper curl -i $QIANWEN_CURL_ARGS "$url" \
-X POST \
-H "Authorization: Bearer $QIANWEN_API_KEY" \
-H 'Content-Type: application/json' $stream_args \
diff --git a/config.example.yaml b/config.example.yaml
index b45514e..f2a8de5 100644
--- a/config.example.yaml
+++ b/config.example.yaml
@@ -73,21 +73,6 @@ clients:
- type: groq
api_key: gsk_xxx # ENV: {client_name}_API_KEY
- # Any platform that is compatible with OpenAI's API can be used here, including:
- # - anyscale: https://docs.anyscale.com/endpoints/model-serving/openai-migration-guide
- # - deepinfra: https://deepinfra.com/docs/advanced/openai_api
- # - fireworks: https://readme.fireworks.ai/docs/openai-compatibility
- # - together: https://docs.together.ai/docs/openai-api-compatibility
- # - huggingface: https://huggingface.co/docs/text-generation-inference/messages_api
- - type: openai-compatible
- name: localai
- api_base: http://localhost:8080/v1
- api_key: sk-xxx # ENV: {client_name}_API_BASE
- chat_endpoint: /chat/completions # Optional
- models: # Required
- - name: llama3
- max_input_tokens: 8192
-
# See https://github.com/jmorganca/ollama
- type: ollama
api_base: http://localhost:11434
@@ -116,6 +101,10 @@ clients:
# Optional field, possible values: BLOCK_NONE, BLOCK_ONLY_HIGH, BLOCK_MEDIUM_AND_ABOVE, BLOCK_LOW_AND_ABOVE
block_threshold: BLOCK_ONLY_HIGH
+ - type: cloudflare
+ account_id: xxx # ENV: {client_name}_ACCOUNT_ID
+ api_key: xxx # ENV: {client_name}_API_BASE
+
# See https://docs.aws.amazon.com/bedrock/latest/userguide/
- type: bedrock
access_key_id: xxx # ENV: {client_name}_ACCESS_KEY_ID
@@ -134,3 +123,19 @@ clients:
# See https://platform.moonshot.cn/docs/intro
- type: moonshot
api_key: sk-xxx # ENV: {client_name}_API_KEY
+
+ # Any platform that is compatible with OpenAI's API can be used here, including:
+ # - localai: https://github.com/mudler/LocalAI
+ # - anyscale: https://docs.anyscale.com/endpoints/model-serving/openai-migration-guide
+ # - deepinfra: https://deepinfra.com/docs/advanced/openai_api
+ # - fireworks: https://readme.fireworks.ai/docs/openai-compatibility
+ # - together: https://docs.together.ai/docs/openai-api-compatibility
+ # - huggingface: https://huggingface.co/docs/text-generation-inference/messages_api
+ - type: openai-compatible
+ name: localai
+ api_base: http://localhost:8080/v1
+ api_key: sk-xxx # ENV: {client_name}_API_BASE
+ chat_endpoint: /chat/completions # Optional
+ models: # Required
+ - name: llama3
+ max_input_tokens: 8192 \ No newline at end of file
diff --git a/models.yaml b/models.yaml
index 5ef72dd..29b499f 100644
--- a/models.yaml
+++ b/models.yaml
@@ -318,8 +318,29 @@
input_price: 8
output_price: 2.4
+- type: cloudflare
+ # docs:
+ # - https://developers.cloudflare.com/workers-ai/models/
+ # - https://developers.cloudflare.com/workers-ai/platform/pricing/
+ models:
+ - name: '@cf/meta/llama-2-7b-chat-fp16'
+ max_input_tokens: 3072
+ max_output_tokens: 2500
+ input_price: 0.56
+ output_price: 6.6
+ - name: '@cf/meta/llama-2-7b-chat-int8'
+ max_input_tokens: 2048
+ max_output_tokens: 1800
+ input_price: 0.16
+ output_price: 0.24
+ - name: '@cf/mistral/mistral-7b-instruct-v0.1'
+ max_input_tokens: 8192
+ max_output_tokens: 8192
+ input_price: 0.11
+ output_price: 0.19
+
- type: ernie
- # docs:
+ # docs:
# - https://cloud.baidu.com/doc/WENXINWORKSHOP/s/Nlks5zkzu
# - https://cloud.baidu.com/doc/WENXINWORKSHOP/s/hlrk4akp7
models:
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
+ ),
);