From 35f306bf2abc8806362450d5c8d1cad4258c900a Mon Sep 17 00:00:00 2001 From: Georg Bauer Date: Tue, 1 Sep 2026 22:37:57 +0200 Subject: [PATCH] Lock tool sessions to their model --- docs/USER_GUIDE.md | 7 ++ .../20260901223000_add_session_model/down.sql | 1 + .../20260901223000_add_session_model/up.sql | 7 ++ src/agent.rs | 6 + src/app.rs | 115 +++++++++++++++++- src/app/generation.rs | 5 + src/app/projects.rs | 22 ++++ src/app/view.rs | 38 ++++++ src/app/view/chat.rs | 4 +- src/app/view/stats.rs | 2 +- src/database.rs | 55 ++++++++- src/engine.rs | 39 ++++++ src/model.rs | 4 + src/schema.rs | 1 + 14 files changed, 299 insertions(+), 7 deletions(-) create mode 100644 migrations/20260901223000_add_session_model/down.sql create mode 100644 migrations/20260901223000_add_session_model/up.sql diff --git a/docs/USER_GUIDE.md b/docs/USER_GUIDE.md index 826d8bf..915e7f5 100644 --- a/docs/USER_GUIDE.md +++ b/docs/USER_GUIDE.md @@ -196,6 +196,13 @@ stored per model; the custom system prompt is shared across profiles. Other sections control A2UI, permission defaults, endpoint settings, Git diff display, Dev Brain, extensions, checkpoint storage, and diagnostics. +A session remains model-independent until its first tool call. That call locks +the session to its model so persisted tool syntax is never mixed. Continuing a +locked session while another model is active asks before switching back. Legacy +sessions recover the exact model from their checkpoint when available; a legacy +tool session whose exact model can no longer be identified remains viewable but +cannot be continued. + ### Main views, sidebar, and branches **Chat** shows the conversation, **A2UI** shows interactive surfaces, **Git** diff --git a/migrations/20260901223000_add_session_model/down.sql b/migrations/20260901223000_add_session_model/down.sql new file mode 100644 index 0000000..3522d01 --- /dev/null +++ b/migrations/20260901223000_add_session_model/down.sql @@ -0,0 +1 @@ +ALTER TABLE sessions DROP COLUMN model; diff --git a/migrations/20260901223000_add_session_model/up.sql b/migrations/20260901223000_add_session_model/up.sql new file mode 100644 index 0000000..4a086d3 --- /dev/null +++ b/migrations/20260901223000_add_session_model/up.sql @@ -0,0 +1,7 @@ +ALTER TABLE sessions ADD COLUMN model TEXT + CHECK (model IS NULL OR model IN ( + 'deepseek-v4-flash-0731', + 'deepseek-v4-pro', + 'glm-5.2', + 'glm-5.3-flash' + )); diff --git a/src/agent.rs b/src/agent.rs index 9da4c55..d3862c9 100644 --- a/src/agent.rs +++ b/src/agent.rs @@ -3627,6 +3627,12 @@ pub(crate) fn parse_tool_calls( .map(|calls| (content, calls)) } +pub(crate) fn tool_protocol_model(text: &str) -> Option { + [ModelChoice::DeepSeekV4Flash0731, ModelChoice::Glm52] + .into_iter() + .find(|model| parse_tool_calls(*model, text).is_ok_and(|(_, calls)| !calls.is_empty())) +} + pub(crate) fn system_prompt(model: ModelChoice, extra: &str, dev_brain: bool) -> String { system_prompt_with_tools(model, extra, dev_brain, false) } diff --git a/src/app.rs b/src/app.rs index decc1ad..5ffe349 100644 --- a/src/app.rs +++ b/src/app.rs @@ -105,6 +105,8 @@ pub(crate) struct App { session_menu: Option, /// Session waiting for explicit confirmation before deletion. pending_session_delete: Option, + /// Locked session model waiting to replace the current global selection. + pending_session_model_switch: Option, /// Session being renamed, with the in-progress title. session_rename: Option<(i32, String)>, /// Projects whose archived sessions are expanded in the sidebar. @@ -500,6 +502,8 @@ pub(crate) enum Message { AllowToolOnce, DenyTool, SubmitPrompt, + ConfirmSessionModelSwitch, + CancelSessionModelSwitch, StopGeneration, GenerationTick, ChatScrolled(scrollable::Viewport), @@ -645,6 +649,7 @@ impl App { background_chats: HashMap::new(), session_menu: None, pending_session_delete: None, + pending_session_model_switch: None, session_rename: None, expanded_archives: HashSet::new(), sidebar_drag: false, @@ -810,6 +815,7 @@ impl App { background_chats: HashMap::new(), session_menu: None, pending_session_delete: None, + pending_session_model_switch: None, session_rename: None, expanded_archives: HashSet::new(), sidebar_drag: false, @@ -1134,7 +1140,7 @@ impl App { let content = export_chat( format, &title, - self.config.model, + self.chat_parser_model(), &self.conversation, &surfaces, &images, @@ -1241,6 +1247,8 @@ impl App { Message::DismissPanel => { if self.quit_confirmation { self.quit_confirmation = false; + } else if self.pending_session_model_switch.is_some() { + self.pending_session_model_switch = None; } else if self.git_diff.is_some() { self.git_diff = None; } else if self.git_commit_all_confirmation { @@ -1559,10 +1567,57 @@ impl App { } } Message::SubmitPrompt => { + if !self.generating && self.selected_session.is_some() { + match required_session_model_switch( + self.config.model, + self.selected_session_model(), + self.chat_tool_protocol_model(), + ) { + Ok(Some(model)) => { + self.pending_session_model_switch = Some(model); + return Task::none(); + } + Err(protocol) => { + self.error = Some(format!( + "This legacy session contains {} tool calls, but its exact model was not recorded and no compatible checkpoint remains. Start a new session to continue safely.", + if protocol.is_glm() { "GLM" } else { "DeepSeek" } + )); + return Task::none(); + } + Ok(None) => {} + } + } self.chat_follow_tail = true; self.start_generation(); return scroll_chat_to_end(); } + Message::ConfirmSessionModelSwitch => { + let Some(model) = self.pending_session_model_switch.take() else { + return Task::none(); + }; + let mut config = self.config.clone(); + config.model = model; + if let Err(error) = config.save(&config_path()) { + self.error = Some(error); + return Task::none(); + } + self.config = config; + self.preference_draft = PreferenceDraft::from_saved(&self.config); + self.context_limit = self.config.active_generation().context_tokens.max(0) as u32; + if model != ModelChoice::Glm53Flash { + self.pending_vision_image = None; + } + #[cfg(target_os = "macos")] + { + self.agent_tools = None; + preferences::update_runtime_config(&self.runtime_config, &self.config); + } + self.error = None; + self.chat_follow_tail = true; + self.start_generation(); + return scroll_chat_to_end(); + } + Message::CancelSessionModelSwitch => self.pending_session_model_switch = None, Message::StopGeneration => { self.stop_requested = true; self.activity = Some("Stopping…".into()); @@ -1938,6 +1993,7 @@ impl App { if self.selected_session == Some(session_id) { return Task::none(); } + self.pending_session_model_switch = None; #[cfg(target_os = "macos")] { self.leave_current_chat(); @@ -1964,6 +2020,7 @@ impl App { session.context_limit, session.last_tokens_per_second, session.permission_mode(), + session.model(), ) }); let Some(database) = &mut self.database else { @@ -1980,6 +2037,27 @@ impl App { self.error = Some(format!("Could not update the session: {error}")); return Task::none(); } + let stored_model = saved_context + .as_ref() + .and_then(|(_, _, _, _, model)| *model); + let protocol_model = messages.iter().find_map(|message| { + crate::agent::tool_protocol_model(&message.content) + }); + if stored_model.is_none() + && let (Some(protocol_model), Some(checkpoint_model)) = ( + protocol_model, + crate::engine::checkpoint_model(&session_checkpoint_path( + session_id, + )), + ) + && protocol_model.is_glm() == checkpoint_model.is_glm() + && let Err(error) = + database.set_session_model(session_id, checkpoint_model) + { + self.error = + Some(format!("Could not restore the session model: {error}")); + return Task::none(); + } self.conversation = messages.into_iter().map(ChatMessage::from).collect(); self.chat_follow_tail = true; generation::promote_legacy_turn_summaries(&mut self.conversation); @@ -2005,8 +2083,8 @@ impl App { self.selected_session = Some(session_id); self.system_prompt_seen_at = 0; self.queued_inputs.clear(); - let (used, limit, tokens_per_second, permission_mode) = - saved_context.unwrap_or((0, 0, None, PermissionMode::default())); + let (used, limit, tokens_per_second, permission_mode, _) = + saved_context.unwrap_or((0, 0, None, PermissionMode::default(), None)); self.context_used = used.max(0) as u32; self.context_limit = if limit > 0 { limit as u32 @@ -3154,6 +3232,18 @@ fn session_checkpoint_path(session_id: i32) -> PathBuf { kv_cache_path().join(format!("{session_id}.bin")) } +fn required_session_model_switch( + current: ModelChoice, + locked: Option, + legacy_protocol: Option, +) -> Result, ModelChoice> { + match locked { + Some(model) if model != current => Ok(Some(model)), + Some(_) => Ok(None), + None => legacy_protocol.map_or(Ok(None), Err), + } +} + fn discard_session_checkpoint_files(directory: &Path, session_id: i32) -> Result { let mut removed = false; for extension in ["bin", "tmp", "compacting"] { @@ -3204,6 +3294,25 @@ pub(crate) fn app_icon() -> window::Icon { mod tests { use super::*; + #[test] + fn tool_calls_lock_continuation_to_the_session_model() { + let flash = ModelChoice::DeepSeekV4Flash0731; + let glm = ModelChoice::Glm53Flash; + assert_eq!( + required_session_model_switch(flash, Some(flash), None), + Ok(None) + ); + assert_eq!( + required_session_model_switch(glm, Some(flash), None), + Ok(Some(flash)) + ); + assert_eq!( + required_session_model_switch(glm, None, Some(flash)), + Err(flash) + ); + assert_eq!(required_session_model_switch(glm, None, None), Ok(None)); + } + #[test] fn preferences_shortcut_and_dspark_support_are_explicit() { let message = shortcut( diff --git a/src/app/generation.rs b/src/app/generation.rs index da2b6e1..cd66e0f 100644 --- a/src/app/generation.rs +++ b/src/app/generation.rs @@ -1584,6 +1584,11 @@ impl App { let session_id = self .selected_session .ok_or_else(|| "The active session is unavailable.".to_owned())?; + self.database + .as_mut() + .ok_or_else(|| "The project database is unavailable.".to_owned())? + .set_session_model(session_id, self.config.model)?; + self.reload_projects(); if self.agent_tools.as_ref().map(|(id, _)| *id) != Some(session_id) { let project_id = self .selected_project diff --git a/src/app/projects.rs b/src/app/projects.rs index 7bc8a2c..7599e81 100644 --- a/src/app/projects.rs +++ b/src/app/projects.rs @@ -1,6 +1,27 @@ use super::*; impl App { + pub(super) fn selected_session_model(&self) -> Option { + let session_id = self.selected_session?; + self.projects + .iter() + .flat_map(|project| &project.sessions) + .find(|session| session.id == session_id) + .and_then(|session| session.model()) + } + + pub(super) fn chat_tool_protocol_model(&self) -> Option { + self.conversation + .iter() + .find_map(|message| crate::agent::tool_protocol_model(&message.content)) + } + + pub(super) fn chat_parser_model(&self) -> ModelChoice { + self.selected_session_model() + .or_else(|| self.chat_tool_protocol_model()) + .unwrap_or(self.config.model) + } + pub(super) fn clear_a2ui(&mut self) { self.a2ui.clear(); self.a2ui_history.clear(); @@ -171,6 +192,7 @@ impl App { self.drafts.entry(project_id).or_insert(title); self.remember_project(project_id); self.selected_session = None; + self.pending_session_model_switch = None; self.permission_mode = self.config.default_permission_mode; self.conversation.clear(); self.chat_follow_tail = true; diff --git a/src/app/view.rs b/src/app/view.rs index 96df342..fe4b0bd 100644 --- a/src/app/view.rs +++ b/src/app/view.rs @@ -110,6 +110,7 @@ impl App { /// tree, so the layers below have to stay out of the dialog's field order. pub(super) fn modal_open(&self) -> bool { self.quit_confirmation + || self.pending_session_model_switch.is_some() || self.git_diff.is_some() || self.git_commit_all_confirmation || self.pending_project_path.is_some() @@ -174,6 +175,8 @@ impl App { #[cfg(target_os = "macos")] if self.quit_confirmation { layers.push(self.quit_confirmation_panel()); + } else if let Some(model) = self.pending_session_model_switch { + layers.push(self.session_model_switch_panel(model)); } else if let Some((prompt, _)) = &self.pending_tool_approval { layers.push(self.tool_approval_panel(prompt)); } else if let Some(path) = &self.pending_project_path { @@ -196,6 +199,8 @@ impl App { #[cfg(not(target_os = "macos"))] if self.quit_confirmation { layers.push(self.quit_confirmation_panel()); + } else if let Some(model) = self.pending_session_model_switch { + layers.push(self.session_model_switch_panel(model)); } else if let Some(path) = &self.pending_project_path { layers.push(self.project_dialog(path)); } else if let Some((_, title)) = &self.session_rename { @@ -252,6 +257,39 @@ impl App { ) } + fn session_model_switch_panel(&self, model: ModelChoice) -> Element<'_, Message> { + let title = self.active_chat_title().unwrap_or("This session"); + let dialog = container( + column![ + text("Switch model to continue?").size(22), + text(format!( + "“{title}” is locked to {model} because it contains that model's tool calls. Continuing will switch the active model from {} to {model}.", + self.config.model + )) + .size(13), + row![ + Space::new().width(Length::Fill), + action_button("Cancel").on_press(Message::CancelSessionModelSwitch), + action_button("Switch and continue") + .on_press(Message::ConfirmSessionModelSwitch), + ] + .spacing(8), + ] + .spacing(12), + ) + .padding(22) + .width(520) + .style(overview_style); + opaque( + container(dialog) + .center_x(Length::Fill) + .center_y(Length::Fill) + .style(|_| { + container::Style::default().background(Color::from_rgba8(0, 0, 0, 0.68)) + }), + ) + } + #[cfg(target_os = "macos")] fn tool_approval_panel<'a>( &'a self, diff --git a/src/app/view/chat.rs b/src/app/view/chat.rs index 0956a95..f0a7e0e 100644 --- a/src/app/view/chat.rs +++ b/src/app/view/chat.rs @@ -80,7 +80,7 @@ impl App { if message.tool && index > 0 && !crate::agent::stored_tool_cards( - self.config.model, + self.chat_parser_model(), &self.conversation[index - 1].content, None, &self.conversation[index - 1].tool_approval_reasons, @@ -191,7 +191,7 @@ impl App { } } else { crate::agent::stored_tool_cards( - self.config.model, + self.chat_parser_model(), &message.content, stored_result, &message.tool_approval_reasons, diff --git a/src/app/view/stats.rs b/src/app/view/stats.rs index 40d3809..4082513 100644 --- a/src/app/view/stats.rs +++ b/src/app/view/stats.rs @@ -5,7 +5,7 @@ use iced::widget::column; impl App { pub(super) fn stats_dashboard(&self) -> Element<'_, Message> { let stats = &self.metrics_snapshot; - let session = SessionStats::from_messages(&self.conversation, self.config.model); + let session = SessionStats::from_messages(&self.conversation, self.chat_parser_model()); let context_fraction = if stats.context_limit == 0 { 0.0 } else { diff --git a/src/database.rs b/src/database.rs index 8e198d2..f86b7fb 100644 --- a/src/database.rs +++ b/src/database.rs @@ -45,6 +45,8 @@ pub struct Session { last_used: i64, /// Raw column value; read it through [`Session::permission_mode`]. permission_mode: String, + /// Locked after the first tool call; read it through [`Session::model`]. + model: Option, } impl Session { @@ -56,6 +58,12 @@ impl Session { PermissionMode::from_id(&self.permission_mode).unwrap_or_default() } + pub(crate) fn model(&self) -> Option { + self.model + .as_deref() + .and_then(crate::model::ModelChoice::from_id) + } + #[cfg(test)] pub fn fixture(id: i32, project_id: i32, title: &str, state: SessionState) -> Self { Self { @@ -69,6 +77,7 @@ impl Session { compacted_summary: None, last_used: 0, permission_mode: PermissionMode::default().as_id().to_owned(), + model: None, } } } @@ -411,6 +420,25 @@ impl Database { .map_err(|error| error.to_string()) } + pub(crate) fn set_session_model( + &mut self, + session_id: i32, + model: crate::model::ModelChoice, + ) -> Result<(), String> { + let updated = diesel::update( + sessions::table + .find(session_id) + .filter(sessions::model.is_null().or(sessions::model.eq(model.id()))), + ) + .set(sessions::model.eq(model.id())) + .execute(&mut self.connection) + .map_err(|error| error.to_string())?; + if updated == 0 { + return Err("The session is locked to a different model.".into()); + } + Ok(()) + } + pub fn touch_session(&mut self, session_id: i32) -> Result<(), String> { touch_session(&mut self.connection, session_id).map_err(|error| error.to_string()) } @@ -820,13 +848,29 @@ mod tests { assert_eq!(loaded[0].sessions[0].title, "Second session"); assert_eq!(loaded[0].sessions[0].state(), SessionState::Normal); assert_eq!(loaded[0].sessions[0].permission_mode(), PermissionMode::Ai); + assert_eq!(loaded[0].sessions[0].model(), None); database .set_session_permission_mode(second.id, PermissionMode::Heuristic) .unwrap(); + database + .set_session_model(second.id, crate::model::ModelChoice::DeepSeekV4Flash0731) + .unwrap(); + database + .set_session_model(second.id, crate::model::ModelChoice::DeepSeekV4Flash0731) + .unwrap(); + assert!( + database + .set_session_model(second.id, crate::model::ModelChoice::Glm53Flash) + .is_err() + ); assert_eq!( database.load_projects().unwrap()[0].sessions[0].permission_mode(), PermissionMode::Heuristic ); + assert_eq!( + database.load_projects().unwrap()[0].sessions[0].model(), + Some(crate::model::ModelChoice::DeepSeekV4Flash0731) + ); let ordinary = loaded[0].sessions[0].id; let pinned = database @@ -1288,6 +1332,9 @@ mod tests { "Tool result 1 (present_svg):\nSVG presented inline: A dark rectangle\n", ) .unwrap(); + database + .set_session_model(session.id, crate::model::ModelChoice::DeepSeekV4Flash0731) + .unwrap(); database .record_compaction(session.id, "Later context summary.", None, None, 100, 1_000) .unwrap(); @@ -1307,8 +1354,14 @@ mod tests { .iter() .position(|message| message.id == assistant.id) .unwrap(); + let session_model = reopened.load_projects().unwrap()[0] + .sessions + .iter() + .find(|stored| stored.id == session.id) + .and_then(Session::model) + .unwrap(); let cards = crate::agent::stored_tool_cards( - crate::model::ModelChoice::DeepSeekV4Flash0731, + session_model, &messages[assistant_index].content, Some(&messages[assistant_index + 1].content), &[], diff --git a/src/engine.rs b/src/engine.rs index 8b7b9e8..fc97417 100644 --- a/src/engine.rs +++ b/src/engine.rs @@ -22,6 +22,8 @@ use kvstore::{KvStore, StoreReason}; use sha2::{Digest, Sha256}; #[cfg(target_os = "macos")] use std::collections::HashMap; +use std::fs::File; +use std::io::{Read, Seek, SeekFrom}; use std::path::Path; #[cfg(target_os = "macos")] use std::path::PathBuf; @@ -43,6 +45,24 @@ const VISION_START_TOKEN: i32 = 154_830; const VISION_END_TOKEN: i32 = 154_831; type VisionOverlays = Vec<(u32, metal::VisionEmbedding)>; +pub(crate) fn checkpoint_model(path: &Path) -> Option { + let mut file = File::open(path).ok()?; + let mut magic = [0; 8]; + file.read_exact(&mut magic).ok()?; + let model_size_offset = match &magic { + b"DS4RKV01" => 40, + b"DS4GLM01" => 44, + _ => return None, + }; + file.seek(SeekFrom::Start(model_size_offset)).ok()?; + let mut bytes = [0; 8]; + file.read_exact(&mut bytes).ok()?; + let size = u64::from_le_bytes(bytes); + crate::model::MODEL_CHOICES + .into_iter() + .find(|model| model.main_artifact_size() == size) +} + pub(crate) fn vision_data_marker(uri: &str) -> String { format!("{VISION_DATA_START}{uri}{VISION_DATA_END}") } @@ -1917,6 +1937,25 @@ impl Rng { mod sampling_tests { use super::*; + #[test] + fn checkpoint_headers_identify_the_exact_model() { + let path = + std::env::temp_dir().join(format!("ds4-checkpoint-model-{}", std::process::id())); + for (magic, offset, model) in [ + (b"DS4RKV01".as_slice(), 40, ModelChoice::DeepSeekV4Flash0731), + (b"DS4RKV01".as_slice(), 40, ModelChoice::DeepSeekV4Pro), + (b"DS4GLM01".as_slice(), 44, ModelChoice::Glm52), + (b"DS4GLM01".as_slice(), 44, ModelChoice::Glm53Flash), + ] { + let mut bytes = magic.to_vec(); + bytes.resize(offset, 0); + bytes.extend_from_slice(&model.main_artifact_size().to_le_bytes()); + std::fs::write(&path, bytes).unwrap(); + assert_eq!(checkpoint_model(&path), Some(model)); + } + std::fs::remove_file(path).unwrap(); + } + #[test] fn zero_temperature_is_greedy() { let mut rng = Rng::new(1); diff --git a/src/model.rs b/src/model.rs index 55173d2..7301d61 100644 --- a/src/model.rs +++ b/src/model.rs @@ -122,6 +122,10 @@ impl ModelChoice { self.is_glm() } + pub(crate) fn main_artifact_size(self) -> u64 { + self.main_artifact().size + } + fn main_artifact(self) -> &'static Artifact { match self { Self::DeepSeekV4Flash0731 => &FLASH_0731, diff --git a/src/schema.rs b/src/schema.rs index e0db417..1b93aab 100644 --- a/src/schema.rs +++ b/src/schema.rs @@ -51,6 +51,7 @@ diesel::table! { compacted_summary -> Nullable, last_used -> BigInt, permission_mode -> Text, + model -> Nullable, } }