diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-14 11:16:55 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-14 11:16:55 +0800 |
| commit | 5284a18248bb8e48eaa4a1e6ddcf73d944d23783 (patch) | |
| tree | fe60e7b0be7c4ea40051f40f88362e0487da6355 /src/client | |
| parent | 154c1e0b4b7fad08094c601893b081581d2ee0c8 (diff) | |
| download | aichat-5284a18248bb8e48eaa4a1e6ddcf73d944d23783.tar.gz | |
refactor: config::Input (#503)
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/common.rs | 8 | ||||
| -rw-r--r-- | src/client/message.rs | 18 |
2 files changed, 4 insertions, 22 deletions
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<Message>) -> Option<String> { } 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\"}" - ); - } -} |
