summaryrefslogtreecommitdiffstats
path: root/src/client/ernie.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-01-13 19:52:07 +0800
committerGitHub <noreply@github.com>2024-01-13 19:52:07 +0800
commitfe35cfd9419302f01baf9672493c0b0a4b41d889 (patch)
tree94e763745fb7989c97af39cc1dfb44250440a5eb /src/client/ernie.rs
parent4e99df4c1bd4028a77251bdb00ff23c664372b5f (diff)
downloadaichat-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.rs32
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())
}