summaryrefslogtreecommitdiffstats
path: root/src/client/model.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-05 09:02:23 +0800
committerGitHub <noreply@github.com>2024-06-05 09:02:23 +0800
commit1ec6abfaee2fdc189b348b7e3a8145bd9a84da74 (patch)
tree6f68a860a39b0fbe87784de925b5a31e00e74e33 /src/client/model.rs
parent71f2e94579511d7524f5534377001ab3f02a9597 (diff)
downloadaichat-1ec6abfaee2fdc189b348b7e3a8145bd9a84da74.tar.gz
feat: support RAG (#560)
* feat: support RAG * support more embeddings models and implement concurrent embedding api * show the progress of addings paths * ignore embedding context when saving message * embedding model max_chunk_size => default_chunk_size * support pdf and pandoc formats (docx, epub, ipynb)
Diffstat (limited to 'src/client/model.rs')
-rw-r--r--src/client/model.rs44
1 files changed, 39 insertions, 5 deletions
diff --git a/src/client/model.rs b/src/client/model.rs
index 65e4143..e16cb4e 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -1,4 +1,7 @@
-use super::message::{Message, MessageContent};
+use super::{
+ message::{Message, MessageContent},
+ EmbeddingsData,
+};
use crate::utils::{estimate_token_length, format_option_value};
@@ -81,6 +84,10 @@ impl Model {
&self.data.name
}
+ pub fn mode(&self) -> &str {
+ &self.data.mode
+ }
+
pub fn data(&self) -> &ModelData {
&self.data
}
@@ -137,6 +144,14 @@ impl Model {
self.data.supports_function_calling
}
+ pub fn default_chunk_size(&self) -> usize {
+ self.data.default_chunk_size.unwrap_or(1000)
+ }
+
+ pub fn max_concurrent_chunks(&self) -> usize {
+ self.data.max_concurrent_chunks.unwrap_or(1)
+ }
+
pub fn max_tokens_param(&self) -> Option<isize> {
if self.data.pass_max_tokens {
self.data.max_output_tokens
@@ -182,30 +197,45 @@ impl Model {
}
}
- pub fn max_input_tokens_limit(&self, messages: &[Message]) -> Result<()> {
+ pub fn guard_max_input_tokens(&self, messages: &[Message]) -> Result<()> {
let total_tokens = self.total_tokens(messages) + BASIS_TOKENS;
if let Some(max_input_tokens) = self.data.max_input_tokens {
if total_tokens >= max_input_tokens {
- bail!("Exceed max input tokens limit")
+ bail!("Exceed max_input_tokens limit")
}
}
Ok(())
}
+
+ pub fn guard_max_concurrent_chunks(&self, data: &EmbeddingsData) -> Result<()> {
+ if data.texts.len() > self.max_concurrent_chunks() {
+ bail!("Exceed max_concurrent_chunks limit");
+ }
+ Ok(())
+ }
}
#[derive(Debug, Clone, Default, Deserialize)]
pub struct ModelData {
pub name: String,
+ #[serde(default = "default_model_mode")]
+ pub mode: String,
pub max_input_tokens: Option<usize>,
+ pub input_price: Option<f64>,
+ pub output_price: Option<f64>,
+
+ // chat-only properties
pub max_output_tokens: Option<isize>,
#[serde(default)]
pub pass_max_tokens: bool,
- pub input_price: Option<f64>,
- pub output_price: Option<f64>,
#[serde(default)]
pub supports_vision: bool,
#[serde(default)]
pub supports_function_calling: bool,
+
+ // embedding-only properties
+ pub default_chunk_size: Option<usize>,
+ pub max_concurrent_chunks: Option<usize>,
}
impl ModelData {
@@ -222,3 +252,7 @@ pub struct BuiltinModels {
pub platform: String,
pub models: Vec<ModelData>,
}
+
+fn default_model_mode() -> String {
+ "chat".into()
+}