From 20cada3d0ecfb3a63bd24138fa1ad329c778c8d1 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 27 Jan 2025 08:13:50 +0800 Subject: feat: ernie migrates to v2 api (#1130) --- src/client/ernie.rs | 329 ---------------------------------------------------- 1 file changed, 329 deletions(-) delete mode 100644 src/client/ernie.rs (limited to 'src/client/ernie.rs') 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, - pub api_key: Option, - pub secret_key: Option, - #[serde(default)] - pub models: Vec, - pub patch: Option, - pub extra: Option, -} - -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 { - 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 { - 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 { - 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 { - 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 { - 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 { - 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 { - 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 { - 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 { - 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, -} - -#[derive(Deserialize)] -struct EmbeddingsResBodyEmbedding { - embedding: Vec, -} - -async fn rerank(builder: RequestBuilder, _model: &Model) -> Result { - 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 = 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 { - 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 { - 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()) -} -- cgit v1.2.3