summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-01-27 08:13:50 +0800
committerGitHub <noreply@github.com>2025-01-27 08:13:50 +0800
commit20cada3d0ecfb3a63bd24138fa1ad329c778c8d1 (patch)
treefef5fc6dc687c5802af716ffbdf68c8d021abac3
parentfb8f117abad66d26b17e2e1a66be4ff48feba3fa (diff)
downloadaichat-20cada3d0ecfb3a63bd24138fa1ad329c778c8d1.tar.gz
feat: ernie migrates to v2 api (#1130)
-rwxr-xr-xArgcfile.sh29
-rw-r--r--config.example.yaml5
-rw-r--r--models.yaml6
-rw-r--r--src/client/common.rs13
-rw-r--r--src/client/ernie.rs329
-rw-r--r--src/client/mod.rs4
-rw-r--r--src/client/openai_compatible.rs8
7 files changed, 16 insertions, 378 deletions
diff --git a/Argcfile.sh b/Argcfile.sh
index cf1e8ff..ea77211 100755
--- a/Argcfile.sh
+++ b/Argcfile.sh
@@ -294,21 +294,6 @@ chat-vertexai() {
-d "$(_build_body vertexai "$@")"
}
-# @cmd Chat with ernie api
-# @meta require-tools jq
-# @env ERNIE_API_KEY!
-# @option -m --model=ernie-tiny-8k $ERNIE_MODEL
-# @flag -S --no-stream
-# @arg text~
-chat-ernie() {
- 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 "$url" \
--X POST \
--d "$(_build_body ernie "$@")"
-}
-
_argc_before() {
OPENAI_COMPATIBLE_PROVIDERS=( \
openai,gpt-4o-mini,https://api.openai.com/v1 \
@@ -316,6 +301,7 @@ _argc_before() {
cloudflare,@cf/meta/llama-3.1-8b-instruct,https://api.cloudflare.com/client/v4/accounts/${CLOUDFLARE_ACCOUNT_ID}/ai/v1 \
deepinfra,meta-llama/Meta-Llama-3.1-8B-Instruct,https://api.deepinfra.com/v1/openai \
deepseek,deepseek-chat,https://api.deepseek.com \
+ ernie,ernie-4.0-turbo-8k-latest,https://qianfan.baidubce.com/v2 \
fireworks,accounts/fireworks/models/llama-v3p1-8b-instruct,https://api.fireworks.ai/inference/v1 \
github,gpt-4o-mini,https://models.inference.ai.azure.com \
groq,llama-3.1-8b-instant,https://api.groq.com/openai/v1 \
@@ -380,7 +366,7 @@ _choice_provider() {
}
_choice_client() {
- printf "%s\n" gemini claude cohere azure-openai vertexai bedrock ernie
+ printf "%s\n" gemini claude cohere azure-openai vertexai bedrock
}
_choice_openai_compatible_provider() {
@@ -442,17 +428,6 @@ _build_body() {
"safetySettings":[{"category":"HARM_CATEGORY_HARASSMENT","threshold":"BLOCK_ONLY_HIGH"},{"category":"HARM_CATEGORY_HATE_SPEECH","threshold":"BLOCK_ONLY_HIGH"},{"category":"HARM_CATEGORY_SEXUALLY_EXPLICIT","threshold":"BLOCK_ONLY_HIGH"},{"category":"HARM_CATEGORY_DANGEROUS_CONTENT","threshold":"BLOCK_ONLY_HIGH"}]
}'
;;
- ernie)
- echo '{
- "messages": [
- {
- "role": "user",
- "content": "'"$*"'"
- }
- ],
- "stream": '$stream'
-}'
- ;;
*)
_die "error: unsupported build body for $kind"
;;
diff --git a/config.example.yaml b/config.example.yaml
index fe41a61..8e81f8a 100644
--- a/config.example.yaml
+++ b/config.example.yaml
@@ -246,9 +246,10 @@ clients:
api_key: xxx
# See https://cloud.baidu.com/doc/WENXINWORKSHOP/index.html
- - type: ernie
+ - type: openai-compatible
+ name: ernie
+ api_base: https://qianfan.baidubce.com/v2
api_key: xxx
- secret_key: xxx
# See https://dashscope.aliyun.com/
- type: openai-compatible
diff --git a/models.yaml b/models.yaml
index a6003aa..4816586 100644
--- a/models.yaml
+++ b/models.yaml
@@ -746,19 +746,19 @@
max_input_tokens: 128000
input_price: 0.042
output_price: 0.084
- - name: bge_large_zh
+ - name: bge-large-zh
type: embedding
input_price: 0.07
max_tokens_per_chunk: 512
default_chunk_size: 1000
max_batch_size: 16
- - name: bge_large_en
+ - name: bge-large-en
type: embedding
input_price: 0.07
max_tokens_per_chunk: 512
default_chunk_size: 1000
max_batch_size: 16
- - name: bce_reranker_base
+ - name: bce-reranker-base
type: reranker
max_input_tokens: 1024
input_price: 0.07
diff --git a/src/client/common.rs b/src/client/common.rs
index 80f585d..322fc5c 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -525,19 +525,6 @@ pub fn json_str_from_map<'a>(
map.get(field_name).and_then(|v| v.as_str())
}
-pub fn maybe_catch_error(data: &Value) -> Result<()> {
- if let (Some(code), Some(message)) = (data["code"].as_str(), data["message"].as_str()) {
- debug!("Invalid response: {}", data);
- bail!("{message} (code: {code})");
- } else if let (Some(error_code), Some(error_msg)) =
- (data["error_code"].as_number(), data["error_msg"].as_str())
- {
- debug!("Invalid response: {}", data);
- bail!("{error_msg} (error_code: {error_code})");
- }
- Ok(())
-}
-
fn set_client_config(list: &[PromptAction], client_config: &mut Value, client: &str) -> Result<()> {
for (key, desc, help_message) in list {
let env_name = format!("{client}_{key}").to_ascii_uppercase();
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
deleted file mode 100644
index 9a6b962..0000000
--- a/src/client/ernie.rs
+++ /dev/null
@@ -1,329 +0,0 @@
-use super::access_token::*;
-use super::openai_compatible::*;
-use super::*;
-
-use anyhow::{anyhow, bail, Context, Result};
-use reqwest::{Client as ReqwestClient, RequestBuilder};
-use serde::Deserialize;
-use serde_json::{json, Value};
-
-const API_BASE: &str = "https://aip.baidubce.com/rpc/2.0/ai_custom/v1";
-const ACCESS_TOKEN_URL: &str = "https://aip.baidubce.com/oauth/2.0/token";
-
-#[derive(Debug, Clone, Deserialize, Default)]
-pub struct ErnieConfig {
- pub name: Option<String>,
- pub api_key: Option<String>,
- pub secret_key: Option<String>,
- #[serde(default)]
- pub models: Vec<ModelData>,
- pub patch: Option<RequestPatch>,
- pub extra: Option<ExtraConfig>,
-}
-
-impl ErnieClient {
- config_get_fn!(api_key, get_api_key);
- config_get_fn!(secret_key, get_secret_key);
- pub const PROMPTS: [PromptAction<'static>; 2] = [
- ("api_key", "API Key", None),
- ("secret_key", "Secret Key", None),
- ];
-}
-
-#[async_trait::async_trait]
-impl Client for ErnieClient {
- client_common_fns!();
-
- async fn chat_completions_inner(
- &self,
- client: &ReqwestClient,
- data: ChatCompletionsData,
- ) -> Result<ChatCompletionsOutput> {
- prepare_access_token(self, client).await?;
- let request_data = prepare_chat_completions(self, data)?;
- let builder = self.request_builder(client, request_data);
- chat_completions(builder, &self.model).await
- }
-
- async fn chat_completions_streaming_inner(
- &self,
- client: &ReqwestClient,
- handler: &mut SseHandler,
- data: ChatCompletionsData,
- ) -> Result<()> {
- prepare_access_token(self, client).await?;
- let request_data = prepare_chat_completions(self, data)?;
- let builder = self.request_builder(client, request_data);
- chat_completions_streaming(builder, handler, &self.model).await
- }
-
- async fn embeddings_inner(
- &self,
- client: &ReqwestClient,
- data: &EmbeddingsData,
- ) -> Result<EmbeddingsOutput> {
- prepare_access_token(self, client).await?;
- let request_data = prepare_embeddings(self, data)?;
- let builder = self.request_builder(client, request_data);
- embeddings(builder, &self.model).await
- }
-
- async fn rerank_inner(
- &self,
- client: &ReqwestClient,
- data: &RerankData,
- ) -> Result<RerankOutput> {
- prepare_access_token(self, client).await?;
- let request_data = prepare_rerank(self, data)?;
- let builder = self.request_builder(client, request_data);
- rerank(builder, &self.model).await
- }
-}
-
-fn prepare_chat_completions(self_: &ErnieClient, data: ChatCompletionsData) -> Result<RequestData> {
- let access_token = get_access_token(self_.name())?;
-
- let url = format!(
- "{API_BASE}/wenxinworkshop/chat/{}?access_token={access_token}",
- self_.model.name(),
- );
-
- let body = build_chat_completions_body(data, &self_.model);
-
- let request_data = RequestData::new(url, body);
-
- Ok(request_data)
-}
-
-fn prepare_embeddings(self_: &ErnieClient, data: &EmbeddingsData) -> Result<RequestData> {
- let access_token = get_access_token(self_.name())?;
-
- let url = format!(
- "{API_BASE}/wenxinworkshop/embeddings/{}?access_token={access_token}",
- self_.model.name(),
- );
-
- let body = json!({
- "input": data.texts,
- });
-
- let request_data = RequestData::new(url, body);
-
- Ok(request_data)
-}
-
-fn prepare_rerank(self_: &ErnieClient, data: &RerankData) -> Result<RequestData> {
- let access_token = get_access_token(self_.name())?;
-
- let url = format!(
- "{API_BASE}/wenxinworkshop/reranker/{}?access_token={access_token}",
- self_.model.name(),
- );
-
- let RerankData {
- query,
- documents,
- top_n,
- } = data;
-
- let body = json!({
- "query": query,
- "documents": documents,
- "top_n": top_n
- });
-
- let request_data = RequestData::new(url, body);
-
- Ok(request_data)
-}
-
-async fn prepare_access_token(self_: &ErnieClient, client: &ReqwestClient) -> Result<()> {
- let client_name = self_.name();
- if !is_valid_access_token(client_name) {
- let api_key = self_.get_api_key()?;
- let secret_key = self_.get_secret_key()?;
-
- let token = fetch_access_token(client, &api_key, &secret_key)
- .await
- .with_context(|| "Failed to fetch access token")?;
- set_access_token(client_name, token, 86400);
- }
- Ok(())
-}
-
-async fn chat_completions(
- builder: RequestBuilder,
- _model: &Model,
-) -> Result<ChatCompletionsOutput> {
- let data: Value = builder.send().await?.json().await?;
- maybe_catch_error(&data)?;
- debug!("non-stream-data: {data}");
- extract_chat_completions_text(&data)
-}
-
-async fn chat_completions_streaming(
- builder: RequestBuilder,
- handler: &mut SseHandler,
- _model: &Model,
-) -> Result<()> {
- let handle = |message: SseMmessage| -> Result<bool> {
- let data: Value = serde_json::from_str(&message.data)?;
- debug!("stream-data: {data}");
- if let Some(function) = data["function_call"].as_object() {
- if let (Some(name), Some(arguments)) = (
- function.get("name").and_then(|v| v.as_str()),
- function.get("arguments").and_then(|v| v.as_str()),
- ) {
- let arguments: Value = arguments.parse().with_context(|| {
- format!("Tool call '{name}' have non-JSON arguments '{arguments}'")
- })?;
- handler.tool_call(ToolCall::new(name.to_string(), arguments, None))?;
- }
- } else if let Some(text) = data["result"].as_str() {
- handler.text(text)?;
- }
- Ok(false)
- };
-
- sse_stream(builder, handle).await
-}
-
-async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result<EmbeddingsOutput> {
- let data: Value = builder.send().await?.json().await?;
- maybe_catch_error(&data)?;
- let res_body: EmbeddingsResBody =
- serde_json::from_value(data).context("Invalid embeddings data")?;
- let output = res_body.data.into_iter().map(|v| v.embedding).collect();
- Ok(output)
-}
-
-#[derive(Deserialize)]
-struct EmbeddingsResBody {
- data: Vec<EmbeddingsResBodyEmbedding>,
-}
-
-#[derive(Deserialize)]
-struct EmbeddingsResBodyEmbedding {
- embedding: Vec<f32>,
-}
-
-async fn rerank(builder: RequestBuilder, _model: &Model) -> Result<RerankOutput> {
- let data: Value = builder.send().await?.json().await?;
- maybe_catch_error(&data)?;
- let res_body: GenericRerankResBody =
- serde_json::from_value(data).context("Invalid rerank data")?;
- Ok(res_body.results)
-}
-
-fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Value {
- let ChatCompletionsData {
- mut messages,
- temperature,
- top_p,
- functions,
- stream,
- } = data;
-
- let system_message = extract_system_message(&mut messages);
-
- let messages: Vec<Value> = messages
- .into_iter()
- .flat_map(|message| {
- let Message { role, content } = message;
- match content {
- MessageContent::ToolCalls(MessageContentToolCalls {
- tool_results, ..
- }) => {
- let mut list = vec![];
- for tool_result in tool_results {
- list.push(json!({
- "role": "assistant",
- "content": format!("Action: {}\nAction Input: {}", tool_result.call.name, tool_result.call.arguments)
- }));
- list.push(json!({
- "role": "user",
- "content": tool_result.output.to_string(),
- }))
-
- }
- list
- }
- _ => vec![json!({ "role": role, "content": content })],
- }
- })
- .collect();
-
- let mut body = json!({
- "messages": messages,
- });
-
- if let Some(v) = system_message {
- body["system"] = v.into();
- }
-
- if let Some(v) = model.max_tokens_param() {
- body["max_output_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();
- }
-
- if let Some(functions) = functions {
- body["functions"] = json!(functions);
- }
-
- body
-}
-
-fn extract_chat_completions_text(data: &Value) -> Result<ChatCompletionsOutput> {
- let text = data["result"].as_str().unwrap_or_default();
-
- let mut tool_calls = vec![];
- if let Some(call) = data["function_call"].as_object() {
- if let (Some(name), Some(arguments)) = (
- call.get("name").and_then(|v| v.as_str()),
- call.get("arguments").and_then(|v| v.as_str()),
- ) {
- let arguments: Value = arguments.parse().with_context(|| {
- format!("Tool call '{name}' have non-JSON arguments '{arguments}'")
- })?;
- tool_calls.push(ToolCall::new(name.to_string(), arguments, None));
- }
- }
-
- if text.is_empty() && tool_calls.is_empty() {
- bail!("Invalid response data: {data}");
- }
- let output = ChatCompletionsOutput {
- text: text.to_string(),
- tool_calls,
- id: data["id"].as_str().map(|v| v.to_string()),
- input_tokens: data["usage"]["prompt_tokens"].as_u64(),
- output_tokens: data["usage"]["completion_tokens"].as_u64(),
- };
- Ok(output)
-}
-
-async fn fetch_access_token(
- client: &reqwest::Client,
- api_key: &str,
- secret_key: &str,
-) -> Result<String> {
- let url = format!("{ACCESS_TOKEN_URL}?grant_type=client_credentials&client_id={api_key}&client_secret={secret_key}");
- let value: Value = client.get(&url).send().await?.json().await?;
- let result = value["access_token"].as_str().ok_or_else(|| {
- if let Some(err_msg) = value["error_description"].as_str() {
- anyhow!("{err_msg}")
- } else {
- anyhow!("Invalid response data")
- }
- })?;
- Ok(result.to_string())
-}
diff --git a/src/client/mod.rs b/src/client/mod.rs
index d8b41c6..c87d2b5 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -31,10 +31,9 @@ register_client!(
),
(vertexai, "vertexai", VertexAIConfig, VertexAIClient),
(bedrock, "bedrock", BedrockConfig, BedrockClient),
- (ernie, "ernie", ErnieConfig, ErnieClient),
);
-pub const OPENAI_COMPATIBLE_PROVIDERS: [(&str, &str); 23] = [
+pub const OPENAI_COMPATIBLE_PROVIDERS: [(&str, &str); 24] = [
("ai21", "https://api.ai21.com/studio/v1"),
(
"cloudflare",
@@ -42,6 +41,7 @@ pub const OPENAI_COMPATIBLE_PROVIDERS: [(&str, &str); 23] = [
),
("deepinfra", "https://api.deepinfra.com/v1/openai"),
("deepseek", "https://api.deepseek.com"),
+ ("ernie", "https://qianfan.baidubce.com/v2"),
("fireworks", "https://api.fireworks.ai/inference/v1"),
("github", "https://models.inference.ai.azure.com"),
("groq", "https://api.groq.com/openai/v1"),
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index ce1eea3..30bc464 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -79,7 +79,11 @@ fn prepare_rerank(self_: &OpenAICompatibleClient, data: &RerankData) -> Result<R
let api_key = self_.get_api_key().ok();
let api_base = get_api_base_ext(self_)?;
- let url = format!("{api_base}/rerank");
+ let url = if self_.name().starts_with("ernie") {
+ format!("{api_base}/rerankers")
+ } else {
+ format!("{api_base}/rerank")
+ };
let body = generic_build_rerank_body(data, &self_.model);
@@ -149,7 +153,7 @@ pub fn generic_build_rerank_body(data: &RerankData, model: &Model) -> Value {
"query": query,
"documents": documents,
});
- if model.client_name() == "voyageai" {
+ if model.client_name().starts_with("voyageai") {
body["top_k"] = (*top_n).into()
} else {
body["top_n"] = (*top_n).into()