diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-29 06:51:03 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-29 06:51:03 +0800 |
| commit | 865be2bf75bb62b6aeee059f684400b4b9938a15 (patch) | |
| tree | b6296ccdde2af49251b02087b9f51c5239cc4d2b /src/client/common.rs | |
| parent | b33e2da75efab4246f9f1d74a0be56aba1d8bd07 (diff) | |
| download | aichat-865be2bf75bb62b6aeee059f684400b4b9938a15.tar.gz | |
feat: non-streaming returns completion stats (#456)
Diffstat (limited to 'src/client/common.rs')
| -rw-r--r-- | src/client/common.rs | 21 |
1 files changed, 16 insertions, 5 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index 695245a..a58c962 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -126,7 +126,7 @@ macro_rules! register_client { client.set_model(model); } else { anyhow::bail!( - "The current model lacks the corresponding capability." + "The current model is incapable of doing that." ); } } @@ -260,7 +260,7 @@ macro_rules! impl_client_trait { &self, client: &reqwest::Client, data: $crate::client::SendData, - ) -> anyhow::Result<String> { + ) -> anyhow::Result<(String, $crate::client::CompletionStats)> { let builder = self.request_builder(client, data)?; $send_message(builder).await } @@ -330,11 +330,11 @@ pub trait Client: Sync + Send { Ok(client) } - async fn send_message(&self, input: Input) -> Result<String> { + async fn send_message(&self, input: Input) -> Result<(String, CompletionStats)> { let global_config = self.config().0; if global_config.read().dry_run { let content = global_config.read().echo_messages(&input); - return Ok(content); + return Ok((content, CompletionStats::default())); } let client = self.build_client()?; let data = global_config.read().prepare_send_data(&input, false)?; @@ -384,7 +384,11 @@ pub trait Client: Sync + Send { } } - async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String>; + async fn send_message_inner( + &self, + client: &ReqwestClient, + data: SendData, + ) -> Result<(String, CompletionStats)>; async fn send_message_streaming_inner( &self, @@ -414,6 +418,13 @@ pub struct SendData { pub stream: bool, } +#[derive(Debug, Clone, Default)] +pub struct CompletionStats { + pub id: Option<String>, + pub input_tokens: Option<u64>, + pub output_tokens: Option<u64>, +} + pub type PromptType<'a> = (&'a str, &'a str, bool, PromptKind); pub fn create_config(list: &[PromptType], client: &str) -> Result<(String, Value)> { |
