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 /src/config | |
| parent | 28f019b72c1027207ec8dc82844401512712464d (diff) | |
| download | aichat-19ce62ab769b74a8e417cd1213b7711a5ed55638.tar.gz | |
feat: check token usage in dry_run mode (#82)
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/message.rs | 9 | ||||
| -rw-r--r-- | src/config/mod.rs | 17 |
2 files changed, 15 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)?; |
