summaryrefslogtreecommitdiffstats
path: root/src/client/model.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-03-06 08:35:40 +0800
committerGitHub <noreply@github.com>2024-03-06 08:35:40 +0800
commit8e5d4e55b1a5158a1f35adaf044dec545159f426 (patch)
treeb360cc4e9ad8acccea3c68b37057798cd1d657c0 /src/client/model.rs
parentbe4e5e569a61c54d8ac8fb77144e9b1d01e3b81f (diff)
downloadaichat-8e5d4e55b1a5158a1f35adaf044dec545159f426.tar.gz
refactor: rename model's `max_tokens` to `max_input_tokens` (#339)
BREAKING CHANGE: rename model's `max_tokens` to `max_input_tokens`
Diffstat (limited to 'src/client/model.rs')
-rw-r--r--src/client/model.rs22
1 files changed, 11 insertions, 11 deletions
diff --git a/src/client/model.rs b/src/client/model.rs
index b29166e..ce181a5 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -11,7 +11,7 @@ pub type TokensCountFactors = (usize, usize); // (per-messages, bias)
pub struct Model {
pub client_name: String,
pub name: String,
- pub max_tokens: Option<usize>,
+ pub max_input_tokens: Option<usize>,
pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
pub tokens_count_factors: TokensCountFactors,
pub capabilities: ModelCapabilities,
@@ -29,7 +29,7 @@ impl Model {
client_name: client_name.into(),
name: name.into(),
extra_fields: None,
- max_tokens: None,
+ max_input_tokens: None,
tokens_count_factors: Default::default(),
capabilities: ModelCapabilities::Text,
}
@@ -83,10 +83,10 @@ impl Model {
self
}
- 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,
+ pub fn set_max_input_tokens(mut self, max_input_tokens: Option<usize>) -> Self {
+ match max_input_tokens {
+ None | Some(0) => self.max_input_tokens = None,
+ _ => self.max_input_tokens = max_input_tokens,
}
self
}
@@ -122,12 +122,12 @@ impl Model {
}
}
- pub fn max_tokens_limit(&self, messages: &[Message]) -> Result<()> {
+ pub fn max_input_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")
+ if let Some(max_input_tokens) = self.max_input_tokens {
+ if total_tokens >= max_input_tokens {
+ bail!("Exceed max input tokens limit")
}
}
Ok(())
@@ -147,7 +147,7 @@ impl Model {
#[derive(Debug, Clone, Deserialize)]
pub struct ModelConfig {
pub name: String,
- pub max_tokens: Option<usize>,
+ pub max_input_tokens: Option<usize>,
pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
#[serde(deserialize_with = "deserialize_capabilities")]
#[serde(default = "default_capabilities")]