summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-02-21 14:59:14 +0800
committerGitHub <noreply@github.com>2025-02-21 14:59:14 +0800
commit3dd90e1e95c478cc9d92b587d44e9215d4b19a7e (patch)
treeda98bf6bd8b9fd46dacb2f6d6e09230c97d45533
parentd1c603c50912873445b7108e7f71263a8da7e955 (diff)
downloadaichat-3dd90e1e95c478cc9d92b587d44e9215d4b19a7e.tar.gz
fix: incorrect model when switching role in session context (#1192)
-rw-r--r--src/config/mod.rs13
-rw-r--r--src/rag/mod.rs2
2 files changed, 8 insertions, 7 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 1ab500f..f8ea258 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -591,7 +591,7 @@ impl Config {
("use_tools", format_option_value(&role.use_tools())),
(
"max_output_tokens",
- self.model
+ role.model()
.max_tokens_param()
.map(|v| format!("{v} (current model)"))
.unwrap_or_else(|| "null".into()),
@@ -867,7 +867,7 @@ impl Config {
pub fn use_prompt(&mut self, prompt: &str) -> Result<()> {
let mut role = Role::new(TEMP_ROLE_NAME, prompt);
- role.set_model(&self.model);
+ role.set_model(self.current_model());
self.use_role_obj(role)
}
@@ -923,16 +923,17 @@ impl Config {
} else {
Role::builtin(name)?
};
+ let current_model = self.current_model();
match role.model_id() {
Some(model_id) => {
- if self.model.id() != model_id {
+ if current_model.id() != model_id {
let model = Model::retrieve_model(self, model_id, ModelType::Chat)?;
role.set_model(&model);
} else {
- role.set_model(&self.model);
+ role.set_model(current_model);
}
}
- None => role.set_model(&self.model),
+ None => role.set_model(current_model),
}
Ok(role)
}
@@ -1795,7 +1796,7 @@ impl Config {
};
} else if cmd == ".set" && args.len() == 2 {
let candidates = match args[0] {
- "max_output_tokens" => match self.model.max_output_tokens() {
+ "max_output_tokens" => match self.current_model().max_output_tokens() {
Some(v) => vec![v.to_string()],
None => vec![],
},
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index 93a43c8..f799b3d 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -841,7 +841,7 @@ impl Debug for DocumentId {
impl DocumentId {
pub fn new(file_index: usize, document_index: usize) -> Self {
- let value = file_index << (usize::BITS / 2) | document_index;
+ let value = (file_index << (usize::BITS / 2)) | document_index;
Self(value)
}