diff options
| author | sigoden <sigoden@gmail.com> | 2024-01-13 19:52:07 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-01-13 19:52:07 +0800 |
| commit | fe35cfd9419302f01baf9672493c0b0a4b41d889 (patch) | |
| tree | 94e763745fb7989c97af39cc1dfb44250440a5eb /src/client/ernie.rs | |
| parent | 4e99df4c1bd4028a77251bdb00ff23c664372b5f (diff) | |
| download | aichat-fe35cfd9419302f01baf9672493c0b0a4b41d889.tar.gz | |
feat: supports model capabilities (#297)
1. automatically switch to the model that has the necessary capabilities.
2. throw an error if the client does not have a model with the necessary capabilities
Diffstat (limited to 'src/client/ernie.rs')
| -rw-r--r-- | src/client/ernie.rs | 32 |
1 files changed, 11 insertions, 21 deletions
diff --git a/src/client/ernie.rs b/src/client/ernie.rs index d7e3a57..7848cf8 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -1,10 +1,6 @@ -use super::{ErnieClient, Client, ExtraConfig, PromptType, SendData, Model, patch_system_message}; +use super::{patch_system_message, Client, ErnieClient, ExtraConfig, Model, PromptType, SendData}; -use crate::{ - config::GlobalConfig, - render::ReplyHandler, - utils::PromptKind, -}; +use crate::{render::ReplyHandler, utils::PromptKind}; use anyhow::{anyhow, bail, Context, Result}; use async_trait::async_trait; @@ -37,9 +33,7 @@ pub struct ErnieConfig { #[async_trait] impl Client for ErnieClient { - fn config(&self) -> (&GlobalConfig, &Option<ExtraConfig>) { - (&self.global_config, &self.config.extra) - } + client_common_fns!(); async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> { self.prepare_access_token().await?; @@ -127,10 +121,7 @@ async fn send_message(builder: RequestBuilder) -> Result<String> { Ok(output.to_string()) } -async fn send_message_streaming( - builder: RequestBuilder, - handler: &mut ReplyHandler, -) -> Result<()> { +async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> { let mut es = builder.eventsource()?; while let Some(event) = es.next().await { match event { @@ -216,13 +207,12 @@ fn build_body(data: SendData, _model: String) -> Value { async fn fetch_access_token(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 = reqwest::get(&url).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") - } - })?; + 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()) } |
