summaryrefslogtreecommitdiffstats
path: root/src/config/message.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-03-09 21:18:28 +0800
committerGitHub <noreply@github.com>2023-03-09 21:18:28 +0800
commit553d0fe55b3e55969e6ca48554e1f63c8e48f97e (patch)
treeb8f40a3b6aacd5207207c41849fb6bbf5ccc5a2f /src/config/message.rs
parent9767c07eeeb20cb2f8386faf66cb88452987988b (diff)
downloadaichat-553d0fe55b3e55969e6ca48554e1f63c8e48f97e.tar.gz
refactor: optimize counting tokens (#53)
Diffstat (limited to 'src/config/message.rs')
-rw-r--r--src/config/message.rs30
1 files changed, 23 insertions, 7 deletions
diff --git a/src/config/message.rs b/src/config/message.rs
index 0a06a73..514a2fa 100644
--- a/src/config/message.rs
+++ b/src/config/message.rs
@@ -1,6 +1,6 @@
use serde::{Deserialize, Serialize};
-pub const MESSAGE_EXTRA_TOKENS: usize = 6;
+use crate::utils::count_tokens;
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct Message {
@@ -25,10 +25,26 @@ pub enum MessageRole {
User,
}
-#[test]
-fn test_serde() {
- assert_eq!(
- serde_json::to_string(&Message::new("Hello World")).unwrap(),
- "{\"role\":\"user\",\"content\":\"Hello World\"}"
- )
+pub fn num_tokens_from_messages(messages: &[Message]) -> usize {
+ let mut num_tokens = 0;
+ for message in messages.iter() {
+ num_tokens += 4;
+ num_tokens += count_tokens(&message.content);
+ num_tokens += 1; // role always take 1 token
+ }
+ num_tokens += 2;
+ num_tokens
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ #[test]
+ fn test_serde() {
+ assert_eq!(
+ serde_json::to_string(&Message::new("Hello World")).unwrap(),
+ "{\"role\":\"user\",\"content\":\"Hello World\"}"
+ )
+ }
}