diff options
| author | sigoden <sigoden@gmail.com> | 2023-03-20 22:51:51 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-03-20 22:51:51 +0800 |
| commit | 19ce62ab769b74a8e417cd1213b7711a5ed55638 (patch) | |
| tree | a47d48798710855b3faca84e54d4e3ffb8776c2f | |
| parent | 28f019b72c1027207ec8dc82844401512712464d (diff) | |
| download | aichat-19ce62ab769b74a8e417cd1213b7711a5ed55638.tar.gz | |
feat: check token usage in dry_run mode (#82)
| -rw-r--r-- | src/config/message.rs | 9 | ||||
| -rw-r--r-- | src/config/mod.rs | 17 | ||||
| -rw-r--r-- | src/main.rs | 1 | ||||
| -rw-r--r-- | src/repl/handler.rs | 1 |
4 files changed, 17 insertions, 11 deletions
diff --git a/src/config/message.rs b/src/config/message.rs index 5717220..2c8a330 100644 --- a/src/config/message.rs +++ b/src/config/message.rs @@ -1,6 +1,5 @@ use crate::utils::count_tokens; -use anyhow::{bail, Result}; use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Deserialize, Serialize)] @@ -26,14 +25,6 @@ pub enum MessageRole { User, } -pub fn within_max_tokens_limit(messages: &[Message], max_tokens: usize) -> Result<()> { - let tokens = num_tokens_from_messages(messages); - if tokens >= max_tokens { - bail!("Exceed max tokens limit") - } - Ok(()) -} - pub fn num_tokens_from_messages(messages: &[Message]) -> usize { let mut num_tokens = 0; for message in messages.iter() { diff --git a/src/config/mod.rs b/src/config/mod.rs index 147a817..e3d81cb 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -2,10 +2,11 @@ mod conversation; mod message; mod role; +use self::conversation::Conversation; use self::message::Message; use self::role::Role; -use self::{conversation::Conversation, message::within_max_tokens_limit}; +use crate::config::message::num_tokens_from_messages; use crate::utils::now; use anyhow::{anyhow, bail, Context, Result}; @@ -292,7 +293,10 @@ impl Config { let message = Message::new(content); vec![message] }; - within_max_tokens_limit(&messages, self.model.1)?; + let tokens = num_tokens_from_messages(&messages); + if tokens >= self.model.1 { + bail!("Exceed max tokens limit") + } Ok(messages) } @@ -434,6 +438,15 @@ impl Config { (self.highlight, self.light_theme) } + pub fn maybe_print_send_tokens(&self, input: &str) { + if self.dry_run { + if let Ok(messages) = self.build_messages(input) { + let tokens = num_tokens_from_messages(&messages); + println!(">>> The following message consumes {tokens} tokens.") + } + } + } + fn open_message_file(&self) -> Result<File> { let path = Config::messages_file()?; ensure_parent_exists(&path)?; diff --git a/src/main.rs b/src/main.rs index 58fa954..f8d2eeb 100644 --- a/src/main.rs +++ b/src/main.rs @@ -91,6 +91,7 @@ fn start_directive( if !stdout().is_terminal() { config.write().highlight = false; } + config.read().maybe_print_send_tokens(input); let output = if no_stream { let (highlight, light_theme) = config.read().get_render_options(); let output = client.send_message(input)?; diff --git a/src/repl/handler.rs b/src/repl/handler.rs index 014a851..8d51439 100644 --- a/src/repl/handler.rs +++ b/src/repl/handler.rs @@ -51,6 +51,7 @@ impl ReplCmdHandler { self.reply.borrow_mut().clear(); return Ok(()); } + self.config.read().maybe_print_send_tokens(&input); let wg = WaitGroup::new(); let ret = render_stream( &input, |
