summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-03-20 22:51:51 +0800
committerGitHub <noreply@github.com>2023-03-20 22:51:51 +0800
commit19ce62ab769b74a8e417cd1213b7711a5ed55638 (patch)
treea47d48798710855b3faca84e54d4e3ffb8776c2f /src/config
parent28f019b72c1027207ec8dc82844401512712464d (diff)
downloadaichat-19ce62ab769b74a8e417cd1213b7711a5ed55638.tar.gz
feat: check token usage in dry_run mode (#82)
Diffstat (limited to 'src/config')
-rw-r--r--src/config/message.rs9
-rw-r--r--src/config/mod.rs17
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)?;