summaryrefslogtreecommitdiffstats
path: root/src/client/model.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-03 06:52:57 +0800
committerGitHub <noreply@github.com>2023-11-03 06:52:57 +0800
commitf9c40e52dabda7b037805c0635b84ccb6d75f5a8 (patch)
tree73fda138c432ed18c5f95b7be571f8298f6284c8 /src/client/model.rs
parentdce6877f5de297803a737387fef7f672593904cf (diff)
downloadaichat-f9c40e52dabda7b037805c0635b84ccb6d75f5a8.tar.gz
refactor: improve code quanity (#203)
- update field name of ModelInfo - rename ModelInfo to Model
Diffstat (limited to 'src/client/model.rs')
-rw-r--r--src/client/model.rs80
1 files changed, 80 insertions, 0 deletions
diff --git a/src/client/model.rs b/src/client/model.rs
new file mode 100644
index 0000000..d00ad46
--- /dev/null
+++ b/src/client/model.rs
@@ -0,0 +1,80 @@
+use super::message::Message;
+
+use crate::utils::count_tokens;
+
+use anyhow::{bail, Result};
+
+pub type TokensCountFactors = (usize, usize); // (per-messages, bias)
+
+#[derive(Debug, Clone)]
+pub struct Model {
+ pub client_index: usize,
+ pub client_name: String,
+ pub llm_name: String,
+ pub max_tokens: Option<usize>,
+ pub tokens_count_factors: TokensCountFactors,
+}
+
+impl Default for Model {
+ fn default() -> Self {
+ Model::new(0, "", "")
+ }
+}
+
+impl Model {
+ pub fn new(client_index: usize, client_name: &str, name: &str) -> Self {
+ Self {
+ client_index,
+ client_name: client_name.into(),
+ llm_name: name.into(),
+ max_tokens: None,
+ tokens_count_factors: Default::default(),
+ }
+ }
+
+ pub fn id(&self) -> String {
+ format!("{}:{}", self.client_name, self.llm_name)
+ }
+
+ pub fn set_max_tokens(mut self, max_tokens: Option<usize>) -> Self {
+ match max_tokens {
+ None | Some(0) => self.max_tokens = None,
+ _ => self.max_tokens = max_tokens,
+ }
+ self
+ }
+
+ pub fn set_tokens_count_factors(mut self, tokens_count_factors: TokensCountFactors) -> Self {
+ self.tokens_count_factors = tokens_count_factors;
+ self
+ }
+
+ pub fn messages_tokens(&self, messages: &[Message]) -> usize {
+ messages.iter().map(|v| count_tokens(&v.content)).sum()
+ }
+
+ pub fn total_tokens(&self, messages: &[Message]) -> usize {
+ if messages.is_empty() {
+ return 0;
+ }
+ let num_messages = messages.len();
+ let message_tokens = self.messages_tokens(messages);
+ let (per_messages, _) = self.tokens_count_factors;
+ if messages[num_messages - 1].role.is_user() {
+ num_messages * per_messages + message_tokens
+ } else {
+ (num_messages - 1) * per_messages + message_tokens
+ }
+ }
+
+ pub fn max_tokens_limit(&self, messages: &[Message]) -> Result<()> {
+ let (_, bias) = self.tokens_count_factors;
+ let total_tokens = self.total_tokens(messages) + bias;
+ if let Some(max_tokens) = self.max_tokens {
+ if total_tokens >= max_tokens {
+ bail!("Exceed max tokens limit")
+ }
+ }
+ Ok(())
+ }
+}