diff --git a/docs/user/tui-and-sessions.md b/docs/user/tui-and-sessions.md index 2eb3ab3..874b2f2 100644 --- a/docs/user/tui-and-sessions.md +++ b/docs/user/tui-and-sessions.md @@ -135,6 +135,14 @@ These local commands are available only while the session is idle. `/agents` tog `/model` opens the model selector. `/effort` opens the advertised ACP reasoning-effort selector; `/effort default|low|medium|high` selects directly. In either dialog, Tab toggles saving the selection to `~/.kit/config.toml`, Enter selects, and Esc closes. Saving `default` removes top-level `reasoning_effort`; other values update it without replacing unrelated TOML. A new or resumed process starts from the resolved CLI/TOML default unless the selection was saved. +Before changing models, Kit compares the latest provider-reported transcript occupancy with the target model's advertised context window. It adds a 20% tokenizer margin (`ceil(tokens × 1.20)`) and warns when that estimate is at least 80% of the target window. This is the latest request's occupancy, not accumulated session usage; cached input tokens are not counted twice. + +The TUI warning offers **Continue anyway**, **Compact**, and **Cancel** (the default). Continue switches without compacting. Compact runs the existing compactor with the **original model** and applies the new selection only after successful completion and durable transcript replacement. A failure or cancellation keeps the original model selected; compaction can already have changed the transcript if cancellation arrives after replacement. Cancel dismisses the warning without changing the model or transcript, and preserves the input draft. Esc cancels an in-progress switch; repeated Esc is harmless, and Ctrl+C while cancellation is pending exits if it is stuck. Stale confirmations are rejected when the session, selected model, target, or transcript changes. + +If the latest token count or target context window is unavailable (including fallback/custom models without catalog metadata), Kit allows an **unchecked switch** rather than inventing a limit or forcing compaction. The estimate is a warning, not a guarantee that the next provider request fits. + +External ACP clients receive a confirmation-required error instead of a Kit dialog. Its extension key in error data is `kit.model_switch`; clients can resubmit the same configuration request with `_meta: {"kit.model_switch": {"token": , "action": "continue"}}`. Both ACP versions also accept `"action": "compact"` to compact transactionally with the original model. Confirmation tokens are session-local, in-memory, one-shot decisions, not persisted configuration. + The ACP server advertises `compact` for every new session. The TUI submits `/compact` unchanged like any other prompt; the runtime consumes exactly one text part beginning with the exact raw token before model dispatch and permits other client-provided context parts. Used alone, it ends after compaction. Whitespace-trimmed text following `/compact` and any other context parts are retained as the latest user message and start the next turn after compaction. Leading whitespace, near-misses such as `/compactness`, prompts containing multiple `/compact` command parts, and unknown slash commands remain ordinary prompts. Local commands win if an advertised command has the same name. ## Persisted transcripts and session files diff --git a/src/compaction.rs b/src/compaction.rs index 47f2703..983a7cc 100644 --- a/src/compaction.rs +++ b/src/compaction.rs @@ -763,12 +763,12 @@ impl LoopMutator for AutomaticCompactor { } } -fn compaction_reason(transcript: &[Item]) -> Option { +pub(crate) fn latest_context_tokens(transcript: &[Item]) -> Option { let usage = transcript .iter() .rev() .find_map(|item| item.usage.as_ref())?; - let used = usage + usage .metadata .get("context_used") .and_then(serde_json::Value::as_u64) @@ -777,7 +777,15 @@ fn compaction_reason(transcript: &[Item]) -> Option { .tokens .as_ref() .and_then(|tokens| tokens.input_tokens.checked_add(tokens.output_tokens)) - })?; + }) +} + +fn compaction_reason(transcript: &[Item]) -> Option { + let used = latest_context_tokens(transcript)?; + let usage = transcript + .iter() + .rev() + .find_map(|item| item.usage.as_ref())?; let window = usage .metadata .get("context_window") diff --git a/src/protocols/acp.rs b/src/protocols/acp.rs index a537971..0ae862a 100644 --- a/src/protocols/acp.rs +++ b/src/protocols/acp.rs @@ -48,6 +48,7 @@ use tokio::{ use tracing::Instrument as _; mod activity; +pub(crate) mod model_switch; mod skill_catalog; pub mod v2; @@ -462,7 +463,9 @@ enum Command { Cancel, SetConfig { request: SetSessionConfigOptionRequest, - reply: oneshot::Sender>, + cancellation_generation: u64, + reply: + oneshot::Sender>, }, Fork { parent_context: Option<(String, String)>, @@ -1581,14 +1584,24 @@ impl Server { async fn set_config( &self, request: SetSessionConfigOptionRequest, - ) -> Result { - let sender = self.sender(&request.session_id).await?; + ) -> Result { + let sender = self.sender(&request.session_id).await.map_err(sdk_error)?; + let cancellation_generation = self + .integration + .cancellation_handle(&request.session_id) + .map_err(sdk_error)? + .generation(); let (tx, rx) = oneshot::channel(); sender - .send(Command::SetConfig { request, reply: tx }) + .send(Command::SetConfig { + request, + reply: tx, + cancellation_generation, + }) .await - .map_err(|_| AcpRuntimeError::ClientClosed)?; - rx.await.map_err(|_| AcpRuntimeError::ClientClosed)? + .map_err(|_| sdk_error(AcpRuntimeError::ClientClosed))?; + rx.await + .map_err(|_| sdk_error(AcpRuntimeError::ClientClosed))? } async fn cancel(&self, notification: CancelNotification) -> Result<(), AcpRuntimeError> { @@ -1649,6 +1662,8 @@ impl Server { &self, session_id: &agentkit_acp::SessionId, ) -> Result, AcpRuntimeError> { + // Publication holds this lock across registry registration and commit. + // A poisoned map cannot establish that the session lifecycle is consistent. self.sessions .lock() .map_err(|_| AcpRuntimeError::ClientClosed)? @@ -1750,6 +1765,7 @@ async fn session_actor(actor: SessionActor) { mut mcp_events, } = actor; let mut binding = Some(binding); + let mut model_switch = model_switch::Guard::default(); loop { tokio::select! { // A queued cancel or close wins over a simultaneously-ready task @@ -1778,8 +1794,35 @@ async fn session_actor(actor: SessionActor) { // The server already interrupted the shared controller; this // marker only establishes its serialized actor position. Some(Command::Cancel) => {} - Some(Command::SetConfig { request, reply }) => { - let result = set_config(&adapter, &catalog, request); + Some(Command::SetConfig { request, reply, cancellation_generation }) => { + let result = async { + let cancellation = integration.cancellation_handle(&session_id).map_err(sdk_error)?; + if cancellation.is_cancelled_since(cancellation_generation) { return Err(model_switch::error("model change cancelled")); } + if request.config_id.to_string() == MODEL_CONFIG_ID { + let target = request.value.as_value_id().ok_or_else(|| model_switch::error("selection requires an id value"))?; + let decision = model_switch.check( + (&adapter.selection().map_err(|error| model_switch::error(&error))?, cancellation_generation), + ModelSelection::from_id(&target.to_string()).map_err(|error| model_switch::error(&error))?, + &catalog, driver.snapshot().transcript, + request.meta.as_ref().and_then(|meta| meta.get(model_switch::META)), + )?; + if decision == model_switch::Decision::Compact { + let marker = model_switch::compact_marker(); + let marker_id = marker.id.clone(); + driver.submit_input(vec![marker]).map_err(|error| sdk_error(record_acp_loop_failure(&session_id, &error)))?; + let reason = activity.execute(activity::ExecutionOrigin::Prompt, + drive_finalized(&session_id, &integration, &mut driver, false, None), + |reason| Some(reason.clone()), + ).await.map_err(sdk_error)?; + if !model_switch::compaction_completed(&reason, &marker_id, &driver.snapshot().transcript, + cancellation.is_cancelled_since(cancellation_generation)) { + return Err(model_switch::error("compaction did not complete; model unchanged")); + } + } + } + if cancellation.is_cancelled_since(cancellation_generation) { return Err(model_switch::error("model change cancelled")); } + set_config(&adapter, &catalog, request).map_err(sdk_error) + }.await; let _ = reply.send(result); } Some(Command::Fork { @@ -2588,8 +2631,7 @@ fn component( async move |request: SetSessionConfigOptionRequest, responder, cx| { let state = Arc::clone(&state); cx.spawn(async move { - responder - .respond_with_result(state.set_config(request).await.map_err(sdk_error)) + responder.respond_with_result(state.set_config(request).await) })?; Ok(()) } @@ -4701,6 +4743,72 @@ pub(super) mod tests { (jobs, tasks) } + #[tokio::test] + async fn set_config_rejects_poisoned_session_map_without_queueing_switch() { + let root = tempfile::tempdir().unwrap(); + let server = Server::new( + Runtime::new(root.path(), "gpt-5.4").unwrap(), + AcpIntegration::builder() + .name("poison-test") + .approval_resolver(AutoDenyResolver) + .build() + .unwrap(), + SessionRegistry::new(), + ); + let session_id = agentkit_acp::SessionId::new("poisoned-session"); + let (client, _messages) = AcpClientHandle::channel(); + server + .integration + .bind_session(AcpSessionBinding::new( + session_id.clone(), + AgentkitSessionId::new("poisoned-session"), + client, + )) + .unwrap(); + let (commands, mut received) = mpsc::channel(1); + server.sessions.lock().unwrap().insert( + session_id.clone(), + SessionHandle { + token: 1, + commands, + background_jobs: BackgroundJobs::default(), + structured_completion: false, + tasks: AsyncTaskManager::new().handle(), + }, + ); + assert!( + std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + let _guard = server.sessions.lock().unwrap(); + panic!("poison session map"); + })) + .is_err() + ); + + for id in [session_id, agentkit_acp::SessionId::new("missing-session")] { + let error = timeout( + Duration::from_secs(1), + server.set_config(SetSessionConfigOptionRequest::new( + id, + MODEL_CONFIG_ID, + "openai-subscription:gpt-5.4-mini", + )), + ) + .await + .expect("poison must return an error without waiting for the actor") + .unwrap_err(); + assert_eq!(error.code, agent_client_protocol::ErrorCode::InternalError); + assert_eq!( + error.data, + Some(serde_json::json!(AcpRuntimeError::ClientClosed.to_string())) + ); + } + assert!(server.sessions.is_poisoned()); + assert!(matches!( + received.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + } + #[tokio::test] async fn depth_zero_cancel_leaves_detached_background_running() { let (jobs, tasks) = start_non_cooperative_background("root-call").await; @@ -5459,6 +5567,7 @@ pub(super) mod tests { let adapter = SelectableAdapter::new(crate::ProviderKind::OpenAiSubscription, "gpt-5.4").unwrap(); let catalog = vec![ModelGroup { + context_windows: Default::default(), provider: crate::ProviderKind::OpenAiSubscription, models: vec!["gpt-5.4".into(), "gpt-5.4-mini".into()], }]; @@ -5484,6 +5593,7 @@ pub(super) mod tests { let adapter = SelectableAdapter::new(crate::ProviderKind::OpenAiSubscription, "gpt-5.4").unwrap(); let catalog = vec![ModelGroup { + context_windows: Default::default(), provider: crate::ProviderKind::OpenAiSubscription, models: vec!["gpt-5.4".into()], }]; diff --git a/src/protocols/acp/model_switch.rs b/src/protocols/acp/model_switch.rs new file mode 100644 index 0000000..e676441 --- /dev/null +++ b/src/protocols/acp/model_switch.rs @@ -0,0 +1,346 @@ +//! Ephemeral, actor-owned confirmation. Never infer a completion from runtime telemetry. +use std::sync::atomic::{AtomicU64, Ordering}; + +use agent_client_protocol::Error; +use agentkit_core::{FinishReason, Item, ItemKind, MessageId}; +use serde::{Deserialize, Serialize}; +use serde_json::{Value, json}; + +use crate::provider::{ModelGroup, ModelSelection}; + +pub(crate) const META: &str = "kit.model_switch"; +static NEXT_TOKEN: AtomicU64 = AtomicU64::new(1); + +#[derive(Clone, Debug, Deserialize, Serialize)] +pub(crate) struct Confirmation { + pub token: u64, + pub action: Decision, +} + +#[derive(Clone, Copy, Debug, Deserialize, Serialize, PartialEq, Eq)] +#[serde(rename_all = "snake_case")] +pub(crate) enum Decision { + Continue, + Compact, +} + +#[derive(Clone, Debug, Deserialize, Serialize)] +pub(crate) struct Warning { + pub token: u64, + pub guarded_tokens: String, + pub target_window: u64, +} + +struct Pending { + token: u64, + current: ModelSelection, + target: ModelSelection, + transcript: Vec, + cancellation_generation: u64, +} + +#[derive(Default)] +pub(super) struct Guard { + pending: Option, +} + +/// The margin applies to the latest request occupancy, not lifetime usage. Cached +/// tokens are already part of input_tokens/context_used and must not be added again. +fn guarded_tokens(transcript: &[Item]) -> Option { + crate::compaction::latest_context_tokens(transcript) + .map(|tokens| (u128::from(tokens) * 120).div_ceil(100)) +} + +pub(super) fn error(message: &str) -> Error { + agent_client_protocol::util::internal_error(message) +} + +impl Guard { + pub(super) fn check( + &mut self, + current: (&ModelSelection, u64), + target: ModelSelection, + catalog: &[ModelGroup], + transcript: Vec, + confirmation: Option<&Value>, + ) -> Result { + let (current, cancellation_generation) = current; + let pending = self.pending.take(); + if let Some(value) = confirmation { + let confirmation: Confirmation = serde_json::from_value(value.clone()) + .map_err(|_| error("invalid model-switch confirmation"))?; + let valid = pending.is_some_and(|pending| { + pending.token == confirmation.token + && pending.cancellation_generation == cancellation_generation + && pending.current == *current + && pending.target == target + && pending.transcript == transcript + }); + return if valid { + Ok(confirmation.action) + } else { + Err(error( + "model-switch confirmation is stale; select the model again", + )) + }; + } + let group = catalog + .iter() + .find(|group| group.provider == target.provider && group.models.contains(&target.model)) + .ok_or_else(|| error("model is not in the advertised catalog"))?; + let window = group + .context_windows + .get(&target.model) + .copied() + .filter(|size| *size > 0); + // Unknown occupancy/window is explicitly unchecked, not an invented limit + // or a mandatory compaction that could prevent selecting a custom model. + let Some((tokens, window)) = guarded_tokens(&transcript).zip(window) else { + return Ok(Decision::Continue); + }; + if current == &target || tokens * 100 < u128::from(window) * 80 { + return Ok(Decision::Continue); + } + let token = NEXT_TOKEN.fetch_add(1, Ordering::Relaxed); + self.pending = Some(Pending { + token, + current: current.clone(), + target, + transcript, + cancellation_generation, + }); + Err(error("target model context is at least 80% occupied after a 20% tokenizer margin; continue explicitly, compact with the current model, or cancel") + .data(json!({ META: Warning { token, guarded_tokens: tokens.to_string(), target_window: window } }))) + } +} + +/// A fresh marker, never an uncorrelated `CompactionFinished` notification. +pub(super) fn compact_marker() -> Item { + Item::text(ItemKind::User, "/compact").with_id(MessageId::new(format!( + "kit-model-switch-{}", + NEXT_TOKEN.fetch_add(1, Ordering::Relaxed) + ))) +} + +pub(super) fn compaction_completed( + reason: &FinishReason, + marker: &Option, + transcript: &[Item], + cancelled: bool, +) -> bool { + !cancelled + && *reason == FinishReason::Completed + && marker.is_some() + && !transcript.iter().any(|item| &item.id == marker) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::provider::ProviderKind; + use agentkit_core::{MetadataMap, TokenUsage, Usage}; + + fn selection(model: &str) -> ModelSelection { + ModelSelection { + provider: ProviderKind::OpenRouter, + model: model.into(), + } + } + fn catalog(window: Option) -> Vec { + vec![ModelGroup { + provider: ProviderKind::OpenRouter, + models: vec!["target".into()], + context_windows: window + .map(|window| ("target".into(), window)) + .into_iter() + .collect(), + }] + } + fn measured(tokens: u64) -> Item { + Item::text(ItemKind::Assistant, "response") + .with_usage(Usage::new(TokenUsage::new(tokens, 0))) + } + fn check( + guard: &mut Guard, + transcript: Vec, + window: Option, + confirmation: Option<&Value>, + ) -> Result { + guard.check( + (&selection("original"), 0), + selection("target"), + &catalog(window), + transcript, + confirmation, + ) + } + fn warn(guard: &mut Guard) -> Warning { + let error = check(guard, vec![measured(100)], Some(150), None).unwrap_err(); + serde_json::from_value(error.data.unwrap()[META].clone()).unwrap() + } + + #[test] + fn model_switch_exact_threshold_and_below_include_rounded_margin() { + let mut guard = Guard::default(); + assert_eq!( + check(&mut guard, vec![measured(99)], Some(150), None).unwrap(), + Decision::Continue + ); + let warning = warn(&mut guard); + assert_eq!(warning.guarded_tokens, "120"); // ceil(100 * 1.20) == 0.80 * 150 + assert_eq!(warning.target_window, 150); + assert_eq!(guarded_tokens(&[measured(101)]), Some(122)); + assert!(check(&mut guard, vec![measured(100)], Some(151), None).is_ok()); + assert!(check(&mut guard, vec![measured(101)], Some(151), None).is_err()); + assert!(check(&mut guard, vec![measured(u64::MAX)], Some(u64::MAX), None).is_err()); + } + + #[test] + fn model_switch_latest_occupancy_not_lifetime_or_double_cached_tokens() { + let last = Item::text(ItemKind::Assistant, "response").with_usage(Usage::new( + TokenUsage::new(80, 20) + .with_cached_input_tokens(70) + .with_cache_write_input_tokens(10), + )); + assert_eq!( + guarded_tokens(&[measured(1_000_000), last.clone()]), + Some(120) + ); + let mut authoritative = last; + authoritative.usage.as_mut().unwrap().metadata = + MetadataMap::from_iter([("context_used".into(), json!(10))]); + assert_eq!(guarded_tokens(&[authoritative]), Some(12)); + let unknown = Item::text(ItemKind::Assistant, "unmeasured").with_usage(Usage::default()); + assert_eq!(guarded_tokens(&[measured(100), unknown]), None); + } + + #[test] + fn model_switch_unknown_values_allow_an_unchecked_switch() { + for (transcript, window) in [ + (vec![], Some(100)), + (vec![measured(100)], None), + (vec![measured(100)], Some(0)), + ] { + assert_eq!( + check(&mut Guard::default(), transcript, window, None).unwrap(), + Decision::Continue + ); + } + } + + #[test] + fn model_switch_actions_are_one_shot_and_bound_to_state() { + for action in [Decision::Continue, Decision::Compact] { + let mut guard = Guard::default(); + let warning = warn(&mut guard); + let confirmation = serde_json::to_value(Confirmation { + token: warning.token, + action, + }) + .unwrap(); + assert_eq!( + check( + &mut guard, + vec![measured(100)], + Some(150), + Some(&confirmation) + ) + .unwrap(), + action + ); + assert!( + check( + &mut guard, + vec![measured(100)], + Some(150), + Some(&confirmation) + ) + .is_err() + ); + } + for stale in [ + "transcript", + "model", + "target", + "session", + "new-request", + "cancel", + ] { + let mut guard = Guard::default(); + let warning = warn(&mut guard); + let confirmation = serde_json::to_value(Confirmation { + token: warning.token, + action: Decision::Continue, + }) + .unwrap(); + if stale == "session" { + guard = Guard::default(); + } + if stale == "new-request" { + let _ = warn(&mut guard); + } + assert!( + guard + .check( + ( + &selection(if stale == "model" { + "another" + } else { + "original" + }), + if stale == "cancel" { 1 } else { 0 } + ), + selection(if stale == "target" { + "another" + } else { + "target" + }), + &catalog(Some(150)), + vec![measured(if stale == "transcript" { 101 } else { 100 })], + Some(&confirmation), + ) + .is_err(), + "{stale}" + ); + } + } + + #[test] + fn model_switch_compaction_requires_exact_consumed_marker_and_success() { + let marker = compact_marker(); + let other = compact_marker(); + assert_ne!(marker.id, other.id); + assert!(!compaction_completed( + &FinishReason::Completed, + &marker.id, + std::slice::from_ref(&marker), + false + )); + assert!(compaction_completed( + &FinishReason::Completed, + &marker.id, + &[other], + false + )); + for reason in [ + FinishReason::Cancelled, + FinishReason::Error, + FinishReason::MaxTokens, + FinishReason::Blocked, + ] { + assert!(!compaction_completed(&reason, &marker.id, &[], false)); + } + assert!(!compaction_completed( + &FinishReason::Completed, + &marker.id, + &[], + true + )); + assert!(!compaction_completed( + &FinishReason::Completed, + &None, + &[], + false + )); + } +} diff --git a/src/protocols/acp/v2.rs b/src/protocols/acp/v2.rs index 3606f0b..b74d9b0 100644 --- a/src/protocols/acp/v2.rs +++ b/src/protocols/acp/v2.rs @@ -34,6 +34,7 @@ use crate::{ }; use super::activity::{ExecutionOrigin, SessionActivity}; +use super::model_switch; use super::{ AuthenticationRequiredData, CancelBackgroundRequest, CancelBackgroundResponse, @@ -477,7 +478,10 @@ enum Command { Prompt(PromptCommand), SetConfig { request: wire::SetSessionConfigOptionRequest, - reply: oneshot::Sender>, + reply: oneshot::Sender< + Result, + >, + cancellation_generation: u64, }, Close { reply: oneshot::Sender<()>, @@ -1053,14 +1057,36 @@ impl Server { async fn set_config( &self, request: wire::SetSessionConfigOptionRequest, - ) -> Result { - let sender = self.sender(&request.session_id)?; + ) -> Result { + let (sender, cancellation_generation) = { + // Publication holds this lock across registry registration and commit. + // A poisoned map cannot establish that the session lifecycle is consistent. + let sessions = self + .sessions + .lock() + .map_err(|_| sdk_error(AcpRuntimeError::ClientClosed))?; + let session = sessions.get(&request.session_id).ok_or_else(|| { + sdk_error(AcpRuntimeError::SessionNotFound( + request.session_id.to_string(), + )) + })?; + ( + session.commands.clone(), + session.integration.cancellation_handle().generation(), + ) + }; let (reply, response) = oneshot::channel(); sender - .send(Command::SetConfig { request, reply }) + .send(Command::SetConfig { + request, + reply, + cancellation_generation, + }) .await - .map_err(|_| AcpRuntimeError::ClientClosed)?; - response.await.map_err(|_| AcpRuntimeError::ClientClosed)? + .map_err(|_| sdk_error(AcpRuntimeError::ClientClosed))?; + response + .await + .map_err(|_| sdk_error(AcpRuntimeError::ClientClosed))? } async fn cancel( @@ -1114,18 +1140,6 @@ impl Server { Ok(wire::CloseSessionResponse::new()) } - fn sender( - &self, - session_id: &wire::SessionId, - ) -> Result, AcpRuntimeError> { - self.sessions - .lock() - .map_err(|_| AcpRuntimeError::ClientClosed)? - .get(session_id) - .map(|session| session.commands.clone()) - .ok_or_else(|| AcpRuntimeError::SessionNotFound(session_id.to_string())) - } - async fn detach_compose( &self, request: DetachComposeRequest, @@ -1202,6 +1216,7 @@ async fn session_actor(actor: SessionActor) mut mcp_events, } = actor; let mut binding = Some(binding); + let mut model_switch = model_switch::Guard::default(); loop { tokio::select! { biased; @@ -1227,8 +1242,37 @@ async fn session_actor(actor: SessionActor) eprintln!("ACP v2 prompt failed for {session_id}: {error}"); } } - Some(Command::SetConfig { request, reply }) => { - let result = set_v2_config(&adapter, &catalog, request); + Some(Command::SetConfig { request, reply, cancellation_generation }) => { + let result = async { + if handle.cancellation_handle().is_cancelled_since(cancellation_generation) { + return Err(model_switch::error("model change cancelled")); + } + if request.config_id.to_string() == super::MODEL_CONFIG_ID { + let target = request.value.as_id().ok_or_else(|| model_switch::error("selection requires an id value"))?; + let decision = model_switch.check( + (&adapter.selection().map_err(|error| model_switch::error(&error))?, cancellation_generation), + crate::provider::ModelSelection::from_id(&target.to_string()).map_err(|error| model_switch::error(&error))?, + &catalog, driver.snapshot().transcript, + request.meta.as_ref().and_then(|meta| meta.get(model_switch::META)), + )?; + if decision == model_switch::Decision::Compact { + claim_prompt(&busy).map_err(sdk_error)?; + handle.prepare_injection_turn(); + handle.start_injection_turn(); + let compacted = compact_for_switch( + &session_id, &integration, &handle, &mut driver, &sink, + cancellation_generation, &activity, + ).await; + handle.stop_injection_turn(); + busy.store(false, Ordering::Release); + compacted?; + } + } + if handle.cancellation_handle().is_cancelled_since(cancellation_generation) { + return Err(model_switch::error("model change cancelled")); + } + set_v2_config(&adapter, &catalog, request).map_err(sdk_error) + }.await; let _ = reply.send(result); } Some(Command::Close { reply }) => { @@ -1395,7 +1439,7 @@ async fn prepare_prompt( .await; integration.finish_prompt(session_id); handle.stop_injection_turn(); - result + result.map(|_| ()) } #[async_trait] @@ -1572,6 +1616,52 @@ where } } +/// Uses the normal manual compactor and lifecycle, while the actor retains the +/// original selection. Returning success is the only path to publishing a switch. +async fn compact_for_switch( + session_id: &wire::SessionId, + integration: &AcpIntegration, + handle: &AcpSessionHandle, + driver: &mut LoopDriver, + sink: &ResponseReplacementSink, + cancellation_generation: u64, + activity: &SessionActivity, +) -> Result<(), agent_client_protocol::Error> { + // Only a successful, persisted standard compaction consumes this marker. + // Runtime telemetry cannot acknowledge a pending model switch. + let marker = model_switch::compact_marker(); + let marker_id = marker.id.clone(); + driver + .submit_input(vec![marker]) + .map_err(|error| sdk_error(map_loop_error(session_id, &error)))?; + let reason = run_active_turn( + session_id, + integration, + handle, + driver, + sink, + cancellation_generation, + None, + activity, + ExecutionOrigin::Prompt, + ) + .await + .map_err(sdk_error)?; + if !model_switch::compaction_completed( + &reason, + &marker_id, + &driver.snapshot().transcript, + handle + .cancellation_handle() + .is_cancelled_since(cancellation_generation), + ) { + return Err(model_switch::error( + "compaction did not complete; model unchanged", + )); + } + Ok(()) +} + #[allow(clippy::too_many_arguments)] async fn run_active_turn( session_id: &wire::SessionId, @@ -1583,7 +1673,7 @@ async fn run_active_turn( structured: Option<(&TaskManagerHandle, &BackgroundJobs)>, activity: &SessionActivity, origin: ExecutionOrigin, -) -> Result<(), AcpRuntimeError> { +) -> Result { activity .execute( origin, @@ -1616,7 +1706,6 @@ async fn run_active_turn( |reason| Some(reason.clone()), ) .await - .map(|_| ()) } async fn drive_autonomous( @@ -1650,7 +1739,7 @@ async fn drive_autonomous( integration.finish_prompt(session_id); handle.stop_injection_turn(); busy.store(false, Ordering::Release); - result + result.map(|_| ()) } fn error_diagnostic_notification( @@ -2107,8 +2196,7 @@ pub(crate) fn component( async move |request: wire::SetSessionConfigOptionRequest, responder, cx| { let state = Arc::clone(&state); cx.spawn(async move { - responder - .respond_with_result(state.set_config(request).await.map_err(sdk_error)) + responder.respond_with_result(state.set_config(request).await) })?; Ok(()) } @@ -3920,6 +4008,68 @@ mod tests { claim_prompt(&busy).unwrap(); } + #[tokio::test] + async fn set_config_rejects_poisoned_session_map_without_queueing_switch() { + let root = tempfile::tempdir().unwrap(); + let server = Server::new( + Runtime::new(root.path(), "gpt-5.4").unwrap(), + SessionRegistry::new(), + ); + let session_id = wire::SessionId::new("poisoned-session"); + let integration = server + .integration + .bind_session(AcpSessionBinding::new( + session_id.clone(), + SessionId::new("poisoned-session"), + RecordingSink::default(), + )) + .unwrap(); + let (commands, mut received) = mpsc::channel(1); + server.sessions.lock().unwrap().insert( + session_id.clone(), + SessionHandle { + token: 1, + commands, + integration, + busy: Arc::new(AtomicBool::new(false)), + background_jobs: BackgroundJobs::default(), + structured_completion: false, + tasks: AsyncTaskManager::new().handle(), + }, + ); + assert!( + std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + let _guard = server.sessions.lock().unwrap(); + panic!("poison session map"); + })) + .is_err() + ); + + for id in [session_id, wire::SessionId::new("missing-session")] { + let error = timeout( + Duration::from_secs(1), + server.set_config(wire::SetSessionConfigOptionRequest::new( + id, + super::super::MODEL_CONFIG_ID, + "openai-subscription:gpt-5.4-mini", + )), + ) + .await + .expect("poison must return an error without waiting for the actor") + .unwrap_err(); + assert_eq!(error.code, agent_client_protocol::ErrorCode::InternalError); + assert_eq!( + error.data, + Some(json!(AcpRuntimeError::ClientClosed.to_string())) + ); + } + assert!(server.sessions.is_poisoned()); + assert!(matches!( + received.try_recv(), + Err(mpsc::error::TryRecvError::Empty) + )); + } + #[test] fn sdk_error_marks_only_missing_credentials_as_authentication_required() { let missing_detail = "openrouter_auth_required: set OPENROUTER_API_KEY or run `kit auth login openrouter` before using the OpenRouter provider"; @@ -4062,8 +4212,9 @@ mod tests { assert!(registry.inner.state.lock().unwrap().v2_sessions.is_empty()); assert_eq!(server.sessions.is_poisoned(), unwind); let error = server - .sender(&wire::SessionId::new("failed-publication")) - .unwrap_err(); + .prompt_route(&wire::SessionId::new("failed-publication")) + .err() + .expect("failed publication must not leave a usable route"); if unwind { assert!(matches!(error, AcpRuntimeError::ClientClosed)); } else { @@ -4154,8 +4305,16 @@ mod tests { original_actor.id() ); assert!(!original_actor.is_finished()); - assert!(server.sender(&wire::SessionId::new("original")).is_ok()); - assert!(server.sender(&wire::SessionId::new("duplicate")).is_err()); + assert!( + server + .prompt_route(&wire::SessionId::new("original")) + .is_ok() + ); + assert!( + server + .prompt_route(&wire::SessionId::new("duplicate")) + .is_err() + ); assert!( timeout(Duration::from_secs(1), duplicate_actor) .await @@ -4194,7 +4353,7 @@ mod tests { server .publish_session(&mut admission, publication, || Ok(())) .unwrap(); - let sender = server.sender(&session_id).unwrap(); + let (sender, _, _) = server.prompt_route(&session_id).unwrap(); let (reply, _ack) = oneshot::channel(); sender.try_send(Command::Close { reply }).unwrap(); assert_eq!(sender.capacity(), 0); @@ -4250,8 +4409,9 @@ mod tests { .unwrap(); let (reply, _ack) = oneshot::channel(); server - .sender(&session_id) + .prompt_route(&session_id) .unwrap() + .0 .try_send(Command::Close { reply }) .unwrap(); assert!( @@ -4709,11 +4869,141 @@ mod tests { )); } + #[derive(Clone)] + struct SwitchSummaryAdapter { + selection: SelectableAdapter, + seen: Arc>>, + outcome: TestOutcome, + interrupt: Option, + } + + #[async_trait] + impl ModelAdapter for SwitchSummaryAdapter { + type Session = TestSession; + async fn start_session(&self, _config: SessionConfig) -> Result { + self.seen + .lock() + .unwrap() + .push(self.selection.selection().unwrap().model); + Ok(TestSession { + outcome: self.outcome, + turns: Arc::new(AtomicU64::new(0)), + interrupt: self.interrupt.clone(), + }) + } + } + + #[tokio::test] + async fn model_switch_compacts_with_original_model_before_selecting_and_keeps_it_on_failure_or_cancel() + { + for mode in ["success", "failure", "cancel"] { + let integration = AcpIntegration::default(); + let sink = ResponseReplacementSink::new(RecordingSink::default()); + let session_id = wire::SessionId::new(format!("switch-{mode}")); + let loop_id = SessionId::new(format!("switch-loop-{mode}")); + let activity = native_activity(session_id.clone(), sink.clone()); + let handle = integration + .bind_session( + AcpSessionBinding::new(session_id.clone(), loop_id.clone(), sink.clone()) + .cancellation(CancellationController::new()), + ) + .unwrap(); + let selection = + SelectableAdapter::new(ProviderKind::OpenAiSubscription, "gpt-5.4").unwrap(); + let seen = Arc::new(Mutex::new(Vec::new())); + let summary_adapter = SwitchSummaryAdapter { + selection: selection.clone(), + seen: seen.clone(), + outcome: if mode == "failure" { + TestOutcome::ProviderError + } else { + TestOutcome::Content + }, + interrupt: (mode == "cancel").then(|| handle.clone()), + }; + let compactor = crate::compaction::automatic( + summary_adapter, + Default::default(), + None, + loop_id.clone(), + ) + .unwrap(); + let mut driver = Agent::builder() + .model(TestAdapter { + outcome: TestOutcome::Content, + turns: Arc::new(AtomicU64::new(0)), + interrupt: None, + }) + .mutator(compactor) + .cancellation(handle.cancellation_handle()) + .transcript(vec![ + Item::text(ItemKind::User, "older content ".repeat(20_000)), + Item::text(ItemKind::Assistant, "recent response"), + ]) + .build() + .unwrap() + .start(SessionConfig::new(loop_id).without_cache()) + .await + .unwrap(); + handle.prepare_injection_turn(); + handle.start_injection_turn(); + let generation = handle.cancellation_handle().generation(); + let compacted = compact_for_switch( + &session_id, + &integration, + &handle, + &mut driver, + &sink, + generation, + &activity, + ) + .await; + assert_eq!(selection.selection().unwrap().model, "gpt-5.4"); + assert_eq!(*seen.lock().unwrap(), ["gpt-5.4"]); + let catalog = [crate::provider::ModelGroup { + provider: ProviderKind::OpenAiSubscription, + models: vec!["gpt-5.4-mini".into()], + context_windows: Default::default(), + }]; + let result = compacted.and_then(|()| { + set_v2_config( + &selection, + &catalog, + wire::SetSessionConfigOptionRequest::new( + session_id, + super::super::MODEL_CONFIG_ID, + "openai-subscription:gpt-5.4-mini", + ), + ) + .map_err(sdk_error) + }); + assert_eq!(result.is_ok(), mode == "success", "{mode}: {result:?}"); + assert_eq!( + selection.selection().unwrap().model, + if mode == "success" { + "gpt-5.4-mini" + } else { + "gpt-5.4" + } + ); + if mode == "success" { + assert!( + driver + .snapshot() + .transcript + .iter() + .any(crate::compaction::is_compaction_summary) + ); + } + } + } + #[test] fn v2_config_mapping_uses_v2_ids_categories_and_values() { let current = crate::provider::ModelSelection::new(crate::ProviderKind::OpenRouter, "test-model"); let catalog = [crate::provider::ModelGroup { + context_windows: Default::default(), provider: crate::ProviderKind::OpenRouter, models: vec!["test-model".into(), "other-model".into()], }]; diff --git a/src/provider/adapter.rs b/src/provider/adapter.rs index c669382..8e098b6 100644 --- a/src/provider/adapter.rs +++ b/src/provider/adapter.rs @@ -121,6 +121,8 @@ pub(super) fn valid_model_id(value: &str) -> bool { pub struct ModelGroup { pub provider: ProviderKind, pub models: Vec, + /// Provider-reported windows only; missing entries are unknown. + pub context_windows: std::collections::HashMap, } #[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq, ValueEnum)] @@ -232,10 +234,10 @@ impl SelectableAdapter { .map_err(|_| "session selection lock is poisoned".into()) } - async fn discovered_openai_models(&self) -> Vec { + async fn discovered_openai_models(&self) -> DiscoveredModels { let config = match SubscriptionConfig::new(OPENAI_FALLBACK[0].to_string()) { Ok(config) => config.with_credential_storage(self.credential_storage.clone()), - Err(_) => return Vec::new(), + Err(_) => return DiscoveredModels::default(), }; let adapter = match OpenAiSubscriptionAdapter::new_with_reasoning_effort_and_catalog( config, @@ -243,12 +245,15 @@ impl SelectableAdapter { self.openai_model_catalog.clone(), ) { Ok(adapter) => adapter, - Err(_) => return Vec::new(), + Err(_) => return DiscoveredModels::default(), }; adapter .model_catalog() .await - .map(|catalog| catalog.visible_models().to_vec()) + .map(|catalog| DiscoveredModels { + models: catalog.visible_models().to_vec(), + context_windows: catalog.context_windows().clone(), + }) .unwrap_or_default() } @@ -1056,13 +1061,16 @@ fn openai_models(discovered: Vec, current: &ModelSelection) -> Vec Vec { match SelectableAdapter::new(current.provider, current.model.clone()) { Ok(adapter) => adapter.model_catalog(current).await, - Err(_) => model_catalog_with_openai(current, std::future::ready(Vec::new())).await, + Err(_) => { + model_catalog_with_openai(current, std::future::ready(DiscoveredModels::default())) + .await + } } } async fn model_catalog_with_openai( current: &ModelSelection, - openai_catalog: impl std::future::Future>, + openai_catalog: impl std::future::Future, ) -> Vec { let openrouter_url = match std::env::var_os("OPENROUTER_BASE_URL") { None => catalog_models_url(None), @@ -1074,28 +1082,32 @@ async fn model_catalog_with_openai( if same_catalog { let models = fetch_model_ids(&public_url) .await - .unwrap_or_else(|_| openrouter_fallback()); + .unwrap_or_else(|_| fallback_catalog()); (models.clone(), models) } else { let openrouter_catalog = async { match openrouter_url { Some(url) => fetch_model_ids(&url) .await - .unwrap_or_else(|_| openrouter_fallback()), - None => openrouter_fallback(), + .unwrap_or_else(|_| fallback_catalog()), + None => fallback_catalog(), } }; let speakeasy_catalog = async { fetch_model_ids(&public_url) .await - .unwrap_or_else(|_| openrouter_fallback()) + .unwrap_or_else(|_| fallback_catalog()) }; tokio::join!(openrouter_catalog, speakeasy_catalog) } }; - let (discovered_openai, (mut openrouter, mut speakeasy)) = - tokio::join!(openai_catalog, other_catalogs); - let openai = openai_models(discovered_openai, current); + let (discovered_openai, (openrouter, speakeasy)) = tokio::join!(openai_catalog, other_catalogs); + let openai_windows = discovered_openai.context_windows; + let openrouter_windows = openrouter.context_windows; + let speakeasy_windows = speakeasy.context_windows; + let mut openrouter = openrouter.models; + let mut speakeasy = speakeasy.models; + let openai = openai_models(discovered_openai.models, current); let current_is_valid = valid_model_id(¤t.model); if current_is_valid && current.provider == ProviderKind::OpenRouter @@ -1137,14 +1149,17 @@ async fn model_catalog_with_openai( ModelGroup { provider: ProviderKind::OpenAiSubscription, models: openai, + context_windows: openai_windows, }, ModelGroup { provider: ProviderKind::OpenRouter, models: openrouter, + context_windows: openrouter_windows, }, ModelGroup { provider: ProviderKind::Speakeasy, models: speakeasy, + context_windows: speakeasy_windows, }, ] } @@ -1156,7 +1171,20 @@ fn openrouter_fallback() -> Vec { .collect() } -async fn fetch_model_ids(url: &str) -> Result, String> { +#[derive(Clone, Default)] +struct DiscoveredModels { + models: Vec, + context_windows: std::collections::HashMap, +} + +fn fallback_catalog() -> DiscoveredModels { + DiscoveredModels { + models: openrouter_fallback(), + ..Default::default() + } +} + +async fn fetch_model_ids(url: &str) -> Result { let client = reqwest::Client::builder() .redirect(reqwest::redirect::Policy::none()) .connect_timeout(Duration::from_secs(5)) @@ -1183,6 +1211,10 @@ async fn fetch_model_ids(url: &str) -> Result, String> { } let value: Value = serde_json::from_slice(&body).map_err(|_| "model catalog is not valid JSON".to_string())?; + parse_discovered_models(&value) +} + +fn parse_discovered_models(value: &Value) -> Result { let entries = value .get("data") .and_then(Value::as_array) @@ -1190,17 +1222,39 @@ async fn fetch_model_ids(url: &str) -> Result, String> { if entries.len() > MAX_MODELS { return Err("model catalog has too many entries".into()); } - Ok(entries + let models: Vec = entries .iter() .filter_map(|entry| entry.get("id").and_then(Value::as_str)) .filter(|id| valid_model_id(id)) .take(MAX_SELECTOR_MODELS) .map(str::to_string) - .collect()) + .collect(); + let context_windows = models + .iter() + .filter_map(|id| parse_context_window(value, id).map(|window| (id.clone(), window))) + .collect(); + Ok(DiscoveredModels { + models, + context_windows, + }) } #[cfg(test)] mod tests { + #[test] + fn model_switch_catalog_retains_only_reported_positive_windows() { + let catalog = super::parse_discovered_models(&serde_json::json!({"data": [ + {"id": "known", "context_length": 200000}, + {"id": "unknown"}, {"id": "zero", "context_length": 0}, + {"id": "bad model", "context_length": 100} + ]})) + .unwrap(); + assert_eq!(catalog.models, ["known", "unknown", "zero"]); + assert_eq!(catalog.context_windows.len(), 1); + assert_eq!(catalog.context_windows["known"], 200000); + assert!(super::fallback_catalog().context_windows.is_empty()); + } + use std::{ collections::BTreeMap, io::{Read, Write}, diff --git a/src/provider/chatgpt.rs b/src/provider/chatgpt.rs index 99ccbe1..e08ac46 100644 --- a/src/provider/chatgpt.rs +++ b/src/provider/chatgpt.rs @@ -101,6 +101,10 @@ pub(crate) struct SubscriptionModelCatalog { } impl SubscriptionModelCatalog { + pub(crate) fn context_windows(&self) -> &HashMap { + &self.context_windows + } + pub(crate) fn visible_models(&self) -> &[String] { &self.visible_models } diff --git a/src/tui/app.rs b/src/tui/app.rs index dc4b5f9..3e6c86d 100644 --- a/src/tui/app.rs +++ b/src/tui/app.rs @@ -208,6 +208,15 @@ pub struct ModelChoice { pub model: String, } +pub struct ModelSwitch { + pub id: u64, + pub choice: ModelChoice, + pub save_defaults: bool, + pub warning: Option, + pub selected: usize, + pub cancelling: bool, +} + pub struct ModelDialog { pub query: String, pub selected: usize, @@ -326,6 +335,7 @@ pub enum Action { choice: ModelChoice, save_defaults: bool, }, + ConfirmModelSwitch(crate::protocols::acp::model_switch::Decision), SelectEffort { effort: String, save_defaults: bool, @@ -640,6 +650,8 @@ pub struct App { pub model: String, pub model_choices: Vec, pub model_dialog: Option, + pub model_switch: Option, + next_model_switch: u64, pub reasoning_effort: String, pub effort_choices: Vec, pub effort_dialog: Option, @@ -945,6 +957,8 @@ impl App { model, model_choices: Vec::new(), model_dialog: None, + model_switch: None, + next_model_switch: 0, reasoning_effort: "default".into(), effort_choices: Vec::new(), effort_dialog: None, @@ -2195,6 +2209,7 @@ impl App { /// history and diagnostics remain useful, while transcript-derived state /// starts empty. pub fn start_session(&mut self, session_id: String) { + self.model_switch = None; self.cancel_steer_edit(); self.selected_steer = None; self.queue_focused = false; @@ -2526,6 +2541,9 @@ impl App { } pub fn paste(&mut self, text: &str) { + if self.model_switch.is_some() { + return; + } // An explicit bracketed paste is not part of the unbracketed key-burst heuristic. self.last_key = None; if let Some(rename) = self @@ -3118,11 +3136,64 @@ impl App { Action::None } + pub fn begin_model_switch(&mut self, choice: ModelChoice, save_defaults: bool) -> Option { + if self.model_switch.is_some() || self.working() { + self.note("wait for the current operation before changing models"); + return None; + } + self.next_model_switch = self.next_model_switch.wrapping_add(1); + let id = self.next_model_switch; + self.model_switch = Some(ModelSwitch { + id, + choice, + save_defaults, + warning: None, + selected: 2, + cancelling: false, + }); + Some(id) + } + /// Applies a key press, returning work for the event loop. pub fn handle_key(&mut self, key: KeyEvent) -> Action { if key.kind != KeyEventKind::Press { return Action::None; } + if let Some(pending) = self.model_switch.as_mut() { + use crate::protocols::acp::model_switch::Decision; + let cancel = key.code == KeyCode::Esc + || (key.code == KeyCode::Char('c') + && key.modifiers.contains(KeyModifiers::CONTROL)); + if pending.warning.is_none() { + if cancel { + if pending.cancelling { + return if key.code == KeyCode::Char('c') { + Action::Quit + } else { + Action::None + }; + } + pending.cancelling = true; + return Action::Cancel; + } + return Action::None; + } + if cancel { + self.model_switch = None; + return Action::Redraw; + } + match key.code { + KeyCode::Up => pending.selected = pending.selected.saturating_sub(1), + KeyCode::Down | KeyCode::Tab => pending.selected = (pending.selected + 1) % 3, + KeyCode::Enter => match pending.selected { + 0 => return Action::ConfirmModelSwitch(Decision::Continue), + 1 => return Action::ConfirmModelSwitch(Decision::Compact), + _ => self.model_switch = None, + }, + _ => {} + } + return Action::Redraw; + } if self.session_dialog.is_some() { // Terminals without bracketed paste deliver a paste as a key burst, so // the arrival gap is the only thing separating it from typing. @@ -3583,6 +3654,9 @@ impl App { } pub fn handle_mouse(&mut self, mouse: MouseEvent) -> Action { + if self.model_switch.is_some() { + return Action::None; + } match mouse.kind { MouseEventKind::ScrollUp if self.mouse_in_agents(mouse.column, mouse.row) => { self.scroll_agents_by(-3); @@ -3958,6 +4032,82 @@ mod tests { ) } + #[test] + fn model_switch_dialog_actions_preserve_input_and_selection() { + use crate::protocols::acp::model_switch::{Decision, Warning}; + for (selected, expected) in [ + (0, Some(Decision::Continue)), + (1, Some(Decision::Compact)), + (2, None), + ] { + let mut app = app(); + app.paste("unsent draft"); + let old_model = app.model.clone(); + let choice = model_choice("openrouter", "target"); + app.begin_model_switch(choice, false).unwrap(); + let pending = app.model_switch.as_mut().unwrap(); + pending.warning = Some(Warning { + token: 1, + guarded_tokens: "120".into(), + target_window: 150, + }); + pending.selected = selected; + app.paste("must not enter composer"); + let result = app.handle_key(press(KeyCode::Enter)); + match expected { + Some(expected) => assert!( + matches!(result, Action::ConfirmModelSwitch(actual) if actual == expected) + ), + None => assert!(app.model_switch.is_none()), + } + assert_eq!(app.editor.text(), "unsent draft"); + assert_eq!(app.model, old_model); + assert!(app.blocks.is_empty()); + } + } + + #[test] + fn model_switch_pending_cancel_and_stale_session_do_not_reuse_operation() { + let mut app = app(); + let first = app + .begin_model_switch(model_choice("openrouter", "target"), false) + .unwrap(); + assert!( + app.begin_model_switch(model_choice("openrouter", "other"), false) + .is_none() + ); + assert!(matches!( + app.handle_key(press(KeyCode::Esc)), + Action::Cancel + )); + assert!(app.model_switch.as_ref().unwrap().cancelling); + assert!(matches!(app.handle_key(press(KeyCode::Esc)), Action::None)); + app.start_session("new-session".into()); + assert!(app.model_switch.is_none()); + let next = app + .begin_model_switch(model_choice("openrouter", "other"), false) + .unwrap(); + assert_ne!(first, next); + } + + #[test] + fn model_switch_dialog_escape_cancels_without_a_request() { + use crate::protocols::acp::model_switch::Warning; + let mut app = app(); + app.begin_model_switch(model_choice("openrouter", "target"), false) + .unwrap(); + app.model_switch.as_mut().unwrap().warning = Some(Warning { + token: 1, + guarded_tokens: "120".into(), + target_window: 150, + }); + assert!(matches!( + app.handle_key(press(KeyCode::Esc)), + Action::Redraw + )); + assert!(app.model_switch.is_none()); + } + #[test] fn typed_eligible_at_opens_picker_but_paste_and_email_do_not() { let mut typed = app(); diff --git a/src/tui/mod.rs b/src/tui/mod.rs index d4c33ab..897a104 100644 --- a/src/tui/mod.rs +++ b/src/tui/mod.rs @@ -58,7 +58,9 @@ use wire::{ use crate::{ events::{self, EVENTS_ENV}, - protocols::acp::{FileSearchRequest, MODEL_CONFIG_ID, REASONING_EFFORT_CONFIG_ID}, + protocols::acp::{ + FileSearchRequest, MODEL_CONFIG_ID, REASONING_EFFORT_CONFIG_ID, model_switch, + }, tools::mcp::CredentialStorage, }; @@ -67,6 +69,91 @@ use app::{ UserImage, }; +struct ModelSwitchCompletion { + generation: u64, + operation: u64, + response: Result, +} + +fn prepare_model_switch_request( + app: &mut App, + action: Action, + session_id: wire::SessionId, +) -> Result<(u64, SetSessionConfigOptionRequest), agent_client_protocol::Error> { + let confirmation = match action { + Action::SelectModel { + choice, + save_defaults, + } => { + app.begin_model_switch(choice, save_defaults) + .ok_or_else(|| { + agent_client_protocol::util::internal_error( + "wait for the current operation before changing models", + ) + })?; + None + } + Action::ConfirmModelSwitch(decision) => { + let warning = app + .model_switch + .as_ref() + .and_then(|pending| pending.warning.as_ref()) + .ok_or_else(|| { + agent_client_protocol::util::internal_error( + "no model-switch warning to confirm", + ) + })?; + Some( + serde_json::to_value(model_switch::Confirmation { + token: warning.token, + action: decision, + }) + .map_err(agent_client_protocol::Error::into_internal_error)?, + ) + } + _ => { + return Err(agent_client_protocol::util::internal_error( + "invalid model-switch action", + )); + } + }; + let pending = app + .model_switch + .as_mut() + .ok_or_else(|| agent_client_protocol::util::internal_error("no pending model switch"))?; + let mut request = + SetSessionConfigOptionRequest::new(session_id, MODEL_CONFIG_ID, pending.choice.id.as_str()); + if let Some(confirmation) = confirmation { + request.meta = Some(serde_json::Map::from_iter([( + model_switch::META.into(), + confirmation, + )])); + // Keep the warning available if preparing its confirmation fails. + pending.warning = None; + } + Ok((pending.id, request)) +} + +fn take_model_switch_completion( + app: &mut App, + route: &Arc>, + generation: u64, + operation: u64, +) -> Result, agent_client_protocol::Error> { + let route = route.lock().map_err(|_| { + agent_client_protocol::util::internal_error("active session route poisoned") + })?; + if route.generation != generation + || app + .model_switch + .as_ref() + .is_none_or(|pending| pending.id != operation) + { + return Ok(None); + } + Ok(app.model_switch.take()) +} + /// Animation and elapsed-time refresh interval. const TICK: Duration = Duration::from_millis(90); /// Terminal events or queued updates applied per frame. @@ -1135,6 +1222,7 @@ pub async fn run_with_reasoning_effort_and_openrouter_key( }; refresh_config_state(&mut app, Some(&config_options)); let mut saved_model_default = current_model_choice(Some(&config_options)); + let (switch_tx, mut switch_rx) = mpsc::unbounded_channel::(); if let Ok(mut active) = transition_session.lock() { active.id = active_session_id.clone(); } @@ -1452,38 +1540,25 @@ pub async fn run_with_reasoning_effort_and_openrouter_key( } } Action::Close => return Ok(()), - Action::SelectModel { choice, save_defaults } => { - let response = connection.send_request( - SetSessionConfigOptionRequest::new( - session_id.clone(), - MODEL_CONFIG_ID, - choice.id.as_str(), - ), - ).block_task().await; - match response { - Ok(response) => { - app.usage = None; - refresh_config_state( - &mut app, - Some(&response.config_options), - ); - app.note(format!("model changed to {} via {}", choice.model, choice.provider)); - if save_defaults { - let saved = { - let choice = choice.clone(); - tokio::task::spawn_blocking(move || save_model_defaults(&choice)).await.map_err(|error| error.to_string()).and_then(|result| result) - }; - match saved { - Ok(()) => { - saved_model_default = Some(choice.clone()); - app.note(config_save_message("model defaults")); - } - Err(error) => app.note(format!("model changed, but defaults were not saved: {error}")), - } - } + action @ (Action::SelectModel { .. } | Action::ConfirmModelSwitch(_)) => { + // A poisoned route may pair a session ID with the wrong generation. + // Stop explicitly rather than send a request with untrusted correlation. + let generation = transition_session.lock().map_err(|_| { + agent_client_protocol::util::internal_error("active session route poisoned") + })?.generation; + let (operation, request) = match prepare_model_switch_request(&mut app, action, session_id.clone()) { + Ok(request) => request, + Err(error) => { + app.note(format!("could not change model: {}", error.message)); + continue; } - Err(error) => app.note(format!("model change failed: {}", error.message)), - } + }; + let connection = connection.clone(); + let completed = switch_tx.clone(); + tokio::task::spawn_local(async move { + let response = connection.send_request(request).block_task().await; + let _ = completed.send(ModelSwitchCompletion { generation, operation, response }); + }); } Action::SelectEffort { effort, save_defaults } => { let response = connection @@ -1677,6 +1752,37 @@ pub async fn run_with_reasoning_effort_and_openrouter_key( Action::None | Action::Redraw => {} } }, + Some(completion) = switch_rx.recv() => { + let Some(mut pending) = take_model_switch_completion(&mut app, &transition_session, completion.generation, completion.operation)? else { continue; }; + match completion.response { + Ok(response) => { + let choice = pending.choice; + let save_defaults = pending.save_defaults; + app.usage = None; + refresh_config_state(&mut app, Some(&response.config_options)); + app.note(format!("model changed to {} via {}", choice.model, choice.provider)); + if save_defaults { + let saved = { + let choice = choice.clone(); + tokio::task::spawn_blocking(move || save_model_defaults(&choice)).await.map_err(|error| error.to_string()).and_then(|result| result) + }; + match saved { + Ok(()) => { saved_model_default = Some(choice); app.note(config_save_message("model defaults")); } + Err(error) => app.note(format!("model changed, but defaults were not saved: {error}")), + } + } + } + Err(error) => { + let warning = error.data.as_ref().and_then(|data| data.get(model_switch::META)).and_then(|value| serde_json::from_value::(value.clone()).ok()); + if let Some(warning) = warning.filter(|_| !pending.cancelling) { + pending.warning = Some(warning); + app.model_switch = Some(pending); + } else { + app.note(format!("model change failed: {}", error.message)); + } + } + } + }, update = updates_rx.recv() => match update { Some(update) => apply_pending_updates( &mut app, @@ -1941,6 +2047,9 @@ fn handle(app: &mut App, event: Event) -> Action { Event::Key(key) => app.handle_key(key), Event::Mouse(mouse) => app.handle_mouse(mouse), Event::Paste(text) => { + if app.model_switch.is_some() { + return Action::None; + } if app.queue_focused && !app.session_rename_active() { return Action::None; } @@ -2941,6 +3050,151 @@ mod tests { ); } + #[test] + fn model_switch_request_rejects_invalid_actions_and_missing_warnings() { + use crate::protocols::acp::model_switch::{Confirmation, Decision, META, Warning}; + let mut app = App::new( + PathBuf::from("/tmp"), + "openrouter".into(), + "original".into(), + String::new(), + ); + let session = wire::SessionId::new("first"); + assert!( + super::prepare_model_switch_request(&mut app, Action::None, session.clone()).is_err() + ); + assert!( + super::prepare_model_switch_request( + &mut app, + Action::ConfirmModelSwitch(Decision::Continue), + session.clone() + ) + .is_err() + ); + assert!(app.model_switch.is_none()); + let choice = ModelChoice { + id: "openrouter:target".into(), + provider: "openrouter".into(), + model: "target".into(), + }; + let (operation, request) = super::prepare_model_switch_request( + &mut app, + Action::SelectModel { + choice, + save_defaults: true, + }, + session.clone(), + ) + .unwrap(); + assert!(request.meta.is_none()); + assert!( + super::prepare_model_switch_request( + &mut app, + Action::ConfirmModelSwitch(Decision::Compact), + session.clone() + ) + .is_err() + ); + assert_eq!(app.model_switch.as_ref().unwrap().id, operation); + for decision in [Decision::Continue, Decision::Compact] { + app.model_switch.as_mut().unwrap().warning = Some(Warning { + token: 42, + guarded_tokens: "120000".into(), + target_window: 150000, + }); + let (confirmed_operation, request) = super::prepare_model_switch_request( + &mut app, + Action::ConfirmModelSwitch(decision), + session.clone(), + ) + .unwrap(); + let confirmation: Confirmation = + serde_json::from_value(request.meta.unwrap()[META].clone()).unwrap(); + assert_eq!(confirmation.token, 42); + assert_eq!(confirmation.action, decision); + assert_eq!(confirmed_operation, operation); + let pending = app.model_switch.as_ref().unwrap(); + assert!(pending.warning.is_none()); + assert!(pending.save_defaults); + assert_eq!(pending.choice.id, "openrouter:target"); + assert_eq!(app.model, "original"); + } + } + + #[test] + fn model_switch_completion_requires_current_session_and_operation() { + let mut app = App::new( + PathBuf::from("/tmp"), + "openrouter".into(), + "original".into(), + String::new(), + ); + let route = Arc::new(Mutex::new(ActiveSessionRoute { + id: "first".into(), + generation: 0, + })); + let choice = ModelChoice { + id: "openrouter:target".into(), + provider: "openrouter".into(), + model: "target".into(), + }; + let operation = app.begin_model_switch(choice, false).unwrap(); + assert!( + super::take_model_switch_completion(&mut app, &route, 0, operation + 1) + .unwrap() + .is_none() + ); + transition_route(&route, "second".into()); + assert!( + super::take_model_switch_completion(&mut app, &route, 0, operation) + .unwrap() + .is_none() + ); + assert!(app.model_switch.is_some()); + assert!( + super::take_model_switch_completion(&mut app, &route, 1, operation) + .unwrap() + .is_some() + ); + assert!( + super::take_model_switch_completion(&mut app, &route, 1, operation) + .unwrap() + .is_none() + ); + assert_eq!(app.model, "original"); + } + + #[test] + fn model_switch_completion_reports_poisoned_route_without_discarding_switch() { + let mut app = App::new( + PathBuf::from("/tmp"), + "openrouter".into(), + "original".into(), + String::new(), + ); + let route = Arc::new(Mutex::new(ActiveSessionRoute { + id: "first".into(), + generation: 0, + })); + let operation = app + .begin_model_switch( + ModelChoice { + id: "openrouter:target".into(), + provider: "openrouter".into(), + model: "target".into(), + }, + false, + ) + .unwrap(); + let _ = std::panic::catch_unwind(|| { + let _guard = route.lock().unwrap(); + panic!("poison session route"); + }); + assert!(super::take_model_switch_completion(&mut app, &route, 0, operation).is_err()); + assert_eq!(app.model_switch.as_ref().unwrap().id, operation); + assert_eq!(app.model, "original"); + } + #[test] fn queued_updates_from_previous_session_generations_are_dropped() { let route = Arc::new(Mutex::new(ActiveSessionRoute { diff --git a/src/tui/ui.rs b/src/tui/ui.rs index 700990a..ea0846b 100644 --- a/src/tui/ui.rs +++ b/src/tui/ui.rs @@ -120,7 +120,9 @@ pub fn draw(frame: &mut Frame<'_>, app: &mut App, images: &mut ImageRuntime) { draw_status(frame, app, status); (prompt, viewport, false) }; - if app.file_picker.is_some() { + if let Some(pending) = &app.model_switch { + draw_model_switch_dialog(frame, pending); + } else if app.file_picker.is_some() { draw_file_picker(frame, app, prompt_area, prompt_viewport, picker_below); } else if app.session_dialog.is_some() { draw_session_dialog(frame, app); @@ -530,6 +532,66 @@ fn draw_session_dialog(frame: &mut Frame<'_>, app: &App) { } } +fn draw_model_switch_dialog(frame: &mut Frame<'_>, pending: &super::app::ModelSwitch) { + let outer = frame.area(); + let width = outer.width.min(76); + let height = outer.height.min(12); + let area = Rect::new( + outer.x + outer.width.saturating_sub(width) / 2, + outer.y + outer.height.saturating_sub(height) / 2, + width, + height, + ); + let panel = Panel::bordered().title(" model context warning "); + let inner = panel.inner(area); + let mut lines = match &pending.warning { + Some(warning) => vec![ + Line::from(format!( + "{} via {}", + pending.choice.model, pending.choice.provider + )), + Line::from(format!( + "Estimated {} / {} tokens.", + warning.guarded_tokens, warning.target_window + )), + Line::from("At least 80% of the target context is occupied."), + Line::from("Compact uses the current model before switching."), + Line::from(""), + ], + None => vec![ + Line::from(if pending.cancelling { + "Cancelling model change…" + } else { + "Preparing model change…" + }), + Line::from("Esc / Ctrl+C cancels; input is preserved."), + ], + }; + if pending.warning.is_some() { + for (index, label) in ["Continue anyway", "Compact", "Cancel"].iter().enumerate() { + lines.push(Line::from(Span::styled( + format!( + "{} {label}", + if index == pending.selected { + "›" + } else { + " " + } + ), + if index == pending.selected { + theme::accent() + } else { + theme::text() + }, + ))); + } + lines.push(Line::from("↑/↓ select · enter confirm · esc cancel")); + } + frame.render_widget(Clear, area); + frame.render_widget(panel, area); + frame.render_widget(Paragraph::new(lines), inner); +} + fn draw_effort_dialog(frame: &mut Frame<'_>, app: &App) { let outer = frame.area(); let width = outer.width.min(48); @@ -2909,6 +2971,45 @@ mod tests { .join("\n") } + #[test] + fn model_switch_popup_renders_all_choices_and_defaults_to_cancel() { + let mut app = App::new( + PathBuf::from("/tmp"), + "openrouter".into(), + "original".into(), + String::new(), + ); + app.begin_model_switch( + crate::tui::app::ModelChoice { + id: "openrouter:target".into(), + provider: "openrouter".into(), + model: "target".into(), + }, + false, + ) + .unwrap(); + app.model_switch.as_mut().unwrap().warning = + Some(crate::protocols::acp::model_switch::Warning { + token: 1, + guarded_tokens: "120000".into(), + target_window: 150000, + }); + let screen = render(&mut app, 90, 24); + for label in [ + "model context warning", + "Estimated 120000 / 150000 tokens.", + "80%", + "Continue anyway", + "Compact", + "› Cancel", + ] { + assert!(screen.contains(label), "missing {label}"); + } + assert!(!screen.contains("margin")); + // Small terminals must not panic when a dialog is clipped. + let _ = render(&mut app, 1, 1); + } + #[test] fn storage_warning_is_visible_without_logs_and_clears_on_recovery() { let mut app = App::new(