diff options
| author | sigoden <sigoden@gmail.com> | 2023-03-09 15:30:39 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-03-09 15:30:39 +0800 |
| commit | 05d20f207fedf03c2a564438a0d44a2d85d4ed2d (patch) | |
| tree | eb7c9ae402632c5827cd4faa323b138510716f22 /src/config/conversation.rs | |
| parent | c7eb261abc32ac4d6fc5364244a6246bcdc462d9 (diff) | |
| download | aichat-05d20f207fedf03c2a564438a0d44a2d85d4ed2d.tar.gz | |
feat: add remain tokens indicator and max tokens guard (#50)
Diffstat (limited to 'src/config/conversation.rs')
| -rw-r--r-- | src/config/conversation.rs | 16 |
1 files changed, 13 insertions, 3 deletions
diff --git a/src/config/conversation.rs b/src/config/conversation.rs index ca50233..23eacb5 100644 --- a/src/config/conversation.rs +++ b/src/config/conversation.rs @@ -2,8 +2,12 @@ use anyhow::Result; use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; +use crate::utils::count_tokens; + +use super::{MAX_TOKENS, MESSAGE_EXTRA_TOKENS}; + #[derive(Debug, Clone, Deserialize, Serialize)] -pub struct Session { +pub struct Conversation { pub tokens: usize, pub messages: Vec<Message>, } @@ -14,7 +18,7 @@ pub struct Message { pub content: String, } -impl Session { +impl Conversation { pub fn new() -> Self { Self { tokens: 0, @@ -22,7 +26,7 @@ impl Session { } } - pub fn add_conversatoin(&mut self, input: &str, output: &str) -> Result<()> { + pub fn add_chat(&mut self, input: &str, output: &str) -> Result<()> { self.messages.push(Message { role: MessageRole::User, content: input.to_string(), @@ -31,6 +35,7 @@ impl Session { role: MessageRole::Assistant, content: output.to_string(), }); + self.tokens += count_tokens(input) + count_tokens(output) + 2 * MESSAGE_EXTRA_TOKENS; Ok(()) } @@ -40,6 +45,7 @@ impl Session { role: MessageRole::System, content: prompt.into(), }); + self.tokens += count_tokens(prompt) + MESSAGE_EXTRA_TOKENS; } pub fn echo_messages(&self, content: &str) -> String { @@ -59,6 +65,10 @@ impl Session { })); json!(messages) } + + pub fn reamind_tokens(&self) -> usize { + MAX_TOKENS.saturating_sub(self.tokens) + } } #[derive(Debug, Clone, Deserialize, Serialize)] |
