diff options
| author | sigoden <sigoden@gmail.com> | 2025-02-21 14:59:14 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-02-21 14:59:14 +0800 |
| commit | 3dd90e1e95c478cc9d92b587d44e9215d4b19a7e (patch) | |
| tree | da98bf6bd8b9fd46dacb2f6d6e09230c97d45533 | |
| parent | d1c603c50912873445b7108e7f71263a8da7e955 (diff) | |
| download | aichat-3dd90e1e95c478cc9d92b587d44e9215d4b19a7e.tar.gz | |
fix: incorrect model when switching role in session context (#1192)
| -rw-r--r-- | src/config/mod.rs | 13 | ||||
| -rw-r--r-- | src/rag/mod.rs | 2 |
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) } |
