From 5284a18248bb8e48eaa4a1e6ddcf73d944d23783 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 14 May 2024 11:16:55 +0800 Subject: refactor: config::Input (#503) --- src/client/common.rs | 8 ++++---- src/client/message.rs | 18 ------------------ 2 files changed, 4 insertions(+), 22 deletions(-) (limited to 'src/client') diff --git a/src/client/common.rs b/src/client/common.rs index 24e6f8e..495160b 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -290,11 +290,11 @@ pub trait Client: Sync + Send { async fn send_message(&self, input: Input) -> Result<(String, CompletionDetails)> { let global_config = self.config().0; if global_config.read().dry_run { - let content = global_config.read().echo_messages(&input); + let content = input.echo_messages(); return Ok((content, CompletionDetails::default())); } let client = self.build_client()?; - let data = global_config.read().prepare_send_data(&input, false)?; + let data = input.prepare_send_data(false)?; self.send_message_inner(&client, data) .await .with_context(|| "Failed to get answer") @@ -315,7 +315,7 @@ pub trait Client: Sync + Send { ret = async { let global_config = self.config().0; if global_config.read().dry_run { - let content = global_config.read().echo_messages(&input); + let content = input.echo_messages(); let tokens = tokenize(&content); for token in tokens { tokio::time::sleep(Duration::from_millis(10)).await; @@ -324,7 +324,7 @@ pub trait Client: Sync + Send { return Ok(()); } let client = self.build_client()?; - let data = global_config.read().prepare_send_data(&input, true)?; + let data = input.prepare_send_data(true)?; self.send_message_streaming_inner(&client, handler, data).await } => { handler.done()?; diff --git a/src/client/message.rs b/src/client/message.rs index bdee750..9621811 100644 --- a/src/client/message.rs +++ b/src/client/message.rs @@ -134,21 +134,3 @@ pub fn extract_system_message(messages: &mut Vec) -> Option { } None } - -#[cfg(test)] -mod tests { - use super::*; - use crate::config::InputContext; - - #[test] - fn test_serde() { - assert_eq!( - serde_json::to_string(&Message::new(&Input::from_str( - "Hello World", - InputContext::default() - ))) - .unwrap(), - "{\"role\":\"user\",\"content\":\"Hello World\"}" - ); - } -} -- cgit v1.2.3