diff --git a/docs/architecture-audit-2026-09-23/ModelSelectionImplementation.md b/docs/architecture-audit-2026-09-23/ModelSelectionImplementation.md new file mode 100644 index 0000000000..ab19856306 --- /dev/null +++ b/docs/architecture-audit-2026-09-23/ModelSelectionImplementation.md @@ -0,0 +1,153 @@ +# 模型选择数据流修复 + +日期:2026-09-23。对应 [原始数据流审计](./ModelSelectionDataFlow.md)。用户要求多 agent 实施,并保持已有选择入口、布局与操作步骤。 + +## 权威与完成条件 + +| 值 | 权威 | 本轮不变量 | +| --- | --- | --- | +| 账户目录 | KeyService / credentials 持久化 | 发现结果原子写入;晚到结果不能覆盖更新后的凭据或目录 | +| 用户偏好 | enabled、手动 alias、显式 family default | 刷新保留最新用户选择;供应商默认与用户覆盖分开 | +| 可选账户/模型 | KeyInfo 与共享前端目录投影 | 空 enabled 保持为空;健康状态不被 executor adapter 提升;wire 元数据优先 | +| 会话身份 | session DB 的 model/account pair | 桌面与手机使用同一写入用例;确认成功时旧 runtime 已失效 | +| 当前轮执行身份 | 已入场 turn 捕获的 runtime 与对应 lease | 已开始的回复继续使用原 provider;取消/清理保持同一 lease;新选择用于后续入场的轮次 | +| creator 默认、recent | 成功接受的完整选择 | 保存失败或作用域已切换时,不记入默认/最近使用 | +| 上下文窗口 | 当前 session 的运行时遥测,然后当前账户目录 | 不读取另一会话的 creator 默认来猜测当前窗口 | + +Market 保持独立权限源。内置产品目录仍是支持资料,不表示账户已获实时授权。既有会话 immutable source/engine 限制保留。 + +```mermaid +flowchart TD + A[供应商发现] --> B[按凭据和目录版本检查的原子写入] + U[用户启用及默认偏好] --> B + B --> K[KeyService / KeyInfo] + K --> D[共享桌面目录投影] + K --> M[手机目录投影] + D --> P[完整 model/account 选择] + M --> P + P --> Q[每会话有序保存 / 即时显示] + Q --> S[共享 patch: 校验、DB、runtime 失效、事件] + S --> C[保存成功后记录默认与最近使用] + S --> N[后续发送读取新身份并初始化] +``` + +## 生产边界与修复 + +| Line | Element | Verdict | Reason | Suggested change | +| --- | --- | --- | --- | --- | +| `key_store/service/model_catalog.rs:21` | discovery 写入 | fix | 原刷新丢 variants/defaults、重写 enabled;读网络请求起点的旧偏好会丢更新 | 已由原子 discovery 命令读取锁内最新偏好,并校验两个 generation | +| `key_store/types.rs`、`commands/crud/models.rs:383`、`save_default_variants.rs:6` | 默认档位来源 | fix | provider 默认和显式选择原来混在一张表;修改某家族会把其余投影默认固定成用户覆盖 | 分开保存 discovered default 与 user override;三个编辑入口只提交所改家族;endpoint/目录重置清除旧 discovery | +| `accountModelCatalog.ts:127`、`crud/models.rs:541`、mobile adapter | 可选目录 | fix | 不同入口各自解释空启用、variant-only 与账户健康 | 共享可选规则、wire-first 变体解析和 registry 兼容关系 | +| `session_directory/patch.rs:491` | 身份更新事务 | fix | 手机绕过 runtime invalidation;失败或取消可留下 DB/runtime 分裂 | 桌面手机共用写入边界、短身份锁与 admission guard | +| `state/session_identity.rs:15`、session init callers | 初始化并发 | fix | 旧初始化可能在改选后安装旧 runtime 或写回旧身份 | 同会话准备/初始化与 patch 串行;现有会话的完整 pair 优先于 seed | +| `sessionPatchQueue.ts:58` | 连续选择和回滚 | fix | 完整旧快照回滚可能覆盖新选择或拆散 model/account | 有序保存、按原子字段组识别回声/外部更新、仅回滚仍持有的字段 | +| `defaultVariantSaveCoordinator.ts` | 默认档位连续保存 | fix | 多入口并发写、全量旧返回与失败快照可能覆盖较新意图 | 共享每账户队列,只提交 family delta;确认值重放 pending,成功回复不替换整行 | +| session turn entry / gateway pipeline | 执行上下文 | fix | 初始化后再次读 cache 可能拿到已失效或已更换的 runtime,lease 也可能错配 | 贯穿准备时的 runtime 和 lease;既有 admission 检查保留 | +| `ModelPill.tsx:250`、`modelSelectionCommit.ts` | 成功时机 | fix | recent/default 先于保存成功 | 保留即时关闭与即时显示;成功且仍属当前意图才记入偏好 | +| `selectionFromSession.ts:15`、`useContextUsageInfo.ts` | fallback | fix | 不同来源的 model/account、context 拼接 | 已有身份整体使用;只有完全未选模型时使用完整 creator fallback | +| `sharedLocalKeyStore.ts` | 账户 store | keep with reason | 已有共享发布与 single-flight,不需要新账户副本 | 保持 | +| Market catalog / prepared source | 权限与生命周期 | keep with reason | 与本地 key 不同的凭据和套餐生命周期 | 保持独立源;晚到准备结果不能替换新的 own-key 选择 | + +## 状态机与边缘情况 + +| 状态/事件 | 显示与写入政策 | 验证边界 | +| --- | --- | --- | +| 刷新中 | 原目录保持;同账户请求合并 | TS 刷新测试 | +| 空目录/发现失败 | 拒绝破坏性覆盖,可重试 | TS + KeyService | +| 刷新期间账户或目录变化 | 拒绝旧 generation 的提交 | KeyService 持久化测试 | +| 显式禁用全部模型 | 选择器保持空候选;管理页仍能管理 | KeyInfo/前端/mobile | +| 选择后保存中 | 保留即时菜单关闭与乐观显示;同会话保存顺序化 | palette + queue | +| 多次快速选择 | 最新意图保持显示;早到自身事件不覆盖后续选择 | queue | +| 默认档位连改/关闭管理页 | 保留原 300ms debounce;关闭时交给共享保存队列;双失败回滚确认值 | default coordinator + rendered table hook | +| 保存失败 | 回到最近确认值,保留外部更新与无关字段;显示错误,可重试 | queue + rendered ModelPill + SQLite | +| 外部 pair 与本地只共享一半字段 | 整组识别与保留,禁止拼接 | queue model/account 与 product/exec 回归 | +| 离开会话/删除/store 替换 | 晚到结果不写新作用域,不复活删除记录;在途队列结束取消订阅 | rendered ModelPill + queue + store owner 检查 | +| 正在回复时改选 | 原回复 provider 保持;下一次入场读取新 pair | captured runtime 与 shared patch 测试 | +| 已入场 A,随后 cache 清空或换为 B | 本轮仍用 A,A 的结束不能清理 B;下一轮用 B | 真实 fake-provider entry/gateway 与 lease 边界回归 | +| prepared Agent Org admission 已固定 runtime | DB 写入前拒绝变更,释放 admission 后可重试 | core runtime mutation tests | +| 手机断线/取消请求 | 等待锁时可取消;已入场 DB 写入仍完成缓存失效和通知,不因响应端消失留下半次更新 | shared patch 断线回归 | +| 后台标题/渠道/压缩初始化、启动失败 | 已有会话身份优先;旧准备工作不得回写旧 pair;失败只更新 status/timestamp | cold-init 与 status-only persistence 回归 | +| 手机列表超过上限 | 在 256 项内保留当前选择 | mobile adapter | +| 无运行时 context | 使用当前会话账户;CLI/imported 原有 unknown 规则保持 | rendered context hook | + +## 兼容性与历史数据 + +新增持久化字段 `model_catalog_generation` 和 `discovered_default_variants` 均有 serde 默认值;旧文件无须迁移。旧 `default_variants` 保守地视为显式用户选择,因为无法可靠恢复其最初来源。未删除历史账户、模型或会话数据,也未自动清理旧默认档位。 + +RPC 新目录 payload 与 `default_variant_overrides` 均为可选字段;旧 health/update 调用保持兼容,delta 保存不允许同时夹带其他账户编辑字段。回退版本会忽略新字段,但不能期待旧代码继续保证本次新增的并发与刷新不变量;若保存凭据文件则应保留备份。无数据库 schema 修改。 + +## 十层架构覆盖 + +| 层 | 结果与范围 | +| --- | --- | +| 1 编译 | frontend 类型检查、定向 lint、Rust owning-crate/应用测试;未跑全 workspace clippy | +| 2 结构/重复 | 合并目录解释与 session 写入入口;未做全库死代码清理 | +| 3 命名 | 明确 discovery default/user override、confirmed/pending、runtime cache | +| 4 语义 | enabled 空集合不代表未配置;catalog 支持不代表授权 | +| 5 默认 | session pair 整体 fallback;已有会话 seed 不覆盖;provider default 与用户覆盖分离 | +| 6 边界 | 健康由凭据边界拥有;身份由 session 写入边界拥有;runtime 为缓存 | +| 7 可理解性 | 生命周期和边缘矩阵;原审计标为修改前基线 | +| 8 Wire | schema/API payload 与 Rust DTO 一致;无凭据加入目录投影;未抓真实供应商请求 | +| 9 入口 | desktop/mobile、key/model first、Org picker、normal send/title/channel/compaction 初始化 | +| 10 resolver | 完整 pair、不从 creator/default 逐字段补齐;wire variant 优先 | + +## 性能与资源生命周期 + +| Area | Verdict | Evidence | Change or reason kept | Verification | +| --- | --- | --- | --- | --- | +| Background work | fix | 手动刷新 promise、入场写入任务;无新增 polling | 同账户 single-flight,保存/失效为有限单次任务 | 刷新重试/并发测试,shared patch | +| Memory | fix | per-store WeakMap、每 session/account pending、Rust weak lock 表 | 每 session/account 最多 64 个 pending;default coordinator 最多 64 个活跃账户;drain 取消订阅并移除;weak lock 下次访问清理 | pending cap/订阅释放与 weak-lock tests | +| Scope/isolation | fix | session ID、store identity、credential/catalog generation | 旧结果拒绝或忽略;默认按发起选择的 category 写入 | stale refresh、scope switch、pair 回归 | +| Rendering/hot path | keep | 只订阅有未完成保存的 session atom | 无流式事件全表扫描;元数据 index 为单次局部派生 | 源码调用链、rendered hook 测试;无 CPU/RSS 声明 | + +| Provider | Raw transition | App/UI state | Topology/boundary | Expected invariant | Observed evidence | +| --- | --- | --- | --- | --- | --- | +| fixture OAuth/API catalog | metadata 刷新/重复/凭据变化 | 原目录已存在 | KeyService 文件持久化 | 偏好不丢、旧响应不覆盖 | 单测;非 live 供应商验证 | +| fake native provider | 同模型换账户、同时换模型账户 | 已有 runtime | mobile send →真实 session DB/runtime | 下条发送使用新 pair | 应用测试;不经过真实手机 UI | +| runtime admission fixture | pin/失败/重试/并发 | prepared turn | core owner | 拒绝时 DB/runtime 均不变 | owning-boundary 测试 | +| real account / device | 刷新、改选、发送、隐藏/重开 | Tauri / iOS | 实机/双实例 | 视觉、授权、CPU/RSS、实际服务响应 | **not run** | + +Performance verdict: **blocked** for real Tauri/iOS CPU/RSS、真实账户及双实例测量;本轮只有源码资源约束与自动化边界证据,不宣称运行时性能提升。 + +## 验证记录 + +Rust 命令在 `src-tauri/` 执行: + +```sh +cargo test -p key_vault --lib --quiet +cargo test -p agent_core --lib -- state::session_ state::commands::session::identity::tests core::session::project_init::tests core::session::turn::entry::runtime_tests terminal_status_write_preserves_a_newer_model_selection --test-threads=1 +cargo test -p org2 --lib -- agent_sessions::session_directory::patch::tests api::mobile_bridge::adapters::model::tests --test-threads=1 +``` + +- KeyVault:491 passed、1 ignored;core 定向:25 passed;app patch/mobile:15 passed。合计 **531 passed、1 ignored、0 failed**。 +- 编译仅观察到 macOS linker 的 `__eh_frame > 16MB` 警告;未跑全 workspace clippy。 +- `pnpm check:test-placement` 未通过:既有 `src/modules/WorkStation/CodeEditor/Panels/EditorMainPane/hooks/__tests__/useFileContentManager.save.test.ts` 与同目录 colocated tests 混用。该未提交目录在本次修改前已存在,没有改动;本轮新测试遵循所在目录约定。 +- changed-file length 检查通过,所有本轮改动的生产 TS 文件均在 700 行内。 +- 34 个修改的生产 TS/TSX 文件经 AST 检查:0 个原生 button/form 绕过,0 个 clickable div/span 替代按钮;6 个 TSX 的 UI 审计为 0 fix / 6 keep with reason / 0 abstract。 + +前端最终整合命令: + +```sh +pnpm test src/hooks/models src/hooks/keyVault/sharedLocalKeyStore.test.ts src/hooks/keyVault/useLocalKeys.test.ts src/hooks/keyVault/defaultVariantSaveCoordinator.test.ts src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette src/modules/MainApp/Integrations/KeyVault/hooks/refreshAccountModels.test.ts src/modules/MainApp/Integrations/KeyVault/Table/useDefaultVariantSaves.test.ts src/modules/MobileRemote/app/useMobileSessionModel.test.ts src/hooks/session/__tests__/sessionPatchQueue.test.ts src/util/session/__tests__/selectionFromSession.test.ts src/util/__tests__/modelVariants.test.ts src/util/__tests__/selectableModelVariants.test.ts src/util/__tests__/variantEditOptions.test.ts src/engines/ChatPanel/InputArea/components/ModelPill.test.ts src/engines/ChatPanel/InputArea/components/ModelPill.memberOwnership.test.ts src/engines/ChatPanel/InputArea/components/useContextUsageInfo.test.ts src/engines/ChatPanel/InputArea/components/useContextUsageInfo.session.test.ts +pnpm typecheck:fast +pnpm check:circular +git diff --check +``` + +整合测试 **39 文件、211 项通过**。`pnpm typecheck:fast` 通过;`pnpm check:circular` 在 8167 个模块中未发现循环依赖;`git diff --check` 通过。本轮 50 个修改的 TS/TSX 文件使用 `pnpm exec eslint --max-warnings 0 ` 检查通过。 + +真实账户/手机视觉流程未执行;保留同一 DOM、shared controls、布局和样式只能支持代码层面的交互保持,不能替代实机验收。 + +## 普通 KeyVault 保存的并发补强 + +后续检查发现,`credentials.json` 原本通过临时文件改名完成整文件替换,但普通 `save_key` 命令先在锁外读取账户快照,再交给 `KeyService.save_key` 加锁写回。目录刷新或另一次编辑若在两步之间落盘,旧快照可能覆盖与本次请求无关的字段。权威源仍是 KeyService 读取的最新持久化账户;不能依赖前端快照代表当前值。 + +实施计划及结果: + +1. 在 KeyService 的同一把锁内读取最新账户、应用普通 RPC 明确提交的字段、校验并写入。校验失败直接返回,不重写文件;原有完整记录保存仍使用同一验证入口。 +2. 模型启用的两个前端入口只提交 `enabled_models`,不再顺带提交列表加载时的 `available_models`。请求协议、界面与操作步骤不变;显式重新验证或更换 endpoint 的目录提交仍保留。 +3. 在生产写入边界验证两个编辑并发、准备请求后健康/目录更新、失败后文件不变且可重试;在前端请求边界验证启用操作不携带旧目录。 + +架构十层复核:编译、结构、命名、语义、默认值、权威边界、可理解性、wire、入口和目录解析均覆盖本次变动。未进行全工作区死代码清理、真实供应商请求或实机 UI 验收。性能与资源生命周期判定为 **通过源码约束,真实运行未测量**:普通保存仍在 `spawn_blocking` 执行,使用既有进程内互斥锁和单次文件替换,没有新增轮询、定时器、缓存或订阅;锁内校验可能延长单次写入等待,未声称 CPU/RSS 改善。跨进程同时修改同一凭据文件不受这把锁保护,同一字段的并发编辑仍按最后进入写入边界者为准。 + +增量验证:`cargo test -p key_vault --lib --quiet` 为 **494 passed、1 ignored**;`cargo clippy -p key_vault --lib -- -D warnings`、所改 Rust 文件的 `rustfmt --check`、所改前端文件的 ESLint、`pnpm typecheck:fast`、`git diff --check` 均通过。定向前端测试 3 个文件、9 项通过。未重复执行全工作区测试;本次只改变 KeyVault 普通保存边界及两个启用请求。 diff --git a/docs/frontend-ui-audit-2026-09-23/ModelCatalogProjection.md b/docs/frontend-ui-audit-2026-09-23/ModelCatalogProjection.md new file mode 100644 index 0000000000..822dea106e --- /dev/null +++ b/docs/frontend-ui-audit-2026-09-23/ModelCatalogProjection.md @@ -0,0 +1,16 @@ +# Model catalog projection UI audit + +The change updates account/catalog derivation only. Existing two-column navigation, search, source scopes, selection controls, dimensions, and styling are retained. No new UI primitive or action control is introduced. + +| Line | Element | Verdict | Reason | Suggested change | +| --- | --- | --- | --- | --- | +| `src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/VariantPill.tsx:192` | Compound effort/thinking/speed trigger | keep with reason | Existing shared `Button` retains its ref, accessible name, disclosure state, and dropdown lifecycle. Custom layout accommodates the independently present brain, effort, speed, separators, and pencil within the established compact pill geometry. Only variant metadata input changes. | None. | +| `src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/keyFirstItems.tsx:88` | Account/model row labels and family counts | keep with reason | Existing Spotlight row renderer owns interaction and focus; these spans provide label/count content and introduce no independent click target. Existing compact sizes remain unchanged as requested. | None. | +| `src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSelectionItems.tsx:90` | Current/recent model label and source trail | keep with reason | Existing icon, truncation, semantic text colors, and Spotlight action are retained. Catalog metadata replaces model-ID inference without changing the rendered control family. | None. | +| `src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/sourceItems.tsx:172` | Source selection row | keep with reason | Existing Spotlight action and `VariantPill` remain the only interactive controls. The fix makes source eligibility agree with model eligibility, including variant-only catalogs. | None. | +| `src/modules/MainApp/AgentOrgs/config/shared/ModelPicker.tsx:83` | Agent Org model selector | keep with reason | Continues using shared searchable `Select` with the same clear-selection option, size, callback, and disabled state; only compatible account projection is shared. | None. | +| `src/engines/ChatPanel/InputArea/components/ModelPill.tsx:250` | Session model selector | keep with reason | Existing selector pills, popover/dropdown entrypoints and visual markup are retained. Async completion now owns default/recent publication; the menu still dismisses immediately and session identity stays optimistic. | None. | + +Verdict totals: **0 fix**, **6 keep with reason**, **0 abstract**. + +Verification: TypeScript AST inspection of all 34 changed production TS/TSX files (including the six TSX files above) found zero raw button/form-control elements, native button creation calls, or clickable `div`/`span` substitutes. Focused behavior tests cover model/source/key-first catalog parity, independent thinking/effort/speed dimensions, and existing palette flows. Rendered ModelPill tests cover successful save, rejection/retry, direct Member ownership and late completion after navigation. No real Tauri screenshots were captured; visual parity is based on unchanged markup/classes plus rendered tests, not a live-account visual claim. diff --git a/src-tauri/crates/agent-core/src/core/session/gateway_pipeline.rs b/src-tauri/crates/agent-core/src/core/session/gateway_pipeline.rs index 4cc69ccdf6..764a86ebe8 100644 --- a/src-tauri/crates/agent-core/src/core/session/gateway_pipeline.rs +++ b/src-tauri/crates/agent-core/src/core/session/gateway_pipeline.rs @@ -11,7 +11,7 @@ use crate::foundation::persistence::images::load_image_as_data_url; use crate::session::persistence as unified_persistence; use crate::session::IdeContext; -use crate::state::{AgentAppState, AgentSession}; +use crate::state::{AgentAppState, AgentSession, SessionRuntime}; use core_types::key_source::KeySource; /// Process a single inbound gateway message. @@ -23,6 +23,22 @@ pub async fn process_gateway_message( session: Arc, ide_context: Option<&IdeContext>, app_handle: Option, +) -> Result, String> { + let runtime = session + .get_runtime() + .await + .ok_or_else(|| format!("Session {} runtime not initialized", session.id))?; + process_gateway_message_with_runtime(msg, session, runtime, ide_context, app_handle).await +} + +/// Keep the dispatcher-captured runtime through gateway preprocessing and the +/// provider call, even when a newer selection clears the session cache. +pub(crate) async fn process_gateway_message_with_runtime( + msg: InboundMessage, + session: Arc, + runtime: Arc, + ide_context: Option<&IdeContext>, + app_handle: Option, ) -> Result, String> { let preview: String = crate::utils::safe_truncate_chars_to_string(&msg.content, 80); info!( @@ -32,11 +48,6 @@ pub async fn process_gateway_message( let session_key = msg.session_key(); - let runtime = session - .get_runtime() - .await - .ok_or_else(|| format!("Session {} runtime not initialized", session.id))?; - let effective_model = runtime.model.clone(); // Ensure a session record exists in the unified `agent_sessions` table. @@ -47,6 +58,7 @@ pub async fn process_gateway_message( let user_input_preview: String = crate::utils::safe_truncate_chars_to_string(&msg.content, 200); let model = effective_model.clone(); + let account_id = runtime.account_id.clone(); if let Err(err) = tokio::task::spawn_blocking(move || match unified_persistence::get_session(&sk) { Ok(Some(_)) => Ok(()), @@ -58,6 +70,7 @@ pub async fn process_gateway_message( name: format!("Channel: {}", channel), status: super::SessionStatus::Running.as_str().to_string(), model: Some(model), + account_id, session_type: session_type.to_string(), channel: Some(channel), chat_id: Some(chat_id), @@ -151,8 +164,13 @@ pub async fn process_gateway_message( turn_intent_id: uuid::Uuid::new_v4().to_string(), }; - let result = - crate::session::process_message(Arc::clone(&session), input, app_handle.clone()).await; + let result = super::turn::entry::process_message_with_runtime( + Arc::clone(&session), + runtime, + input, + app_handle.clone(), + ) + .await; // Compact-fork redirect. if let Ok(ref pr) = result { diff --git a/src-tauri/crates/agent-core/src/core/session/launch/launch_helpers.rs b/src-tauri/crates/agent-core/src/core/session/launch/launch_helpers.rs index 8fb03ceafe..2fea8f8353 100644 --- a/src-tauri/crates/agent-core/src/core/session/launch/launch_helpers.rs +++ b/src-tauri/crates/agent-core/src/core/session/launch/launch_helpers.rs @@ -67,7 +67,7 @@ pub(super) async fn handle_background_launch_failure( } match mark_session_failed(session_id.to_string()).await { - Ok(()) => crate::lifecycle::emit_session_status_changed( + Ok(_terminal_at) => crate::lifecycle::emit_session_status_changed( app_handle, session_id, crate::persistence::db_helpers::AgentSessionStatus::Failed, @@ -249,18 +249,21 @@ pub(super) fn broadcast_launch_send_error(session_id: &str, message: &str) { broadcast_agent_error_structured(session_id, &error); } -pub(super) async fn mark_session_failed(session_id: String) -> Result<(), String> { +pub(super) async fn mark_session_failed(session_id: String) -> Result { tokio::task::spawn_blocking(move || { - let Some(mut record) = - crate::session::persistence::get_session(&session_id).map_err(|err| err.to_string())? - else { + let terminal_at = chrono::Utc::now().to_rfc3339(); + let changed = crate::session::persistence::update_status_at( + &session_id, + crate::session::SessionStatus::Failed, + &terminal_at, + ) + .map_err(|err| err.to_string())?; + if !changed { return Err(format!( "session {session_id} disappeared before first-turn failure could be persisted" )); - }; - record.status = crate::session::SessionStatus::Failed.as_str().to_string(); - record.updated_at = chrono::Utc::now().to_rfc3339(); - crate::session::persistence::upsert_session(&record).map_err(|err| err.to_string()) + } + Ok(terminal_at) }) .await .map_err(|err| err.to_string())? diff --git a/src-tauri/crates/agent-core/src/core/session/launch/launch_org.rs b/src-tauri/crates/agent-core/src/core/session/launch/launch_org.rs index e148504d7d..074a3e70b5 100644 --- a/src-tauri/crates/agent-core/src/core/session/launch/launch_org.rs +++ b/src-tauri/crates/agent-core/src/core/session/launch/launch_org.rs @@ -327,8 +327,11 @@ pub(super) async fn send_initial_turn( content, None, crate::state::commands::session::identity::IdentityOverrides { - model, - account_id, + // The launch row already owns the chosen pair. Send resolves + // it under its own identity lock so a newer picker change + // cannot be replaced by this task's captured launch defaults. + model: None, + account_id: None, workspace_root: Some(workspace_root), native_harness_type, }, @@ -350,8 +353,20 @@ pub(super) async fn send_initial_turn( return Ok(()); } + let identity_guard = crate::state::session_identity_lock(session_id) + .await + .lock_owned() + .await; + let (model, account_id) = + crate::state::commands::session::identity::resolve_initialization_model_pair( + state, + session_id, + model.as_deref(), + account_id.as_deref(), + ) + .await?; let model = model.ok_or_else(|| "model is required for sub-agent launch".to_string())?; - let launch_spec = crate::init::launch_spec::AgentLaunchSpec::work_item_session( + let mut launch_spec = crate::init::launch_spec::AgentLaunchSpec::work_item_session( state, session_id, &model, @@ -361,7 +376,12 @@ pub(super) async fn send_initial_turn( &sub_agent_ids, ) .await?; + // Preserve credential-owned None rather than the constructor's empty + // account placeholder, which would conflict with a Market selector. + launch_spec.account_id = account_id; crate::init::init_session(state, launch_spec).await?; + // `send_message_impl` acquires this same non-reentrant lock itself. + drop(identity_guard); crate::state::commands::session::message::send_message_impl( state, @@ -369,8 +389,8 @@ pub(super) async fn send_initial_turn( content, None, crate::state::commands::session::identity::IdentityOverrides { - model: Some(model), - account_id, + model: None, + account_id: None, workspace_root: Some(workspace_root), native_harness_type, }, diff --git a/src-tauri/crates/agent-core/src/core/session/launch/mod.rs b/src-tauri/crates/agent-core/src/core/session/launch/mod.rs index 1eba9f7c7f..06166439c2 100644 --- a/src-tauri/crates/agent-core/src/core/session/launch/mod.rs +++ b/src-tauri/crates/agent-core/src/core/session/launch/mod.rs @@ -255,6 +255,28 @@ async fn generate_title_before_first_turn( return; } + // This task is spawned independently of the initial turn. Its captured + // launch defaults may be older than an already acknowledged picker edit. + let identity_guard = crate::state::session_identity_lock(session_id) + .await + .lock_owned() + .await; + let (model, account_id) = + match crate::state::commands::session::identity::resolve_initialization_model_pair( + state, + session_id, + model.as_deref(), + account_id.as_deref(), + ) + .await + { + Ok(pair) => pair, + Err(error) => { + tracing::warn!(session_id = %session_id, %error, "[session_title] failed to load current model selection"); + return; + } + }; + let launch_spec = match AgentLaunchSpec::from_session_sources( state, session_id, @@ -287,6 +309,9 @@ async fn generate_title_before_first_turn( return; } }; + // Only runtime installation is serialized; title generation is an + // independent provider request and must not block a next-turn selection. + drop(identity_guard); let title = crate::session::title::generate_and_persist_session_title( session_id, diff --git a/src-tauri/crates/agent-core/src/core/session/persistence/crud/mod.rs b/src-tauri/crates/agent-core/src/core/session/persistence/crud/mod.rs index 9019cbfda6..6f08f34218 100644 --- a/src-tauri/crates/agent-core/src/core/session/persistence/crud/mod.rs +++ b/src-tauri/crates/agent-core/src/core/session/persistence/crud/mod.rs @@ -33,7 +33,7 @@ pub use ops::{ upsert_session, }; pub(crate) use ops::{ - delete_session_with_connection, finish_session_delete, prepare_session_delete, + delete_session_with_connection, finish_session_delete, prepare_session_delete, update_status_at, }; pub(super) use record::{row_to_record, UNIFIED_SESSION_SELECT}; pub use record::{session_type, UnifiedSessionRecord}; diff --git a/src-tauri/crates/agent-core/src/core/session/persistence/crud/ops.rs b/src-tauri/crates/agent-core/src/core/session/persistence/crud/ops.rs index 85bc8d7602..223410b849 100644 --- a/src-tauri/crates/agent-core/src/core/session/persistence/crud/ops.rs +++ b/src-tauri/crates/agent-core/src/core/session/persistence/crud/ops.rs @@ -390,12 +390,21 @@ fn settles_linked_session(status: SessionStatus) -> bool { /// Update session status. pub fn update_status(session_id: &str, status: SessionStatus) -> SqliteResult { + update_status_at(session_id, status, &Utc::now().to_rfc3339()) +} + +/// Status-only write with the caller's event timestamp. This must never carry +/// a previously loaded model/account pair back into the authoritative row. +pub(crate) fn update_status_at( + session_id: &str, + status: SessionStatus, + updated_at: &str, +) -> SqliteResult { let changed = with_sessions_writer(|| -> SqliteResult { let conn = get_connection()?; - let now = Utc::now().to_rfc3339(); let updated = conn.execute( "UPDATE agent_sessions SET status = ?2, updated_at = ?3 WHERE session_id = ?1", - params![session_id, status.as_str(), now], + params![session_id, status.as_str(), updated_at], )?; Ok(updated > 0) })?; diff --git a/src-tauri/crates/agent-core/src/core/session/persistence/crud/ops_tests.rs b/src-tauri/crates/agent-core/src/core/session/persistence/crud/ops_tests.rs index 8da6c238ba..da6bada315 100644 --- a/src-tauri/crates/agent-core/src/core/session/persistence/crud/ops_tests.rs +++ b/src-tauri/crates/agent-core/src/core/session/persistence/crud/ops_tests.rs @@ -108,6 +108,22 @@ fn seed_session(session_id: &str, status: SessionStatus) { upsert_session(&record).expect("seed upsert"); } +#[test] +fn terminal_status_write_preserves_a_newer_model_selection() { + let _sandbox = test_env::sandbox(); + let sid = "status-after-identity-change"; + seed_session(sid, SessionStatus::Running); + super::ops::update_model_and_account(sid, "new-model", Some("new-account")).unwrap(); + let terminal_at = "2026-09-23T10:00:00Z"; + assert!(super::ops::update_status_at(sid, SessionStatus::Failed, terminal_at).unwrap()); + let record = super::ops::get_session(sid).unwrap().unwrap(); + assert_eq!(record.model.as_deref(), Some("new-model")); + assert_eq!(record.account_id.as_deref(), Some("new-account")); + assert_eq!(record.status, SessionStatus::Failed.as_str()); + assert_eq!(record.updated_at, terminal_at); + assert!(!super::ops::update_status_at("missing", SessionStatus::Failed, terminal_at).unwrap()); +} + #[test] #[serial_test::serial] fn delete_session_refuses_active_shell_before_removing_session_row() { diff --git a/src-tauri/crates/agent-core/src/core/session/persistence/mod.rs b/src-tauri/crates/agent-core/src/core/session/persistence/mod.rs index 96f8c47d31..0292f22b0c 100644 --- a/src-tauri/crates/agent-core/src/core/session/persistence/mod.rs +++ b/src-tauri/crates/agent-core/src/core/session/persistence/mod.rs @@ -35,7 +35,7 @@ pub use crud::{ update_worktree_merge_status, upsert_session, UnifiedSessionRecord, }; pub(crate) use crud::{ - delete_session_with_connection, finish_session_delete, prepare_session_delete, + delete_session_with_connection, finish_session_delete, prepare_session_delete, update_status_at, }; pub use sidebar::{ list_agent_org_root_sessions_page, list_standalone_coding_sessions_page, diff --git a/src-tauri/crates/agent-core/src/core/session/project_init.rs b/src-tauri/crates/agent-core/src/core/session/project_init.rs index abac6b8cf6..e76bf6963c 100644 --- a/src-tauri/crates/agent-core/src/core/session/project_init.rs +++ b/src-tauri/crates/agent-core/src/core/session/project_init.rs @@ -7,6 +7,17 @@ use crate::session::persistence as unified_persistence; use crate::state::{AgentAppState, SessionRuntime}; use core_types::key_source::KeySource; +fn seed_missing_workspace_identity( + record: &mut unified_persistence::UnifiedSessionRecord, + model: String, + account_id: Option, +) { + if record.model.as_deref().is_none_or(str::is_empty) { + record.model = Some(model); + record.account_id = account_id; + } +} + /// Initialize a workspace-scoped session's runtime using the unified init path. /// /// Agent resolve contract (design doc §11.4): coding sessions resolve against the session's @@ -21,6 +32,10 @@ use core_types::key_source::KeySource; /// `message_pipeline` fallback branch (which can only synthesize a /// generic OS-typed row with an empty `workspace_path`) ever sees this /// session id. +/// +/// The production channel dispatcher owns `session_identity_lock` across +/// resolving the current model/account, this initialization, and the eager +/// persistence below. Do not reacquire that non-reentrant lock here. pub async fn init_workspace_session( state: &AgentAppState, session_id: &str, @@ -59,10 +74,9 @@ pub async fn init_workspace_session( // don't overwrite channel / chat_id / parent metadata that a // previous dispatch established. Some(mut existing) => { - existing.model = Some(model_owned); - if existing.account_id.is_none() { - existing.account_id = account_owned; - } + // The caller resolved the chosen pair while holding the + // identity lock. Preserve an existing persisted selection. + seed_missing_workspace_identity(&mut existing, model_owned, account_owned); let needs_workspace = existing .workspace_path .as_deref() @@ -133,3 +147,59 @@ pub async fn init_workspace_session( Ok(runtime) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn eager_workspace_write_preserves_selected_model_and_account() { + let mut record = unified_persistence::UnifiedSessionRecord { + model: Some("selected".into()), + account_id: Some("selected-account".into()), + ..Default::default() + }; + seed_missing_workspace_identity( + &mut record, + "stale-default".into(), + Some("stale-account".into()), + ); + assert_eq!(record.model.as_deref(), Some("selected")); + assert_eq!(record.account_id.as_deref(), Some("selected-account")); + } + + #[test] + fn eager_workspace_write_does_not_fill_a_credential_owned_account() { + let mut record = unified_persistence::UnifiedSessionRecord { + model: Some("market-model".into()), + credential_source: Some("market:selection".into()), + ..Default::default() + }; + seed_missing_workspace_identity( + &mut record, + "gateway-model".into(), + Some("personal-account".into()), + ); + assert_eq!(record.model.as_deref(), Some("market-model")); + assert_eq!(record.account_id, None); + assert_eq!( + record.credential_source.as_deref(), + Some("market:selection") + ); + } + + #[test] + fn eager_workspace_write_seeds_identity_as_a_complete_pair() { + let mut record = unified_persistence::UnifiedSessionRecord { + account_id: Some("identity-less-old-account".into()), + ..Default::default() + }; + seed_missing_workspace_identity( + &mut record, + "gateway-model".into(), + Some("gateway-account".into()), + ); + assert_eq!(record.model.as_deref(), Some("gateway-model")); + assert_eq!(record.account_id.as_deref(), Some("gateway-account")); + } +} diff --git a/src-tauri/crates/agent-core/src/core/session/turn/entry.rs b/src-tauri/crates/agent-core/src/core/session/turn/entry.rs index d9e57e22c3..e2e2272ff4 100644 --- a/src-tauri/crates/agent-core/src/core/session/turn/entry.rs +++ b/src-tauri/crates/agent-core/src/core/session/turn/entry.rs @@ -7,7 +7,7 @@ use std::sync::Arc; use tauri::Manager; -use crate::state::AgentSession; +use crate::state::{AgentSession, SessionRuntime}; use super::super::types::{ProcessingContext, ProcessingResult}; use super::event_handler::EventHandlerConfig; @@ -129,6 +129,19 @@ pub async fn process_message( .await .ok_or_else(|| format!("Session {} runtime not initialized", session.id))?; + process_message_with_runtime(session, runtime, input, app_handle).await +} + +/// Execute with the runtime captured while the caller prepared this turn. +/// A later model selection may invalidate or replace the session cache, but +/// applies to the next turn rather than changing this admitted provider call. +/// Callers retain responsibility for their existing admission checks. +pub(crate) async fn process_message_with_runtime( + session: Arc, + runtime: Arc, + input: TurnInput, + app_handle: Option, +) -> Result { let workspace_path = runtime.workspace_state.read().working_dir().to_path_buf(); let lsp_manager = extract_lsp_manager(&app_handle); @@ -243,6 +256,10 @@ pub async fn process_message( .await } +#[cfg(test)] +#[path = "entry_runtime_tests.rs"] +mod runtime_tests; + #[cfg(test)] mod tests { use super::expand_skill_slash_command; diff --git a/src-tauri/crates/agent-core/src/core/session/turn/entry_runtime_tests.rs b/src-tauri/crates/agent-core/src/core/session/turn/entry_runtime_tests.rs new file mode 100644 index 0000000000..3109e21043 --- /dev/null +++ b/src-tauri/crates/agent-core/src/core/session/turn/entry_runtime_tests.rs @@ -0,0 +1,179 @@ +//! Provider execution must retain the runtime admitted before a model edit. + +use std::sync::Arc; + +use async_trait::async_trait; +use parking_lot::Mutex; +use serde_json::Value; + +use super::{process_message, process_message_with_runtime, TurnInput}; +use crate::definitions::{resolved::ResolvedAgent, AgentDefinition}; +use crate::providers::traits::{LLMProvider, LLMResponse, ProviderError}; +use crate::session::persistence::{self, UnifiedSessionRecord}; +use crate::state::{AgentSession, SessionRuntime}; + +struct RecordingProvider { + account: &'static str, + calls: Arc>>, +} + +#[async_trait] +impl LLMProvider for RecordingProvider { + async fn chat( + &self, + _: &[Value], + _: Option<&[Value]>, + model: &str, + _: u32, + _: f32, + ) -> Result { + self.calls + .lock() + .push((self.account.to_string(), model.to_string())); + Ok(LLMResponse::text(&format!("{}:{model}", self.account))) + } + + fn default_model(&self) -> &str { + "model-a" + } + + fn provider_name(&self) -> &str { + "captured-runtime-fixture" + } +} + +fn runtime( + model: &str, + account: &'static str, + workspace: &std::path::Path, + calls: Arc>>, +) -> Arc { + let definition = AgentDefinition { + selected_model_id: Some(model.to_string()), + ..Default::default() + }; + let resolved = ResolvedAgent::resolve(&definition, None, &Default::default()).unwrap(); + Arc::new(SessionRuntime { + provider: Arc::new(RecordingProvider { account, calls }), + tool_registry: Arc::new(crate::tools::registry::ToolRegistry::new()), + policy: Arc::new(crate::tools::policy::ResolvedToolPolicy::permissive()), + model: model.to_string(), + account_id: Some(account.to_string()), + native_harness_type: None, + workspace_state: Arc::new(parking_lot::RwLock::new( + crate::session::workspace::SessionWorkspace::new(workspace.to_path_buf()), + )), + mcp_auto_approved: vec![], + resolved, + integrations_snapshot: Default::default(), + overrides: Default::default(), + agent_soul: None, + sovereign_prompt: false, + policy_context_activator: None, + agent_org_context: None, + agent_org_current_member_id: None, + agent_definition_id: None, + }) +} + +fn init_schema() { + let connection = database::db::get_connection().unwrap(); + crate::persistence::test_schema::ensure_agent_sessions_schema(&connection); + crate::persistence::session_snapshots::ensure_tables_with(&connection).unwrap(); + persistence::init(&connection).unwrap(); + crate::coordination::init_agent_org_schemas(&connection).unwrap(); +} + +fn input(content: &str) -> TurnInput { + TurnInput { + content: content.to_string(), + turn_intent_id: uuid::Uuid::new_v4().to_string(), + ..Default::default() + } +} + +#[tokio::test(flavor = "multi_thread")] +async fn admitted_turn_executes_captured_provider_after_cache_invalidation_or_replacement() { + let sandbox = test_helpers::test_env::sandbox(); + init_schema(); + for replace_cache in [false, true] { + let sid = format!("captured-runtime-{replace_cache}"); + let session = Arc::new(AgentSession::new(sid.clone(), AgentDefinition::default())); + persistence::upsert_session(&UnifiedSessionRecord { + session_id: sid.clone(), + model: Some("model-a".to_string()), + account_id: Some("account-a".to_string()), + ..Default::default() + }) + .unwrap(); + let calls = Arc::new(Mutex::new(Vec::new())); + let admitted = runtime("model-a", "account-a", sandbox.path(), Arc::clone(&calls)); + let next = runtime("model-b", "account-b", sandbox.path(), Arc::clone(&calls)); + session.set_runtime(Arc::clone(&admitted)).await.unwrap(); + + // Reproduce the model edit after preparation releases its identity lock. + let mutation = session.begin_identity_mutation().await.unwrap(); + persistence::update_model_and_account(&sid, "model-b", Some("account-b")).unwrap(); + mutation.invalidate_runtime().await; + if replace_cache { + session.set_runtime(Arc::clone(&next)).await.unwrap(); + } + + let result = process_message_with_runtime( + Arc::clone(&session), + admitted, + input("First admitted request"), + None, + ) + .await + .expect("a later model edit cannot uninitialize the admitted request"); + assert_eq!(result.content, "account-a:model-a"); + assert_eq!(*calls.lock(), vec![("account-a".into(), "model-a".into())]); + let saved = persistence::get_session(&sid).unwrap().unwrap(); + assert_eq!(saved.model.as_deref(), Some("model-b")); + assert_eq!(saved.account_id.as_deref(), Some("account-b")); + + if !replace_cache { + session.set_runtime(next).await.unwrap(); + } + let result = process_message(session, input("Next request"), None) + .await + .unwrap(); + assert_eq!(result.content, "account-b:model-b"); + assert_eq!( + calls.lock().last().unwrap(), + &("account-b".into(), "model-b".into()) + ); + } +} + +#[tokio::test(flavor = "multi_thread")] +async fn gateway_preprocessing_preserves_captured_provider_and_persists_complete_seed() { + let sandbox = test_helpers::test_env::sandbox(); + init_schema(); + let sid = "captured-gateway-runtime"; + let session = Arc::new(AgentSession::new(sid.into(), AgentDefinition::default())); + let calls = Arc::new(Mutex::new(Vec::new())); + let admitted = runtime("model-a", "account-a", sandbox.path(), Arc::clone(&calls)); + session.set_runtime(Arc::clone(&admitted)).await.unwrap(); + session + .begin_identity_mutation() + .await + .unwrap() + .invalidate_runtime() + .await; + let mut message = crate::bus::InboundMessage::new("fixture", "user", "chat", "Request"); + message.session_key_override = Some(sid.into()); + + let response = crate::session::gateway_pipeline::process_gateway_message_with_runtime( + message, session, admitted, None, None, + ) + .await + .unwrap() + .unwrap(); + assert_eq!(response.content, "account-a:model-a"); + assert_eq!(*calls.lock(), vec![("account-a".into(), "model-a".into())]); + let saved = persistence::get_session(sid).unwrap().unwrap(); + assert_eq!(saved.model.as_deref(), Some("model-a")); + assert_eq!(saved.account_id.as_deref(), Some("account-a")); +} diff --git a/src-tauri/crates/agent-core/src/state/commands/channel_handler/dispatch.rs b/src-tauri/crates/agent-core/src/state/commands/channel_handler/dispatch.rs index 69029e03bb..5dd8b1c75e 100644 --- a/src-tauri/crates/agent-core/src/state/commands/channel_handler/dispatch.rs +++ b/src-tauri/crates/agent-core/src/state/commands/channel_handler/dispatch.rs @@ -289,8 +289,22 @@ async fn dispatch_to_session( _question_manager: &Arc, _permission_manager: &Arc, ) -> Result, String> { + let identity_guard = crate::state::session_identity_lock(target_session_id) + .await + .lock_owned() + .await; let (gw_account, gw_model) = resolve_gateway_model_and_account(state).await; - let effective_account = account_id.or(gw_account.as_deref()); + // Gateway configuration/global account remain new-session seeds. An + // existing per-chat selection always supplies the whole model/account pair. + let (session_model, session_account) = + crate::state::commands::session::identity::resolve_initialization_model_pair( + state, + target_session_id, + gw_model.as_deref(), + account_id.or(gw_account.as_deref()), + ) + .await?; + let effective_account = session_account.as_deref(); // SDE sessions have a session-specific `workspace_path` that MUST NOT // be overwritten by the generic `channels.workspace_path()`. Route @@ -310,7 +324,7 @@ async fn dispatch_to_session( match persisted_workspace { Some(path_str) if !path_str.is_empty() => { let workspace_path = std::path::PathBuf::from(&path_str); - let model = gw_model + let model = session_model .as_deref() .ok_or_else(|| "gateway.model not configured".to_string())?; crate::session::init_workspace_session( @@ -330,7 +344,7 @@ async fn dispatch_to_session( state, target_session_id, effective_account, - gw_model.as_deref(), + session_model.as_deref(), ) .await? } @@ -340,11 +354,20 @@ async fn dispatch_to_session( state, target_session_id, effective_account, - gw_model.as_deref(), + session_model.as_deref(), ) .await? }; + let session_arc = state + .get_session(target_session_id) + .await + .ok_or_else(|| format!("Session {} not found after init", target_session_id))?; + + // Both runtime installation and the SDE helper's eager identity write + // are complete. Do not hold this lock across gateway/LLM processing. + drop(identity_guard); + let (origin_channel, origin_chat_id) = if is_reinject { let src_channel = msg .metadata @@ -394,14 +417,10 @@ async fn dispatch_to_session( m }; - let session_arc = state - .get_session(target_session_id) - .await - .ok_or_else(|| format!("Session {} not found after init", target_session_id))?; - - let outbound = crate::session::gateway_pipeline::process_gateway_message( + let outbound = crate::session::gateway_pipeline::process_gateway_message_with_runtime( enriched, session_arc, + runtime, None, state.app_handle.clone(), ) diff --git a/src-tauri/crates/agent-core/src/state/commands/channel_handler/lifecycle.rs b/src-tauri/crates/agent-core/src/state/commands/channel_handler/lifecycle.rs index d65a4e5c6c..39041241ad 100644 --- a/src-tauri/crates/agent-core/src/state/commands/channel_handler/lifecycle.rs +++ b/src-tauri/crates/agent-core/src/state/commands/channel_handler/lifecycle.rs @@ -75,6 +75,8 @@ pub async fn agent_toggle_channel( /// Channel session init: resolve definition from the registered session, /// use `personal_workspace()` as the working directory, delegate to `init_session`. +/// The dispatcher owns the identity lock and passes its resolved session pair; +/// these parameters are not independent gateway overrides for an existing row. pub(super) async fn init_channel_session( state: &AgentAppState, session_id: &str, diff --git a/src-tauri/crates/agent-core/src/state/commands/session/channel.rs b/src-tauri/crates/agent-core/src/state/commands/session/channel.rs index dcb23ec1aa..09abf58f99 100644 --- a/src-tauri/crates/agent-core/src/state/commands/session/channel.rs +++ b/src-tauri/crates/agent-core/src/state/commands/session/channel.rs @@ -26,77 +26,47 @@ pub async fn channel_process_message( ); let session_key = session_id.unwrap_or_else(|| "tauri:direct".to_string()); info!("[channel_process_message] session_key={}", session_key); + let identity_guard = crate::state::session_identity_lock(&session_key) + .await + .lock_owned() + .await; - // Model override is threaded per-request into `init_session`; we do NOT - // mutate any shared state. Account changes do still require invalidation - // so the next turn re-picks the provider. The comparison baseline is - // THIS session's runtime account — never a global (an unrelated - // session's switch must not affect us, and ours must not affect them). - { - let runtime_account = match state.get_session(&session_key).await { - Some(session) => session - .get_runtime() - .await - .and_then(|r| r.account_id.clone()), + // Explicit command edits still win. Otherwise keep the complete current + // session pair instead of reconstructing a model from agent defaults. + let (current_model, current_account) = + super::identity::resolve_initialization_model_pair(&state, &session_key, None, None) + .await?; + let previous_account = account_id + .as_ref() + .filter(|account| current_account.as_deref() != Some(account.as_str())) + .map(|_| current_account.clone()); + if model.is_some() || account_id.is_some() { + let session = state.get_session(&session_key).await; + // A prepared Member turn pins its provider. Reject before writing and + // keep the admission guard through the atomic identity edit. + let mutation = match session.as_ref() { + Some(session) => Some(session.begin_identity_mutation().await?), None => None, }; - let account_changed = if let Some(ref new_account_id) = account_id { - runtime_account.as_deref() != Some(new_account_id.as_str()) - } else { - false - }; - - if account_changed { - state.invalidate_session(&session_key).await; - if let Some(ref new_account_id) = account_id { - session_persistence::update_account_id(&session_key, new_account_id).map_err( - |err| format!("[channel] Failed to persist account switch: {}", err), - )?; - crate::lifecycle::emit_session_account_switched( - state.app_handle.as_ref(), - &session_key, - runtime_account.as_deref(), - new_account_id, - model.as_deref(), - ); + match (model.as_deref(), account_id.as_deref()) { + (Some(model), account) => { + session_persistence::update_model_and_account(&session_key, model, account) + .map_err(|err| format!("[channel] Failed to persist model switch: {err}"))?; } + (None, Some(account)) => { + session_persistence::update_account_id(&session_key, account) + .map_err(|err| format!("[channel] Failed to persist account switch: {err}"))?; + } + (None, None) => unreachable!("identity edit requires an override"), } - if let Some(ref new_model) = model { - session_persistence::update_model(&session_key, new_model) - .map_err(|err| format!("[channel] Failed to persist model switch: {}", err))?; + if let Some(mutation) = mutation { + // Release the admission guard before init, which validates the + // same domain while installing the replacement runtime. + mutation.invalidate_runtime().await; } } - - // Effective model: the caller-supplied override takes precedence; - // otherwise the session's agent definition resolves it at - // `init_session` time. - let requested_model_override = model.clone(); - - // Account chain mirrors `resolve_session_identity`: override → this - // session's runtime → DB row. No global fallback. - let effective_account_id = if account_id.is_some() { - account_id - } else { - let runtime_account = match state.get_session(&session_key).await { - Some(session) => session - .get_runtime() - .await - .and_then(|r| r.account_id.clone()), - None => None, - }; - if runtime_account.is_some() { - runtime_account - } else { - let sk = session_key.clone(); - tokio::task::spawn_blocking(move || { - session_persistence::get_session(&sk) - .map_err(|err| format!("[channel] DB error loading account_id: {}", err)) - .map(|opt| opt.and_then(|s| s.account_id)) - }) - .await - .map_err(|err| format!("[channel] Task panic loading account_id: {}", err))?? - } - }; + let requested_model_override = model.or(current_model); + let effective_account_id = account_id.or(current_account); let ide_context = if active_repo_path.is_some() || active_branch.is_some() { Some(crate::session::IdeContext { @@ -121,6 +91,19 @@ pub async fn channel_process_message( .await?; let runtime = crate::init::init_session(&state, launch_spec).await?; let effective_model = runtime.model.clone(); + if let (Some(previous_account), Some(account)) = + (previous_account, runtime.account_id.as_deref()) + { + // Account-only edits still publish the resolved model. Consumers must + // receive this complete pair instead of merging it with optimistic UI. + crate::lifecycle::emit_session_account_switched( + state.app_handle.as_ref(), + &session_key, + previous_account.as_deref(), + account, + Some(&effective_model), + ); + } let session_arc = state .get_session(&session_key) @@ -141,10 +124,19 @@ pub async fn channel_process_message( turn_intent_id: uuid::Uuid::new_v4().to_string(), }; + // The turn captures its provider separately. Never serialize the whole + // provider request behind the model-picker identity lock. + drop(identity_guard); + const CALLER_TIMEOUT_SECS: u64 = 180; let response = tokio::time::timeout( std::time::Duration::from_secs(CALLER_TIMEOUT_SECS), - crate::session::process_message(session_arc, input, state.app_handle.clone()), + crate::session::turn::entry::process_message_with_runtime( + session_arc, + runtime, + input, + state.app_handle.clone(), + ), ) .await .map_err(|_| format!("Request timed out after {}s", CALLER_TIMEOUT_SECS))? diff --git a/src-tauri/crates/agent-core/src/state/commands/session/compaction.rs b/src-tauri/crates/agent-core/src/state/commands/session/compaction.rs index a539aa8caf..4d96453b57 100644 --- a/src-tauri/crates/agent-core/src/state/commands/session/compaction.rs +++ b/src-tauri/crates/agent-core/src/state/commands/session/compaction.rs @@ -97,6 +97,13 @@ pub async fn prepare_session_for_scheduler_maintenance( if session_id.starts_with(core_types::session::CLI_SESSION_PREFIX) { return Err("Native CLI compaction is owned by the provider runtime".to_string()); } + // Cold maintenance initialization reads the same identity as a send. Keep + // an overlapping picker edit from clearing the cache before this helper + // installs a runtime assembled with the previous model/account pair. + let _identity_guard = crate::state::session_identity_lock(session_id) + .await + .lock_owned() + .await; let needs_init = match state.get_session(session_id).await { Some(session) => session.get_runtime().await.is_none(), None => true, diff --git a/src-tauri/crates/agent-core/src/state/commands/session/identity.rs b/src-tauri/crates/agent-core/src/state/commands/session/identity.rs index 844b5db7d0..a2246f706e 100644 --- a/src-tauri/crates/agent-core/src/state/commands/session/identity.rs +++ b/src-tauri/crates/agent-core/src/state/commands/session/identity.rs @@ -63,6 +63,51 @@ pub struct IdentityOverrides { pub workspace_root: Option, } +/// Cold-init defaults are seeds, not overrides of a session's chosen pair. +/// The caller holds `session_identity_lock` through runtime installation and +/// any identity writeback. A missing account on an existing pair is intentional +/// (for example a Market credential) and must not inherit a default account. +pub(crate) async fn resolve_initialization_model_pair( + state: &AgentAppState, + session_id: &str, + default_model: Option<&str>, + default_account_id: Option<&str>, +) -> Result<(Option, Option), String> { + if let Some(session) = state.get_session(session_id).await { + if let Some(runtime) = session.get_runtime().await { + if !runtime.model.is_empty() { + return Ok((Some(runtime.model.clone()), runtime.account_id.clone())); + } + } + } + let sid = session_id.to_string(); + let persisted = tokio::task::spawn_blocking(move || session_persistence::get_session(&sid)) + .await + .map_err(|error| format!("Initialization identity lookup failed: {error}"))? + .map_err(|error| format!("Initialization identity DB lookup failed: {error}"))?; + Ok(initialization_model_pair( + persisted.map(|record| (record.model, record.account_id)), + default_model, + default_account_id, + )) +} + +fn initialization_model_pair( + stored: Option<(Option, Option)>, + default_model: Option<&str>, + default_account_id: Option<&str>, +) -> (Option, Option) { + if let Some((Some(model), account)) = stored { + if !model.is_empty() { + return (Some(model), account); + } + } + ( + default_model.map(str::to_string), + default_account_id.map(str::to_string), + ) +} + /// Resolve session identity with a strict priority chain (same for all /// three fields): /// 1. Caller-supplied overrides @@ -264,6 +309,47 @@ fn workspace_paths_to_working_dir( mod tests { use super::*; + #[test] + fn initialization_preserves_selected_pair_over_stale_launch_defaults() { + assert_eq!( + initialization_model_pair( + Some(( + Some("selected-model".into()), + Some("selected-account".into()) + )), + Some("old-launch-model"), + Some("old-launch-account"), + ), + ( + Some("selected-model".into()), + Some("selected-account".into()) + ) + ); + } + + #[test] + fn initialization_preserves_credential_owned_pair_without_an_account() { + assert_eq!( + initialization_model_pair( + Some((Some("market-model".into()), None)), + Some("gateway-model"), + Some("personal-account"), + ), + (Some("market-model".into()), None) + ); + } + + #[test] + fn initialization_keeps_new_channel_default_behavior_without_a_stored_model() { + for stored in [None, Some((None, Some("unused-account".into())))] { + assert_eq!( + initialization_model_pair(stored, Some("gateway-model"), Some("gateway-account")), + (Some("gateway-model".into()), Some("gateway-account".into())) + ); + } + assert_eq!(initialization_model_pair(None, None, None), (None, None)); + } + #[test] fn worktree_session_resolves_to_worktree_path() { // Worktree session: workspace_path is the user's project, diff --git a/src-tauri/crates/agent-core/src/state/commands/session/message/send.rs b/src-tauri/crates/agent-core/src/state/commands/session/message/send.rs index fac14a49c8..92a3790d7a 100644 --- a/src-tauri/crates/agent-core/src/state/commands/session/message/send.rs +++ b/src-tauri/crates/agent-core/src/state/commands/session/message/send.rs @@ -311,6 +311,13 @@ pub(crate) async fn send_message_impl( ); // ── 1. Resolve session identity (unified — single code path) ───────── + // Finish admission's runtime/DB writeback before a picker can commit the + // next-turn identity. Release before execution: an active reply keeps its + // captured provider while the picker can configure the following turn. + let identity_guard = crate::state::session_identity_lock(&session_id) + .await + .lock_owned() + .await; let identity = resolve_session_identity(state, &session_id, overrides).await?; let explicit_org_run_id = match (org_wake_run_id.as_deref(), intent_org_run_id.as_deref()) { @@ -521,6 +528,7 @@ pub(crate) async fn send_message_impl( .get_session(&session_id) .await .ok_or_else(|| format!("Session not found after init: {}", session_id))?; + let admitted_runtime_lease_id = session_handle.runtime_lease_for(&runtime).await?; session_handle.refresh_last_active().await; @@ -829,6 +837,8 @@ pub(crate) async fn send_message_impl( } } + drop(identity_guard); + // ── 4b. Project root WorkItem bootstrap (orgtrack/v1 §7.2) ────────── // // The first accepted non-empty submission of a Project session with @@ -1109,7 +1119,11 @@ pub(crate) async fn send_message_impl( } let turn_id = session - .begin_turn_with_intent(content.clone(), Some(turn_intent_id.clone())) + .begin_turn_with_runtime_lease( + content.clone(), + Some(turn_intent_id.clone()), + Some(admitted_runtime_lease_id), + ) .await; if let Some(reservation) = direct_runtime_admission.as_ref() { session.release_runtime_admission(reservation).await; @@ -1138,9 +1152,13 @@ pub(crate) async fn send_message_impl( turn_intent_id: turn_intent_id.clone(), }; - let response = - crate::session::process_message(Arc::clone(&session), input, app_handle.clone()) - .await; + let response = crate::session::turn::entry::process_message_with_runtime( + Arc::clone(&session), + Arc::clone(&runtime), + input, + app_handle.clone(), + ) + .await; let final_turn_state = if cancel_flag.load(std::sync::atomic::Ordering::SeqCst) { crate::session::DialogTurnState::Cancelled diff --git a/src-tauri/crates/agent-core/src/state/mod.rs b/src-tauri/crates/agent-core/src/state/mod.rs index 1343fe66a3..ba99453446 100644 --- a/src-tauri/crates/agent-core/src/state/mod.rs +++ b/src-tauri/crates/agent-core/src/state/mod.rs @@ -13,6 +13,7 @@ pub mod commands; pub mod control_flow; pub mod integrations_store; +mod session_identity; mod session_runtime; mod unified; @@ -20,5 +21,6 @@ mod unified; // deeper `state::integrations_store::*` path; the module is public so // background subsystems and the main crate can call the process-wide // `integrations_store()` accessor. -pub use session_runtime::{AgentSession, SessionRuntime}; +pub use session_identity::session_identity_lock; +pub use session_runtime::{AgentSession, SessionIdentityMutationGuard, SessionRuntime}; pub use unified::AgentAppState; diff --git a/src-tauri/crates/agent-core/src/state/session_identity.rs b/src-tauri/crates/agent-core/src/state/session_identity.rs new file mode 100644 index 0000000000..cba508ed27 --- /dev/null +++ b/src-tauri/crates/agent-core/src/state/session_identity.rs @@ -0,0 +1,52 @@ +//! Serialize identity admission and picker writes for one session. +//! +//! Rust turns hold this only through resolve/init/persistence, then release it +//! before execution. CLI runners retain it through provider-native publication, +//! whose files are account-bound. Idle sessions retain no mutex; dead weak +//! entries are pruned on every access. + +use std::collections::HashMap; +use std::sync::{Arc, LazyLock, Weak}; +use tokio::sync::Mutex; + +static LOCKS: LazyLock>>>> = + LazyLock::new(|| Mutex::new(HashMap::new())); + +pub async fn session_identity_lock(session_id: &str) -> Arc> { + let mut locks = LOCKS.lock().await; + locks.retain(|_, lock| lock.strong_count() > 0); + if let Some(lock) = locks.get(session_id).and_then(Weak::upgrade) { + return lock; + } + let lock = Arc::new(Mutex::new(())); + locks.insert(session_id.to_string(), Arc::downgrade(&lock)); + lock +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn same_session_serializes_and_unrelated_sessions_remain_independent() { + let first = session_identity_lock("identity-one").await; + let second = session_identity_lock("identity-one").await; + let other = session_identity_lock("identity-two").await; + assert!(Arc::ptr_eq(&first, &second)); + let guard = first.lock().await; + assert!(second.try_lock().is_err()); + assert!(other.try_lock().is_ok()); + drop(guard); + assert!(second.try_lock().is_ok()); + } + + #[tokio::test] + async fn idle_identity_locks_are_released_and_pruned_on_next_access() { + let lock = session_identity_lock("identity-released").await; + let weak = Arc::downgrade(&lock); + drop(lock); + assert!(weak.upgrade().is_none()); + let _next = session_identity_lock("identity-next").await; + assert!(!LOCKS.lock().await.contains_key("identity-released")); + } +} diff --git a/src-tauri/crates/agent-core/src/state/session_identity_mutation_tests.rs b/src-tauri/crates/agent-core/src/state/session_identity_mutation_tests.rs new file mode 100644 index 0000000000..936d15ac90 --- /dev/null +++ b/src-tauri/crates/agent-core/src/state/session_identity_mutation_tests.rs @@ -0,0 +1,153 @@ +use super::*; + +async fn fixture() -> (Arc, Arc) { + let definition = AgentDefinition { + selected_model_id: Some("e2e-fake-provider".into()), + ..Default::default() + }; + let runtime = Arc::new(SessionRuntime { + provider: Arc::new(crate::providers::e2e_fake::E2eFakeProvider), + tool_registry: Arc::new(ToolRegistry::new()), + policy: Arc::new(ResolvedToolPolicy::permissive()), + model: "e2e-fake-provider".into(), + account_id: Some("account-a".into()), + native_harness_type: None, + workspace_state: Arc::new(parking_lot::RwLock::new(SessionWorkspace::new( + std::env::temp_dir(), + ))), + mcp_auto_approved: vec![], + resolved: crate::definitions::resolved::ResolvedAgent::resolve( + &definition, + None, + &Default::default(), + ) + .unwrap(), + integrations_snapshot: Default::default(), + overrides: Default::default(), + agent_soul: None, + sovereign_prompt: false, + policy_context_activator: None, + agent_org_context: None, + agent_org_current_member_id: None, + agent_definition_id: None, + }); + let session = Arc::new(AgentSession::new("identity-mutation".into(), definition)); + session.set_runtime(runtime.clone()).await.unwrap(); + (session, runtime) +} + +#[tokio::test] +async fn pinned_admission_rejects_identity_edit_and_release_allows_retry() { + let (session, runtime) = fixture().await; + let reservation = session + .reserve_runtime_admission(&runtime, "prepared-turn") + .await + .unwrap(); + let error = match session.begin_identity_mutation().await { + Ok(_) => panic!("a pinned runtime must reject identity edits before persistence"), + Err(error) => error, + }; + assert!(error.starts_with("agent_org_runtime_admission_conflict:")); + assert!(Arc::ptr_eq(&runtime, &session.get_runtime().await.unwrap())); + assert!( + session + .runtime_admission_is_current(&reservation, &runtime) + .await + ); + session.release_runtime_admission(&reservation).await; + session + .begin_identity_mutation() + .await + .unwrap() + .invalidate_runtime() + .await; + assert!(session.get_runtime().await.is_none()); + assert_eq!(runtime.account_id.as_deref(), Some("account-a")); +} + +#[tokio::test] +async fn failed_identity_write_keeps_runtime_and_allows_new_admission() { + let (session, runtime) = fixture().await; + let mutation = session.begin_identity_mutation().await.unwrap(); + drop(mutation); // The persistence write failed; no invalidation is committed. + assert!(Arc::ptr_eq(&runtime, &session.get_runtime().await.unwrap())); + assert!(session + .reserve_runtime_admission(&runtime, "retry") + .await + .is_ok()); +} + +#[tokio::test] +async fn new_admission_waits_for_identity_write_then_rejects_stale_provider() { + let (session, runtime) = fixture().await; + let mutation = session.begin_identity_mutation().await.unwrap(); + let admission = { + let session = session.clone(); + let runtime = runtime.clone(); + tokio::spawn(async move { + session + .reserve_runtime_admission(&runtime, "competing-turn") + .await + }) + }; + tokio::task::yield_now().await; + assert!(!admission.is_finished()); + mutation.invalidate_runtime().await; + let error = admission.await.unwrap().unwrap_err(); + assert!(error.starts_with("agent_org_runtime_admission_stale:")); + assert!(session.get_runtime().await.is_none()); +} + +#[tokio::test] +async fn admitted_turn_control_keeps_old_lease_and_cannot_clear_replacement() { + let (session, admitted) = fixture().await; + let admitted_lease = session.runtime_lease_for(&admitted).await.unwrap(); + session + .begin_identity_mutation() + .await + .unwrap() + .invalidate_runtime() + .await; + let (_, replacement) = fixture().await; + let replacement_lease = session.set_runtime(replacement.clone()).await.unwrap(); + assert_ne!(admitted_lease, replacement_lease); + assert!(session.runtime_lease_for(&admitted).await.is_err()); + let generation = session + .begin_turn_with_runtime_lease( + "queued turn".into(), + Some("turn-a".into()), + Some(admitted_lease.clone()), + ) + .await; + let turn = session.runtime_turn_identity().await.unwrap(); + assert_eq!(turn.runtime_lease_id, admitted_lease); + assert_eq!( + session + .turn_process_control() + .unwrap() + .owner + .runtime_lease_id, + admitted_lease + ); + assert!( + !session + .release_runtime_if_current(&admitted_lease, &generation) + .await + ); + assert!(Arc::ptr_eq( + &replacement, + &session.get_runtime().await.unwrap() + )); + session + .end_turn(DialogTurnState::Completed, TurnStats::default()) + .await; + assert!( + !session + .release_runtime_lease_if_current(&admitted_lease) + .await + ); + assert!(Arc::ptr_eq( + &replacement, + &session.get_runtime().await.unwrap() + )); +} diff --git a/src-tauri/crates/agent-core/src/state/session_runtime.rs b/src-tauri/crates/agent-core/src/state/session_runtime.rs index c6e3f77105..e50bbbb5e8 100644 --- a/src-tauri/crates/agent-core/src/state/session_runtime.rs +++ b/src-tauri/crates/agent-core/src/state/session_runtime.rs @@ -126,6 +126,20 @@ struct RuntimeAdmissionSlot { holders: usize, } +/// Holds the admission CAS domain while a model/account edit is persisted. +/// Dropping this guard on a failed write leaves the current runtime untouched. +/// A successful write must clear the cached runtime before releasing the guard. +pub struct SessionIdentityMutationGuard<'a> { + session: &'a AgentSession, + _admissions: tokio::sync::MutexGuard<'a, Vec>, +} + +impl SessionIdentityMutationGuard<'_> { + pub async fn invalidate_runtime(self) { + *self.session.runtime.write().await = None; + } +} + /// Exact in-memory identity of the Turn currently using a runtime lease. #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) struct RuntimeTurnIdentity { @@ -495,6 +509,23 @@ impl AgentSession { .map(|slot| Arc::clone(&slot.runtime)) } + /// Capture the lease alongside the admitted provider, before a picker may + /// replace the cache. Queued turns must not bind control to a later lease. + pub(crate) async fn runtime_lease_for( + &self, + runtime: &Arc, + ) -> Result { + self.runtime + .read() + .await + .as_ref() + .filter(|slot| Arc::ptr_eq(&slot.runtime, runtime)) + .map(|slot| slot.lease_id.clone()) + .ok_or_else(|| { + "session_runtime_admission_stale: admitted provider has no current lease".into() + }) + } + /// Prepare an exact runtime lease before durable DirectMember admission. /// Exact retries reuse their token; a different runtime or missing slot /// fails closed before any database row claims that the turn was accepted. @@ -567,6 +598,25 @@ impl AgentSession { release_runtime_admission_slot(&mut admissions, reservation) } + /// Reject identity edits before persistence when an accepted Member turn + /// still pins the provider. Holding this guard prevents a new admission + /// from appearing between validation and runtime invalidation. + pub async fn begin_identity_mutation( + &self, + ) -> Result, String> { + let admissions = self.runtime_admissions.lock().await; + if !admissions.is_empty() { + return Err( + "agent_org_runtime_admission_conflict: a prepared Member turn pins the current runtime" + .to_string(), + ); + } + Ok(SessionIdentityMutationGuard { + session: self, + _admissions: admissions, + }) + } + /// Clear whichever runtime is current. This remains the ordinary SDE /// invalidation path; Pause uses the conditional lease method below. pub(crate) async fn invalidate_runtime(&self) { @@ -779,6 +829,18 @@ impl AgentSession { .await .as_ref() .map(|slot| slot.lease_id.clone()); + self.begin_turn_with_runtime_lease(user_input, turn_intent_id, runtime_lease_id) + .await + } + + /// Start an admitted turn with its captured runtime lease, even when the + /// next-turn cache was invalidated or replaced while it waited in the queue. + pub(crate) async fn begin_turn_with_runtime_lease( + &self, + user_input: String, + turn_intent_id: Option, + runtime_lease_id: Option, + ) -> String { let turn = DialogTurn::new(user_input, Arc::clone(&self.cancel_flag)); let turn_id = turn.turn_id.clone(); let process_control = runtime_lease_id.as_ref().zip(turn_intent_id.as_ref()).map( @@ -1021,6 +1083,10 @@ fn runtime_lease_identity_matches(current_lease_id: Option<&str>, expected_lease current_lease_id == Some(expected_lease_id) } +#[cfg(test)] +#[path = "session_identity_mutation_tests.rs"] +mod identity_mutation_tests; + #[cfg(test)] mod runtime_lease_tests { use std::time::Duration; diff --git a/src-tauri/crates/key-vault/src/auto_detect/suggestions/cc_switch.rs b/src-tauri/crates/key-vault/src/auto_detect/suggestions/cc_switch.rs index 5d8fc5cf6a..4fea213edc 100644 --- a/src-tauri/crates/key-vault/src/auto_detect/suggestions/cc_switch.rs +++ b/src-tauri/crates/key-vault/src/auto_detect/suggestions/cc_switch.rs @@ -270,10 +270,23 @@ pub(super) mod fixtures { /// Representative rows: two Official profiles with empty blobs, one /// Claude relay, one Codex relay with a TOML config, one Gemini key. - pub(crate) fn sample_rows() -> Vec<(&'static str, &'static str, &'static str, &'static str, bool)> { + pub(crate) fn sample_rows( + ) -> Vec<(&'static str, &'static str, &'static str, &'static str, bool)> { vec![ - ("official", "claude", "Claude Official", r#"{"env":{}}"#, true), - ("official", "codex", "OpenAI Official", r#"{"auth":{},"config":""}"#, true), + ( + "official", + "claude", + "Claude Official", + r#"{"env":{}}"#, + true, + ), + ( + "official", + "codex", + "OpenAI Official", + r#"{"auth":{},"config":""}"#, + true, + ), ( "longcat", "claude", diff --git a/src-tauri/crates/key-vault/src/auto_detect/suggestions/codex_config.rs b/src-tauri/crates/key-vault/src/auto_detect/suggestions/codex_config.rs index 8bd62b9535..5576f794e6 100644 --- a/src-tauri/crates/key-vault/src/auto_detect/suggestions/codex_config.rs +++ b/src-tauri/crates/key-vault/src/auto_detect/suggestions/codex_config.rs @@ -14,7 +14,10 @@ pub(super) struct CodexModelProvider { pub env_key: Option, } -pub(super) fn codex_config_path_in(codex_home: Option<&str>, home: Option<&Path>) -> Option { +pub(super) fn codex_config_path_in( + codex_home: Option<&str>, + home: Option<&Path>, +) -> Option { if let Some(dir) = codex_home.map(str::trim).filter(|d| !d.is_empty()) { return Some(PathBuf::from(dir).join("config.toml")); } diff --git a/src-tauri/crates/key-vault/src/commands/crud.rs b/src-tauri/crates/key-vault/src/commands/crud.rs index d14ef27166..bcf3cb244e 100644 --- a/src-tauri/crates/key-vault/src/commands/crud.rs +++ b/src-tauri/crates/key-vault/src/commands/crud.rs @@ -9,6 +9,7 @@ mod models; mod projections; mod read_delete; mod save; +mod save_default_variants; pub use clipboard::{clipboard_write_image, clipboard_write_text}; pub use dtos::{ @@ -25,7 +26,7 @@ pub use read_delete::{ pub use save::save_key; pub(super) use models::oauth_model_metadata; -pub(super) use projections::key_info_from_entry; +pub use projections::key_info_from_entry; // Re-exported here so consumers keep the established key_vault::commands path // (matching model_supports_output_config_effort). diff --git a/src-tauri/crates/key-vault/src/commands/crud/dtos.rs b/src-tauri/crates/key-vault/src/commands/crud/dtos.rs index 3d5fb9d1e5..97f97943c0 100644 --- a/src-tauri/crates/key-vault/src/commands/crud/dtos.rs +++ b/src-tauri/crates/key-vault/src/commands/crud/dtos.rs @@ -108,6 +108,8 @@ pub struct SaveKeyRequest { pub model_aliases: Option>, pub model_variants: Option>, pub default_variants: Option>, + /// Family-scoped user choices; exclusive of the full account edit fields. + pub default_variant_overrides: Option>, pub quota_info: Option, pub has_local_key: Option, pub is_listed: Option, @@ -119,6 +121,8 @@ pub struct SaveKeyRequest { /// Full key response (unmasked, for internal use) #[derive(serde::Serialize)] pub struct FullKeyResponse { + pub credential_generation: u64, + pub model_catalog_generation: u64, pub id: String, pub name: Option, pub agent_type: String, diff --git a/src-tauri/crates/key-vault/src/commands/crud/models.rs b/src-tauri/crates/key-vault/src/commands/crud/models.rs index 559f09437c..7b065671d9 100644 --- a/src-tauri/crates/key-vault/src/commands/crud/models.rs +++ b/src-tauri/crates/key-vault/src/commands/crud/models.rs @@ -390,6 +390,18 @@ pub(super) fn default_variants_for_key(entry: &ModelKey) -> Vec Vec } out } + +/// Shared selectable catalog for desktop, mobile and external harness setup. +/// Empty enablement is an explicit empty selection; discovery never grants it. +impl super::KeyInfo { + pub fn selectable_model_ids(&self) -> Vec { + if !self.enabled { + return Vec::new(); + } + let enabled: std::collections::HashSet<&str> = + self.enabled_models.iter().map(String::as_str).collect(); + let bases: std::collections::HashMap<&str, &str> = self + .model_variants + .iter() + .map(|variant| (variant.model.as_str(), variant.base_model.as_str())) + .collect(); + let mut seen = std::collections::HashSet::new(); + self.available_models + .iter() + .chain(self.model_variants.iter().map(|variant| &variant.model)) + .filter(|model| { + (enabled.contains(model.as_str()) + || bases + .get(model.as_str()) + .is_some_and(|base| enabled.contains(base))) + && seen.insert(model.as_str()) + }) + .cloned() + .collect() + } +} diff --git a/src-tauri/crates/key-vault/src/commands/crud/projections.rs b/src-tauri/crates/key-vault/src/commands/crud/projections.rs index 31c4d63c23..a2cf127643 100644 --- a/src-tauri/crates/key-vault/src/commands/crud/projections.rs +++ b/src-tauri/crates/key-vault/src/commands/crud/projections.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use super::models::{default_variants_for_key, model_variants_for_key}; -use super::{DefaultVariantInfo, FullKeyResponse, KeyInfo, ModelAliasInfo, ModelVariantInfo}; +use super::{FullKeyResponse, KeyInfo, ModelAliasInfo, ModelVariantInfo}; use crate::commands::validate::key_can_refresh_quota; use crate::key_store::{AuthMethod, HealthStatus, ModelKey, ModelType}; @@ -105,22 +105,10 @@ fn enrich_cursor_native_models(info: &mut KeyInfo) -> Result<(), String> { merge_unique_models(&mut info.available_models, models); } - if info.enabled_models.is_empty() { - for model in CURSOR_NATIVE_FALLBACK_MODELS { - if info - .available_models - .iter() - .any(|available| available == model) - { - info.enabled_models.push(model.to_string()); - } - } - } - Ok(()) } -pub(in crate::commands) fn key_info_from_entry(entry: ModelKey) -> Result { +pub fn key_info_from_entry(entry: ModelKey) -> Result { let mut info = KeyInfo::from(entry); enrich_cursor_native_models(&mut info)?; Ok(info) @@ -224,7 +212,10 @@ impl From for KeyInfo { impl From for FullKeyResponse { fn from(entry: ModelKey) -> Self { + let default_variants = default_variants_for_key(&entry); FullKeyResponse { + credential_generation: entry.credential_generation, + model_catalog_generation: entry.model_catalog_generation, id: entry.id, name: entry.name, agent_type: entry.model_type.as_str().to_string(), @@ -255,14 +246,7 @@ impl From for FullKeyResponse { context_window: variant.context_window.filter(|ctx| *ctx > 0), }) .collect(), - default_variants: entry - .default_variants - .into_iter() - .map(|variant| DefaultVariantInfo { - base_model: variant.base_model, - model: variant.model, - }) - .collect(), + default_variants, auth_method: match entry.auth_method { AuthMethod::ApiKey => "api_key", AuthMethod::Oauth => "oauth", diff --git a/src-tauri/crates/key-vault/src/commands/crud/read_delete.rs b/src-tauri/crates/key-vault/src/commands/crud/read_delete.rs index 18ec3ceee6..169c2ad2d8 100644 --- a/src-tauri/crates/key-vault/src/commands/crud/read_delete.rs +++ b/src-tauri/crates/key-vault/src/commands/crud/read_delete.rs @@ -2,7 +2,7 @@ use std::collections::HashMap; use super::{key_info_from_entry, FullKeyResponse, KeyInfo}; use crate::commands::validate::invalidate_key_quota_runtime; -use crate::key_store::{HealthStatus, ModelType, KEY_SERVICE}; +use crate::key_store::{HealthStatus, ModelCatalogRefresh, ModelType, KEY_SERVICE}; /// List all stored keys (masked) #[tauri::command] @@ -100,7 +100,8 @@ pub async fn delete_key_by_id(key_id: String) -> Result { .map_err(|err| format!("Task join error: {}", err))? } -/// Update key health status after validation +/// Update key health status after validation, or commit an independent discovery snapshot. +#[allow(clippy::too_many_arguments)] #[tauri::command] pub async fn update_key_health( key_id: String, @@ -110,8 +111,15 @@ pub async fn update_key_health( enabled_models: Option>, quota_info: Option, model_context_lengths: Option>, + catalog_refresh: Option, ) -> Result, String> { tokio::task::spawn_blocking(move || { + if let Some(refresh) = catalog_refresh { + return KEY_SERVICE + .refresh_model_catalog(&key_id, refresh) + .and_then(key_info_from_entry) + .map(Some); + } let status = match health_status.as_str() { "valid" => HealthStatus::Valid, "degraded" => HealthStatus::Degraded, diff --git a/src-tauri/crates/key-vault/src/commands/crud/save.rs b/src-tauri/crates/key-vault/src/commands/crud/save.rs index 6c2742fde4..361fc246e4 100644 --- a/src-tauri/crates/key-vault/src/commands/crud/save.rs +++ b/src-tauri/crates/key-vault/src/commands/crud/save.rs @@ -53,119 +53,162 @@ fn normalize_claude_official_oauth_routing(entry: &mut ModelKey) { entry.protocol = None; } +fn clear_discovery_defaults_for_catalog_reset(entry: &mut ModelKey, request: &SaveKeyRequest) { + if request.available_models.is_some() + && request.model_variants.is_some() + && request.default_variants.as_ref().is_some_and(Vec::is_empty) + { + entry.discovered_default_variants.clear(); + } +} + /// Save or update a key #[tauri::command] pub async fn save_key(request: SaveKeyRequest) -> Result { tokio::task::spawn_blocking(move || { - let agent_type = - ModelType::from_str(&request.agent_type).ok_or("Unknown agent type".to_string())?; - - // Load existing key if updating - let existing = match request.id.as_deref() { - Some(id) => KEY_SERVICE.get_key_by_id_checked(id)?, - None => None, - }; - - let prior_quota_revision = existing.as_ref().map(quota_credential_revision); - let mut entry = if let Some(existing) = existing { - existing - } else { - ModelKey::new(agent_type.clone()) - }; - let mut received_oauth_material = false; - - // Update fields - if let Some(id) = request.id { - entry.id = id; - } - if let Some(name) = request.name { - entry.name = Some(name); - } - if let Some(desc) = request.description { - entry.description = if desc.is_empty() { None } else { Some(desc) }; - } - entry.model_type = agent_type; - if let Some(key) = request.api_key { - let key = key.trim().to_string(); - entry.api_key = if key.is_empty() { None } else { Some(key) }; - } - if let Some(token) = request.session_token { - let token = token.trim().to_string(); - received_oauth_material = !token.is_empty(); - entry.session_token = if token.is_empty() { None } else { Some(token) }; - } - if let Some(url) = request.base_url { - entry.base_url = Some(url); - } - if let Some(protocol) = request.protocol { - entry.protocol = match protocol.as_str() { - "openai" => Some(ProviderProtocol::OpenAi), - "anthropic" => Some(ProviderProtocol::Anthropic), - _ => return Err(format!("Unknown provider protocol: {}", protocol)), - }; - } - if let Some(env) = request.env_vars { - received_oauth_material = - received_oauth_material || env.values().any(|value| !value.trim().is_empty()); - entry.env_vars = env; - } - if let Some(metadata) = request.account_metadata { - entry.account_metadata = metadata; - } - if let Some(models) = request.available_models { - entry.available_models = models; - } - if let Some(enabled) = request.enabled_models { - entry.enabled_models = enabled; - } - if let Some(aliases) = request.model_aliases { - entry.model_aliases = aliases - .into_iter() - .map(|a| crate::key_store::ModelAlias { - display_name: a.display_name, - alias: a.alias, - icon: a.icon, - }) - .collect(); - } - if let Some(variants) = request.model_variants { - entry.model_variants = variants.into_iter().map(ModelVariant::from).collect(); - } - if let Some(default_variants) = request.default_variants { - entry.default_variants = default_variants - .into_iter() - .map(|variant| DefaultVariant { - base_model: variant.base_model, - model: variant.model, - }) - .collect(); - } - if let Some(quota) = request.quota_info { - entry.quota_info = Some(quota); - } - if let Some(local) = request.has_local_key { - entry.has_local_key = local; - } - if let Some(listed) = request.is_listed { - entry.is_listed = listed; - } - if let Some(auth) = request.auth_method { - entry.auth_method = match auth.as_str() { - "oauth" => AuthMethod::Oauth, - _ => AuthMethod::ApiKey, - }; + if request.default_variant_overrides.is_some() { + return super::save_default_variants::save_default_variant_overrides(request); } - if let Some(listing) = request.listing_id { - entry.listing_id = if listing.is_empty() { - None - } else { - Some(listing) - }; - } - if let Some(enabled) = request.enabled { - entry.enabled = enabled; - entry.oauth_auto_disabled = false; - if enabled && entry.auth_method == AuthMethod::Oauth { + save_key_with_service(&KEY_SERVICE, request) + }) + .await + .map_err(|err| format!("Task join error: {}", err))? +} + +fn save_key_with_service( + service: &crate::key_store::KeyService, + request: SaveKeyRequest, +) -> Result { + let agent_type = + ModelType::from_str(&request.agent_type).ok_or("Unknown agent type".to_string())?; + + let request_id = request.id.clone(); + let (previous, saved) = service.edit_key( + request_id.as_deref(), + agent_type.clone(), + move |mut entry| { + clear_discovery_defaults_for_catalog_reset(&mut entry, &request); + let mut received_oauth_material = false; + + // Update fields + if let Some(id) = request.id { + entry.id = id; + } + if let Some(name) = request.name { + entry.name = Some(name); + } + if let Some(desc) = request.description { + entry.description = if desc.is_empty() { None } else { Some(desc) }; + } + entry.model_type = agent_type; + if let Some(key) = request.api_key { + let key = key.trim().to_string(); + entry.api_key = if key.is_empty() { None } else { Some(key) }; + } + if let Some(token) = request.session_token { + let token = token.trim().to_string(); + received_oauth_material = !token.is_empty(); + entry.session_token = if token.is_empty() { None } else { Some(token) }; + } + if let Some(url) = request.base_url { + entry.base_url = Some(url); + } + if let Some(protocol) = request.protocol { + entry.protocol = match protocol.as_str() { + "openai" => Some(ProviderProtocol::OpenAi), + "anthropic" => Some(ProviderProtocol::Anthropic), + _ => return Err(format!("Unknown provider protocol: {}", protocol)), + }; + } + if let Some(env) = request.env_vars { + received_oauth_material = + received_oauth_material || env.values().any(|value| !value.trim().is_empty()); + entry.env_vars = env; + } + if let Some(metadata) = request.account_metadata { + entry.account_metadata = metadata; + } + if let Some(models) = request.available_models { + entry.available_models = models; + } + if let Some(enabled) = request.enabled_models { + entry.enabled_models = enabled; + } + if let Some(aliases) = request.model_aliases { + entry.model_aliases = aliases + .into_iter() + .map(|a| crate::key_store::ModelAlias { + display_name: a.display_name, + alias: a.alias, + icon: a.icon, + }) + .collect(); + } + if let Some(variants) = request.model_variants { + entry.model_variants = variants.into_iter().map(ModelVariant::from).collect(); + } + if let Some(default_variants) = request.default_variants { + entry.default_variants = default_variants + .into_iter() + .map(|variant| DefaultVariant { + base_model: variant.base_model, + model: variant.model, + }) + .collect(); + } + if let Some(quota) = request.quota_info { + entry.quota_info = Some(quota); + } + if let Some(local) = request.has_local_key { + entry.has_local_key = local; + } + if let Some(listed) = request.is_listed { + entry.is_listed = listed; + } + if let Some(auth) = request.auth_method { + entry.auth_method = match auth.as_str() { + "oauth" => AuthMethod::Oauth, + _ => AuthMethod::ApiKey, + }; + } + if let Some(listing) = request.listing_id { + entry.listing_id = if listing.is_empty() { + None + } else { + Some(listing) + }; + } + if let Some(enabled) = request.enabled { + entry.enabled = enabled; + entry.oauth_auto_disabled = false; + if enabled && entry.auth_method == AuthMethod::Oauth { + entry.oauth_refresh_failure_count = 0; + entry.last_oauth_refresh_failed_at = None; + entry.last_validation_error = None; + entry.temporary_unavailable_until = None; + entry.temporary_unavailable_reason = None; + entry.last_upstream_status = None; + entry.last_upstream_error_type = None; + entry.rate_limit_reset_at = None; + if entry.health_status == HealthStatus::Invalid { + entry.health_status = HealthStatus::Unknown; + } + } + } + + // Normalize OAuth keys: only keep session_token, clear api_key. + // Cursor is the exception: we persist both credentials and let each + // runtime entry point choose the one it needs. + if entry.auth_method == AuthMethod::Oauth && entry.model_type != ModelType::CursorCli { + if entry.api_key.is_some() && entry.session_token.is_none() { + entry.session_token = entry.api_key.take(); + } + entry.api_key = None; + } + + normalize_claude_official_oauth_routing(&mut entry); + + if entry.auth_method == AuthMethod::Oauth && received_oauth_material { entry.oauth_refresh_failure_count = 0; entry.last_oauth_refresh_failed_at = None; entry.last_validation_error = None; @@ -174,66 +217,45 @@ pub async fn save_key(request: SaveKeyRequest) -> Result { entry.last_upstream_status = None; entry.last_upstream_error_type = None; entry.rate_limit_reset_at = None; - if entry.health_status == HealthStatus::Invalid { - entry.health_status = HealthStatus::Unknown; - } - } - } - - // Normalize OAuth keys: only keep session_token, clear api_key. - // Cursor is the exception: we persist both credentials and let each - // runtime entry point choose the one it needs. - if entry.auth_method == AuthMethod::Oauth && entry.model_type != ModelType::CursorCli { - if entry.api_key.is_some() && entry.session_token.is_none() { - entry.session_token = entry.api_key.take(); } - entry.api_key = None; - } - normalize_claude_official_oauth_routing(&mut entry); - - if entry.auth_method == AuthMethod::Oauth && received_oauth_material { - entry.oauth_refresh_failure_count = 0; - entry.last_oauth_refresh_failed_at = None; - entry.last_validation_error = None; - entry.temporary_unavailable_until = None; - entry.temporary_unavailable_reason = None; - entry.last_upstream_status = None; - entry.last_upstream_error_type = None; - entry.rate_limit_reset_at = None; - } - - if entry.model_type == ModelType::CursorCli { - if let Some(api_key) = entry.api_key.as_deref() { - if !(api_key.starts_with("key_") || api_key.starts_with("crsr_")) - || api_key.len() <= 20 - { - return Err("Cursor API key should start with 'key_' or 'crsr_'".to_string()); + if entry.model_type == ModelType::CursorCli { + if let Some(api_key) = entry.api_key.as_deref() { + if !(api_key.starts_with("key_") || api_key.starts_with("crsr_")) + || api_key.len() <= 20 + { + return Err( + "Cursor API key should start with 'key_' or 'crsr_'".to_string() + ); + } } - } - let session_token = entry.session_token.as_deref().unwrap_or_default(); - if session_token.is_empty() { - return Err("Cursor requires a session token before saving".to_string()); - } - if is_cursor_web_session_token(session_token) { - return Err( + let session_token = entry.session_token.as_deref().unwrap_or_default(); + if session_token.is_empty() { + return Err("Cursor requires a session token before saving".to_string()); + } + if is_cursor_web_session_token(session_token) { + return Err( "Cursor web login tokens cannot be used for native chat; please sign in again" .to_string(), ); + } } - } - let saved = KEY_SERVICE.save_key(entry)?; - let saved_quota_revision = quota_credential_revision(&saved); - if prior_quota_revision.as_deref() != Some(saved_quota_revision.as_str()) { - invalidate_key_quota_runtime(&saved.id); - } - key_info_from_entry(saved) - }) - .await - .map_err(|err| format!("Task join error: {}", err))? + Ok(entry) + }, + )?; + let prior_quota_revision = previous.as_ref().map(quota_credential_revision); + let saved_quota_revision = quota_credential_revision(&saved); + if prior_quota_revision.as_deref() != Some(saved_quota_revision.as_str()) { + invalidate_key_quota_runtime(&saved.id); + } + key_info_from_entry(saved) } +#[cfg(test)] +#[path = "save_atomic_tests.rs"] +mod atomic_tests; + #[cfg(test)] mod tests { use super::*; @@ -351,3 +373,38 @@ mod tests { ); } } + +#[cfg(test)] +mod catalog_reset_tests { + use super::*; + + #[test] + fn explicit_catalog_reset_removes_discovery_defaults_but_preference_reset_does_not() { + let dir = tempfile::tempdir().unwrap(); + let service = crate::key_store::KeyService::new(Some(dir.path().to_path_buf())); + let mut key = ModelKey::new(ModelType::CustomApi); + key.discovered_default_variants.push(DefaultVariant { + base_model: "model".into(), + model: "old-effort".into(), + }); + let mut key = service.save_key(key).unwrap(); + let preference_reset: SaveKeyRequest = serde_json::from_value(serde_json::json!({ + "agent_type": "custom_api", "default_variants": [] + })) + .unwrap(); + clear_discovery_defaults_for_catalog_reset(&mut key, &preference_reset); + assert_eq!(key.discovered_default_variants.len(), 1); + let catalog_reset: SaveKeyRequest = serde_json::from_value(serde_json::json!({ + "agent_type": "custom_api", "available_models": ["model"], + "model_variants": [], "default_variants": [] + })) + .unwrap(); + clear_discovery_defaults_for_catalog_reset(&mut key, &catalog_reset); + service.save_key(key.clone()).unwrap(); + assert!(service + .get_key_by_id(&key.id) + .unwrap() + .discovered_default_variants + .is_empty()); + } +} diff --git a/src-tauri/crates/key-vault/src/commands/crud/save_atomic_tests.rs b/src-tauri/crates/key-vault/src/commands/crud/save_atomic_tests.rs new file mode 100644 index 0000000000..dec97ee3f3 --- /dev/null +++ b/src-tauri/crates/key-vault/src/commands/crud/save_atomic_tests.rs @@ -0,0 +1,121 @@ +use std::sync::{mpsc, Arc, Barrier}; + +use super::*; +use crate::key_store::{HealthStatus, KeyService}; + +fn request(id: &str, extra: serde_json::Value) -> SaveKeyRequest { + let mut value = serde_json::json!({"id": id, "agent_type": "custom_api"}); + value + .as_object_mut() + .unwrap() + .extend(extra.as_object().unwrap().clone()); + serde_json::from_value(value).unwrap() +} + +#[test] +fn concurrent_account_edits_keep_both_independent_fields() { + let dir = tempfile::tempdir().unwrap(); + let service = Arc::new(KeyService::new(Some(dir.path().to_path_buf()))); + let mut key = ModelKey::new(ModelType::CustomApi); + key.available_models = vec!["model-a".into()]; + key.enabled_models = vec!["model-a".into()]; + let key = service.save_key(key).unwrap(); + let ready = Arc::new(Barrier::new(2)); + let release = Arc::new(Barrier::new(2)); + let (started_tx, started_rx) = mpsc::channel(); + + let writer_service = Arc::clone(&service); + let writer_id = key.id.clone(); + let writer_ready = Arc::clone(&ready); + let writer_release = Arc::clone(&release); + let first = std::thread::spawn(move || { + writer_service.edit_key( + Some(&writer_id), + ModelType::CustomApi, + move |mut current| { + writer_ready.wait(); + writer_release.wait(); + current.name = Some("renamed".into()); + Ok(current) + }, + ) + }); + ready.wait(); // The first editor owns the same lock used by the RPC save. + + let second_service = Arc::clone(&service); + let second_id = key.id.clone(); + let second = std::thread::spawn(move || { + started_tx.send(()).unwrap(); + save_key_with_service( + &second_service, + request(&second_id, serde_json::json!({"enabled_models": []})), + ) + }); + started_rx.recv().unwrap(); + release.wait(); + first.join().unwrap().unwrap(); + second.join().unwrap().unwrap(); + + let reloaded = KeyService::new(Some(dir.path().to_path_buf())) + .get_key_by_id(&key.id) + .unwrap(); + assert_eq!(reloaded.name.as_deref(), Some("renamed")); + assert!(reloaded.enabled_models.is_empty()); +} + +#[test] +fn ordinary_edit_keeps_current_discovery_and_health() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let key = service + .save_key(ModelKey::new(ModelType::CustomApi)) + .unwrap(); + let prepared = request(&key.id, serde_json::json!({"name": "renamed"})); + service + .update_key_health( + &key.id, + HealthStatus::Valid, + None, + Some(vec!["new-model".into()]), + None, + None, + None, + ) + .unwrap(); + + save_key_with_service(&service, prepared).unwrap(); + let reloaded = KeyService::new(Some(dir.path().to_path_buf())) + .get_key_by_id(&key.id) + .unwrap(); + assert_eq!(reloaded.name.as_deref(), Some("renamed")); + assert_eq!(reloaded.health_status, HealthStatus::Valid); + assert!(reloaded.available_models.contains(&"new-model".into())); +} + +#[test] +fn rejected_edit_leaves_file_unchanged_then_retry_succeeds() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let key = service + .save_key(ModelKey::new(ModelType::CustomApi)) + .unwrap(); + let file = dir.path().join("credentials.json"); + let before = std::fs::read(&file).unwrap(); + + assert!(save_key_with_service( + &service, + request(&key.id, serde_json::json!({"protocol": "invalid"})) + ) + .is_err()); + assert_eq!(std::fs::read(&file).unwrap(), before); + + save_key_with_service( + &service, + request(&key.id, serde_json::json!({"name": "retry"})), + ) + .unwrap(); + assert_eq!( + service.get_key_by_id(&key.id).unwrap().name.as_deref(), + Some("retry") + ); +} diff --git a/src-tauri/crates/key-vault/src/commands/crud/save_default_variants.rs b/src-tauri/crates/key-vault/src/commands/crud/save_default_variants.rs new file mode 100644 index 0000000000..d8bddb256e --- /dev/null +++ b/src-tauri/crates/key-vault/src/commands/crud/save_default_variants.rs @@ -0,0 +1,62 @@ +//! A default pick is a family-scoped intent, not a replacement of the +//! effective (user + provider + product fallback) defaults shown by the UI. +use super::{key_info_from_entry, KeyInfo, SaveKeyRequest}; +use crate::key_store::{DefaultVariant, ModelType, KEY_SERVICE}; + +pub(super) fn save_default_variant_overrides(request: SaveKeyRequest) -> Result { + // Keep this patch contract narrow. Silently ignoring another account edit + // would look successful even though it never reached persistence. + if request.name.is_some() + || request.description.is_some() + || request.api_key.is_some() + || request.session_token.is_some() + || request.base_url.is_some() + || request.protocol.is_some() + || request.env_vars.is_some() + || request.account_metadata.is_some() + || request.available_models.is_some() + || request.enabled_models.is_some() + || request.model_aliases.is_some() + || request.model_variants.is_some() + || request.default_variants.is_some() + || request.quota_info.is_some() + || request.has_local_key.is_some() + || request.is_listed.is_some() + || request.auth_method.is_some() + || request.listing_id.is_some() + || request.enabled.is_some() + { + return Err("Save default variant overrides separately from other account edits".into()); + } + let id = request + .id + .ok_or("Default variant overrides require an existing account")?; + let model_type = ModelType::from_str(&request.agent_type).ok_or("Unknown agent type")?; + let overrides = request + .default_variant_overrides + .unwrap_or_default() + .into_iter() + .map(|variant| DefaultVariant { + base_model: variant.base_model, + model: variant.model, + }) + .collect(); + KEY_SERVICE + .update_default_variant_overrides(&id, &model_type, overrides) + .and_then(key_info_from_entry) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn mixed_delta_and_full_account_edits_are_rejected_before_persistence() { + let request = serde_json::from_value(serde_json::json!({ + "id": "fixture", "agent_type": "custom_api", "name": "other edit", + "default_variant_overrides": [{ "base_model": "model", "model": "model-high" }] + })) + .unwrap(); + assert!(save_default_variant_overrides(request).is_err()); + } +} diff --git a/src-tauri/crates/key-vault/src/commands/prompt_polish/client.rs b/src-tauri/crates/key-vault/src/commands/prompt_polish/client.rs index 44251ebb7a..dea563c237 100644 --- a/src-tauri/crates/key-vault/src/commands/prompt_polish/client.rs +++ b/src-tauri/crates/key-vault/src/commands/prompt_polish/client.rs @@ -94,7 +94,11 @@ fn model_candidates(key: &ModelKey) -> Vec { for variant in &key.model_variants { push_unique(&mut candidates, &variant.model); } - for default_variant in &key.default_variants { + for default_variant in key + .default_variants + .iter() + .chain(&key.discovered_default_variants) + { push_unique(&mut candidates, &default_variant.model); } diff --git a/src-tauri/crates/key-vault/src/commands/validate.rs b/src-tauri/crates/key-vault/src/commands/validate.rs index f09c5bc568..9b5e156460 100644 --- a/src-tauri/crates/key-vault/src/commands/validate.rs +++ b/src-tauri/crates/key-vault/src/commands/validate.rs @@ -23,14 +23,14 @@ pub use oauth::{ }; pub use opencode::{validate_opencode_key, OPENCODE_GO_BASE_URL, OPENCODE_ZEN_BASE_URL}; pub use quota_dispatch::fetch_key_quota; -pub use suggestions::{ - import_credential_suggestions, list_credential_suggestions, CredentialImportItemReport, - CredentialImportReport, CredentialImportStatus, -}; pub use quota_refresh::{ get_key_quota_refresh_status, invalidate_key_quota_runtime, key_quota_refresh_status, refresh_key_quota, KeyQuotaRefreshAttemptInfo, KeyQuotaRefreshStatusInfo, }; +pub use suggestions::{ + import_credential_suggestions, list_credential_suggestions, CredentialImportItemReport, + CredentialImportReport, CredentialImportStatus, +}; // Only the `commands/tests` suite reaches this projection helper. #[cfg(test)] diff --git a/src-tauri/crates/key-vault/src/commands/validate/suggestions.rs b/src-tauri/crates/key-vault/src/commands/validate/suggestions.rs index d1d6a716c7..ac667756cf 100644 --- a/src-tauri/crates/key-vault/src/commands/validate/suggestions.rs +++ b/src-tauri/crates/key-vault/src/commands/validate/suggestions.rs @@ -174,6 +174,7 @@ async fn import_via_detector(selection: &CredentialSuggestion) -> Result Result { + /// Save a complete internal record, validating against the record read + /// under the same lock as the eventual file replacement. + pub fn save_key(&self, key: ModelKey) -> Result { + let key_id = key.id.clone(); + self.try_update_store_for_key(Some(&key_id), |store| Self::save_key_in_store(store, key)) + } + + /// Apply an RPC field patch to the latest persisted account. The caller + /// provides only explicitly edited fields; discovery and token refreshes + /// committed earlier are preserved even when the UI's snapshot is old. + pub(crate) fn edit_key( + &self, + key_id: Option<&str>, + model_type: ModelType, + edit: F, + ) -> Result<(Option, ModelKey), String> + where + F: FnOnce(ModelKey) -> Result, + { + self.try_update_store_for_key(key_id, |store| { + let previous = key_id.and_then(|id| store.get_by_id(id)).cloned(); + let entry = edit( + previous + .clone() + .unwrap_or_else(|| ModelKey::new(model_type)), + )?; + let saved = Self::save_key_in_store(store, entry)?; + Ok((previous, saved)) + }) + } + + fn save_key_in_store( + store: &mut super::super::store::KeyStore, + mut key: ModelKey, + ) -> Result { // Explicit aliases are user-owned request IDs, including IDs absent // from discovery. Validate before touching the persisted credential. // @@ -93,7 +126,7 @@ impl KeyService { // (rename, description, endpoint edits) carry the stored aliases along // unchanged, and a historical record that predates these rules must // not block every later write to the account. - let previous_key = self.get_key_by_id(&key.id); + let previous_key = store.get_by_id(&key.id).cloned(); let retained: HashMap> = previous_key .as_ref() .map(|existing| { @@ -161,31 +194,62 @@ impl KeyService { } super::codex_cli_auth::bind_codex_cli_source(&mut key); let key_id = key.id.clone(); - self.update_store(|store| { - if let Some(previous) = store.get_by_id(&key_id) { - key.credential_generation = previous.credential_generation; - key.codex_pending_source_token_hash = - previous.codex_pending_source_token_hash.clone(); - if !key.same_credential_material(previous) { - key.credential_generation = key - .credential_generation - .checked_add(1) - .ok_or_else(|| "Credential generation exhausted".to_string())?; - key.oauth_auto_disabled = false; - key.codex_pending_source_token_hash = None; - } - if previous.enabled && !key.enabled { - key.oauth_auto_disabled = false; - } - } else { - key.credential_generation = 0; + if let Some(previous) = store.get_by_id(&key_id) { + if key.model_catalog_generation != previous.model_catalog_generation { + return Err("Model catalog changed during save; retry the account edit".into()); } - store.set(key); - Ok(store - .get_by_id(&key_id) - .cloned() - .expect("KeyStore::set must retain the inserted key")) - })? + key.credential_generation = previous.credential_generation; + // Discovery defaults belong to a provider route. A new endpoint, + // protocol or provider must not inherit the old route's ladder. + if key.model_type != previous.model_type + || key.base_url != previous.base_url + || key.protocol != previous.protocol + { + key.discovered_default_variants.clear(); + } + let aliases: HashSet<&str> = key + .model_aliases + .iter() + .chain(&previous.model_aliases) + .map(|alias| alias.alias.as_str()) + .collect(); + let discovered_ids = |entry: &ModelKey| { + entry + .available_models + .iter() + .filter(|id| !aliases.contains(id.as_str())) + .cloned() + .collect::>() + }; + if discovered_ids(&key) != discovered_ids(previous) + || key.model_variants != previous.model_variants + || key.discovered_default_variants != previous.discovered_default_variants + { + key.model_catalog_generation = key + .model_catalog_generation + .checked_add(1) + .ok_or_else(|| "Model catalog generation exhausted".to_string())?; + } + key.codex_pending_source_token_hash = previous.codex_pending_source_token_hash.clone(); + if !key.same_credential_material(previous) { + key.credential_generation = key + .credential_generation + .checked_add(1) + .ok_or_else(|| "Credential generation exhausted".to_string())?; + key.oauth_auto_disabled = false; + key.codex_pending_source_token_hash = None; + } + if previous.enabled && !key.enabled { + key.oauth_auto_disabled = false; + } + } else { + key.credential_generation = 0; + } + store.set(key); + Ok(store + .get_by_id(&key_id) + .cloned() + .expect("KeyStore::set must retain the inserted key")) } /// Record behaviorally-observed reasoning capability for `model` on key @@ -279,6 +343,14 @@ impl KeyService { ) -> Result, String> { self.update_store(|store| { if let Some(entry) = store.keys.get_mut(key_id) { + if available_models.is_some() + || model_context_lengths.is_some_and(|contexts| !contexts.is_empty()) + { + entry.model_catalog_generation = entry + .model_catalog_generation + .checked_add(1) + .ok_or_else(|| "Model catalog generation exhausted".to_string())?; + } entry.health_status = health_status; entry.last_validation_error = error_message; entry.last_validated_at = Some(Utc::now()); @@ -359,11 +431,11 @@ impl KeyService { entry.updated_at = Utc::now(); store.updated_at = Utc::now(); - Some(entry.clone()) + Ok(Some(entry.clone())) } else { - None + Ok(None) } - }) + })? } /// Delete key by agent type and optional ID. diff --git a/src-tauri/crates/key-vault/src/key_store/service/mod.rs b/src-tauri/crates/key-vault/src/key_store/service/mod.rs index 8a93173daf..a88ec7c354 100644 --- a/src-tauri/crates/key-vault/src/key_store/service/mod.rs +++ b/src-tauri/crates/key-vault/src/key_store/service/mod.rs @@ -15,6 +15,8 @@ mod claude_oauth; mod codex_cli_auth; mod codex_oauth; mod keys; +mod model_catalog; +pub use model_catalog::ModelCatalogRefresh; mod oauth_health; mod persistence; mod token_sync; diff --git a/src-tauri/crates/key-vault/src/key_store/service/model_catalog.rs b/src-tauri/crates/key-vault/src/key_store/service/model_catalog.rs new file mode 100644 index 0000000000..5134a4ad42 --- /dev/null +++ b/src-tauri/crates/key-vault/src/key_store/service/model_catalog.rs @@ -0,0 +1,191 @@ +//! Discovery-only write boundary. User choices are read under the store lock, +//! never replayed from the UI snapshot that started a network request. +use std::collections::{HashMap, HashSet}; + +use super::super::{DefaultVariant, ModelKey, ModelVariant}; +use super::KeyService; + +#[derive(Debug, Clone, serde::Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct ModelCatalogRefresh { + pub expected_credential_generation: u64, + pub expected_catalog_generation: u64, + pub available_models: Vec, + /// None means this discovery protocol does not describe effort variants. + pub model_variants: Option>, + pub default_variants: Option>, + pub model_context_lengths: HashMap, +} + +impl KeyService { + pub fn refresh_model_catalog( + &self, + key_id: &str, + refresh: ModelCatalogRefresh, + ) -> Result { + if refresh.available_models.is_empty() { + return Err("Provider returned an empty model list".into()); + } + self.update_store(|store| { + let entry = store + .keys + .get_mut(key_id) + .ok_or("Account no longer exists")?; + if entry.credential_generation != refresh.expected_credential_generation + || entry.model_catalog_generation != refresh.expected_catalog_generation + { + return Err( + "Account or model catalog changed during discovery; refresh again".into(), + ); + } + let next_generation = entry + .model_catalog_generation + .checked_add(1) + .ok_or("Model catalog generation exhausted")?; + let aliases: HashSet<&str> = entry + .model_aliases + .iter() + .map(|alias| alias.alias.as_str()) + .collect(); + let incoming: HashSet<&str> = refresh + .available_models + .iter() + .map(String::as_str) + .collect(); + let discovered_contexts: HashMap = refresh + .model_variants + .as_ref() + .into_iter() + .flatten() + .filter_map(|variant| { + variant + .context_window + .filter(|value| *value > 0) + .map(|value| (variant.model.clone(), value)) + }) + .collect(); + if let Some(variants) = refresh.model_variants { + // A live ladder replaces the old ladder. Keep manual IDs and + // behaviorally observed bare records; discovery must not erase + // explicit custom requests or runtime capability observations. + entry.model_variants.retain(|variant| { + aliases.contains(variant.model.as_str()) + || (incoming.contains(variant.model.as_str()) + && variant.model == variant.base_model + && variant.reasoning.is_some() + && !variants.iter().any(|new| new.model == variant.model)) + }); + for mut variant in variants { + if aliases.contains(variant.model.as_str()) { + continue; + } + variant.context_window = variant.context_window.filter(|value| *value > 0); + entry.model_variants.push(variant); + } + } else { + entry.model_variants.retain(|variant| { + aliases.contains(variant.model.as_str()) + || incoming.contains(variant.model.as_str()) + || incoming.contains(variant.base_model.as_str()) + }); + } + for variant in &mut entry.model_variants { + if !aliases.contains(variant.model.as_str()) + && incoming.contains(variant.model.as_str()) + { + variant.context_window = refresh + .model_context_lengths + .get(&variant.model) + .copied() + .filter(|value| *value > 0) + .or_else(|| discovered_contexts.get(&variant.model).copied()); + } + } + for (model, context_window) in refresh.model_context_lengths { + if context_window == 0 + || !incoming.contains(model.as_str()) + || aliases.contains(model.as_str()) + { + continue; + } + if let Some(variant) = entry + .model_variants + .iter_mut() + .find(|variant| variant.model == model) + { + variant.context_window = Some(context_window); + } else { + entry.model_variants.push(ModelVariant { + base_model: model.clone(), + model, + reasoning: None, + fast: false, + context_window: Some(context_window), + }); + } + } + if let Some(defaults) = refresh.default_variants { + entry.discovered_default_variants = defaults; + } + entry.available_models = refresh.available_models; + for alias in &entry.model_aliases { + if !entry.available_models.contains(&alias.alias) { + entry.available_models.push(alias.alias.clone()); + } + } + // Product-supported completion remains distinct from live discovery. + entry.normalize_model_catalog(); + entry.model_catalog_generation = next_generation; + entry.updated_at = chrono::Utc::now(); + store.updated_at = entry.updated_at; + Ok(entry.clone()) + })? + } +} + +impl KeyService { + /// Merge only explicitly chosen families, reading the current user layer + /// under the persistence lock so concurrent family picks cannot erase one + /// another or turn provider defaults into pinned user choices. + pub fn update_default_variant_overrides( + &self, + key_id: &str, + model_type: &super::super::ModelType, + overrides: Vec, + ) -> Result { + self.update_store(|store| { + let entry = store + .keys + .get_mut(key_id) + .ok_or("Account no longer exists")?; + if &entry.model_type != model_type { + return Err("Account provider changed; choose the model again".into()); + } + let mut families = HashSet::new(); + for choice in &overrides { + if !families.insert(&choice.base_model) { + return Err("Duplicate model family default".into()); + } + if choice.base_model.is_empty() + || choice.model.is_empty() + || choice + .base_model + .chars() + .chain(choice.model.chars()) + .any(|character| character.is_whitespace() || character.is_control()) + { + return Err("Default model and family must be non-empty request IDs".into()); + } + } + for choice in overrides { + entry + .default_variants + .retain(|previous| previous.base_model != choice.base_model); + entry.default_variants.push(choice); + } + entry.updated_at = chrono::Utc::now(); + store.updated_at = entry.updated_at; + Ok(entry.clone()) + })? + } +} diff --git a/src-tauri/crates/key-vault/src/key_store/service/persistence.rs b/src-tauri/crates/key-vault/src/key_store/service/persistence.rs index 7941d99a9d..afadfda818 100644 --- a/src-tauri/crates/key-vault/src/key_store/service/persistence.rs +++ b/src-tauri/crates/key-vault/src/key_store/service/persistence.rs @@ -147,6 +147,32 @@ impl KeyService { Ok(result) } + + /// Fallible account edit. The caller's read, validation, and mutation all + /// happen under one lock; a rejected edit never rewrites credentials.json. + pub(crate) fn try_update_store_for_key( + &self, + key_id: Option<&str>, + updater: F, + ) -> Result + where + F: FnOnce(&mut KeyStore) -> Result, + { + let _guard = self.lock.lock().map_err(|e| format!("Lock error: {}", e))?; + let mut loaded = self.load_store_checked()?; + if let Some(key_id) = key_id { + if let Some(error) = loaded.invalid_credential_error(key_id) { + return Err(error); + } + } + loaded.ensure_any_valid_credentials(&self.storage_file)?; + let result = updater(&mut loaded.store)?; + if !loaded.invalid_credentials.is_empty() { + loaded.log_diagnostics(&self.storage_file); + } + self.save_store(&loaded.store, &loaded.invalid_credentials)?; + Ok(result) + } } fn deserialize_key_store(contents: &str) -> Result { diff --git a/src-tauri/crates/key-vault/src/key_store/tests/mod.rs b/src-tauri/crates/key-vault/src/key_store/tests/mod.rs index 04729d5640..1f8c7787b8 100644 --- a/src-tauri/crates/key-vault/src/key_store/tests/mod.rs +++ b/src-tauri/crates/key-vault/src/key_store/tests/mod.rs @@ -8,3 +8,5 @@ mod model_type_tests; mod tests; mod oauth_generation_tests; + +mod model_catalog_tests; diff --git a/src-tauri/crates/key-vault/src/key_store/tests/model_catalog_tests.rs b/src-tauri/crates/key-vault/src/key_store/tests/model_catalog_tests.rs new file mode 100644 index 0000000000..1f8096ac4c --- /dev/null +++ b/src-tauri/crates/key-vault/src/key_store/tests/model_catalog_tests.rs @@ -0,0 +1,391 @@ +use crate::commands::KeyInfo; +use crate::key_store::{ + AuthMethod, DefaultVariant, HealthStatus, KeyService, ModelAlias, ModelCatalogRefresh, + ModelKey, ModelType, ModelVariant, +}; +use std::collections::HashMap; + +fn variant(model: &str, base: &str, context_window: Option) -> ModelVariant { + ModelVariant { + model: model.into(), + base_model: base.into(), + reasoning: Some("high".into()), + fast: false, + context_window, + } +} +fn default(model: &str, base: &str) -> DefaultVariant { + DefaultVariant { + model: model.into(), + base_model: base.into(), + } +} +fn discovery(key: &ModelKey) -> ModelCatalogRefresh { + ModelCatalogRefresh { + expected_credential_generation: key.credential_generation, + expected_catalog_generation: key.model_catalog_generation, + available_models: vec!["model".into(), "new-family".into()], + model_variants: Some(vec![ + variant("model-high", "model", Some(9000)), + variant("new-family-high", "new-family", None), + ]), + default_variants: Some(vec![ + default("model-high", "model"), + default("new-family-high", "new-family"), + ]), + model_context_lengths: HashMap::from([("model".into(), 9000)]), + } +} + +#[test] +fn refresh_commits_full_metadata_and_preserves_latest_user_choices_after_reload() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let mut key = ModelKey::new(ModelType::CustomApi); + key.api_key = Some("fixture-key".into()); + key.available_models = vec!["model".into(), "removed".into()]; + key.enabled_models = vec!["model".into()]; + key.model_variants = vec![variant("model-low", "model", Some(1000))]; + let snapshot = service.save_key(key).unwrap(); + // The request began before the user disabled the family, chose a default, + // and added a manual request ID. Merge these values under the write lock. + let mut latest = snapshot.clone(); + latest.enabled_models = vec![]; + latest.health_status = HealthStatus::Degraded; + latest.default_variants = vec![default("model-manual", "model")]; + latest.model_aliases.push(ModelAlias { + alias: "manual".into(), + display_name: "Manual".into(), + icon: None, + }); + service.save_key(latest).unwrap(); + service + .refresh_model_catalog(&snapshot.id, discovery(&snapshot)) + .unwrap(); + let reloaded = KeyService::new(Some(dir.path().to_path_buf())) + .get_key_by_id(&snapshot.id) + .unwrap(); + assert!(reloaded.enabled_models.is_empty()); + assert!(matches!(reloaded.health_status, HealthStatus::Degraded)); + assert!(reloaded.available_models.contains(&"manual".into())); + assert!(!reloaded.available_models.contains(&"removed".into())); + assert!(!reloaded + .model_variants + .iter() + .any(|v| v.model == "model-low")); + assert!(reloaded + .model_variants + .iter() + .any(|v| v.model == "model-high" && v.context_window == Some(9000))); + assert_eq!(reloaded.discovered_default_variants.len(), 2); + let info = KeyInfo::from(reloaded); + assert!(info.selectable_model_ids().is_empty()); + assert!(info + .default_variants + .iter() + .any(|v| v.base_model == "model" && v.model == "model-manual")); + assert!(info + .default_variants + .iter() + .any(|v| v.base_model == "new-family" && v.model == "new-family-high")); +} + +#[test] +fn refresh_rejects_late_catalog_and_replaced_credential_results() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let key = service + .save_key(ModelKey::new(ModelType::CustomApi)) + .unwrap(); + let first = discovery(&key); + service + .refresh_model_catalog(&key.id, first.clone()) + .unwrap(); + assert!(service.refresh_model_catalog(&key.id, first).is_err()); + let snapshot = service.get_key_by_id(&key.id).unwrap(); + let mut replaced = snapshot.clone(); + replaced.api_key = Some("new-fixture-key".into()); + service.save_key(replaced).unwrap(); + assert!(service + .refresh_model_catalog(&key.id, discovery(&snapshot)) + .is_err()); + let stored = service.get_key_by_id(&key.id).unwrap(); + assert_eq!(stored.model_catalog_generation, 1); + assert_eq!(stored.api_key.as_deref(), Some("new-fixture-key")); +} + +#[test] +fn empty_discovery_leaves_stored_catalog_unchanged_and_retry_succeeds() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let key = service + .save_key(ModelKey::new(ModelType::CustomApi)) + .unwrap(); + let before = std::fs::read(service.get_storage_file()).unwrap(); + let mut empty = discovery(&key); + empty.available_models.clear(); + assert!(service.refresh_model_catalog(&key.id, empty).is_err()); + assert_eq!(before, std::fs::read(service.get_storage_file()).unwrap()); + service + .refresh_model_catalog(&key.id, discovery(&key)) + .unwrap(); + assert_eq!( + service + .get_key_by_id(&key.id) + .unwrap() + .model_catalog_generation, + 1 + ); +} + +#[test] +fn selectable_catalog_inherits_metadata_base_but_never_enables_empty_accounts() { + let mut key = ModelKey::new(ModelType::CustomApi); + key.available_models = vec!["base".into(), "manual".into()]; + // An opaque ID proves this uses metadata, not suffix parsing. + key.model_variants = vec![variant("opaque-variant", "base", None)]; + assert!(KeyInfo::from(key.clone()).selectable_model_ids().is_empty()); + key.enabled_models = vec!["base".into(), "manual".into()]; + assert_eq!( + KeyInfo::from(key.clone()).selectable_model_ids(), + vec!["base", "manual", "opaque-variant"] + ); + key.enabled = false; + assert!(KeyInfo::from(key).selectable_model_ids().is_empty()); +} + +#[test] +fn codex_product_completion_and_cursor_projection_do_not_enable_disabled_models() { + let mut key = ModelKey::new(ModelType::Codex); + key.auth_method = AuthMethod::Oauth; + key.session_token = Some("fixture-token".into()); + key.normalize_model_catalog(); + assert!(!key.available_models.is_empty()); + assert!(KeyInfo::from(key).selectable_model_ids().is_empty()); + let mut cursor = ModelKey::new(ModelType::CursorCli); + cursor.session_token = Some("fixture-token".into()); + let info = crate::commands::key_info_from_entry(cursor).unwrap(); + assert!(info.available_models.contains(&"composer-2".into())); + assert!(info.enabled_models.is_empty()); + assert!(info.selectable_model_ids().is_empty()); +} + +#[test] +fn live_variant_context_survives_without_parallel_context_map_and_provider_default_refreshes() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let mut key = ModelKey::new(ModelType::CustomApi); + key.discovered_default_variants = vec![default("model-old", "model")]; + let key = service.save_key(key).unwrap(); + let mut refresh = discovery(&key); + refresh.model_context_lengths.clear(); + refresh + .model_variants + .as_mut() + .unwrap() + .push(variant("model", "model", Some(12345))); + let stored = service.refresh_model_catalog(&key.id, refresh).unwrap(); + assert!(stored + .model_variants + .iter() + .any(|v| v.model == "model" && v.context_window == Some(12345))); + assert_eq!( + KeyInfo::from(stored.clone()).default_variants[0].model, + "model-high" + ); + assert_eq!( + crate::commands::FullKeyResponse::from(stored).default_variants[0].model, + "model-high" + ); +} + +#[test] +fn catalog_writers_invalidate_old_refresh_but_preference_edits_do_not() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let snapshot = service + .save_key(ModelKey::new(ModelType::CustomApi)) + .unwrap(); + service + .update_key_health( + &snapshot.id, + HealthStatus::Valid, + None, + Some(vec!["validated".into()]), + None, + None, + None, + ) + .unwrap(); + assert!(service + .refresh_model_catalog(&snapshot.id, discovery(&snapshot)) + .is_err()); + let snapshot = service.get_key_by_id(&snapshot.id).unwrap(); + let mut edited = snapshot.clone(); + edited.available_models.push("user-configured-model".into()); + service.save_key(edited).unwrap(); + assert!(service + .refresh_model_catalog(&snapshot.id, discovery(&snapshot)) + .is_err()); + let snapshot = service.get_key_by_id(&snapshot.id).unwrap(); + let mut edited = snapshot.clone(); + edited.enabled_models.clear(); + edited.default_variants.push(default("model-user", "model")); + service.save_key(edited).unwrap(); + assert!(service + .refresh_model_catalog(&snapshot.id, discovery(&snapshot)) + .is_ok()); +} + +#[test] +fn stale_account_edit_cannot_erase_a_completed_discovery() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let snapshot = service + .save_key(ModelKey::new(ModelType::CustomApi)) + .unwrap(); + service + .refresh_model_catalog(&snapshot.id, discovery(&snapshot)) + .unwrap(); + let mut stale = snapshot.clone(); + stale.name = Some("rename from stale snapshot".into()); + assert!(service.save_key(stale).is_err()); + assert!(service + .get_key_by_id(&snapshot.id) + .unwrap() + .available_models + .contains(&"model".into())); +} + +#[test] +fn family_delta_keeps_other_provider_defaults_discovered_across_refresh() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let key = service + .save_key(ModelKey::new(ModelType::CustomApi)) + .unwrap(); + let key = service + .refresh_model_catalog(&key.id, discovery(&key)) + .unwrap(); + let saved = service + .update_default_variant_overrides( + &key.id, + &ModelType::CustomApi, + vec![default("model", "model")], + ) + .unwrap(); + assert_eq!(saved.default_variants, vec![default("model", "model")]); + let mut refreshed = discovery(&saved); + refreshed.default_variants = Some(vec![ + default("model-high", "model"), + default("new-family", "new-family"), + ]); + let saved = service.refresh_model_catalog(&key.id, refreshed).unwrap(); + let reloaded = KeyService::new(Some(dir.path().to_path_buf())) + .get_key_by_id(&key.id) + .unwrap(); + assert_eq!(reloaded.default_variants, vec![default("model", "model")]); + let info = KeyInfo::from(saved); + assert!(info + .default_variants + .iter() + .any(|v| v.base_model == "model" && v.model == "model")); + assert!(info + .default_variants + .iter() + .any(|v| v.base_model == "new-family" && v.model == "new-family")); +} + +#[test] +fn concurrent_family_deltas_merge_at_persistence_boundary() { + let dir = tempfile::tempdir().unwrap(); + let service = std::sync::Arc::new(KeyService::new(Some(dir.path().to_path_buf()))); + let key = service + .save_key(ModelKey::new(ModelType::CustomApi)) + .unwrap(); + let key = service + .refresh_model_catalog(&key.id, discovery(&key)) + .unwrap(); + let mut threads = Vec::new(); + for family in ["model", "new-family"] { + let service = service.clone(); + let id = key.id.clone(); + threads.push(std::thread::spawn(move || { + service + .update_default_variant_overrides( + &id, + &ModelType::CustomApi, + vec![default(family, family)], + ) + .unwrap() + })); + } + for thread in threads { + thread.join().unwrap(); + } + let stored = service.get_key_by_id(&key.id).unwrap(); + assert_eq!(stored.default_variants.len(), 2); + assert!(stored.default_variants.contains(&default("model", "model"))); + assert!(stored + .default_variants + .contains(&default("new-family", "new-family"))); +} + +#[test] +fn changed_provider_route_clears_discovered_defaults_but_preserves_explicit_choices() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let key = service + .save_key(ModelKey::new(ModelType::CustomApi)) + .unwrap(); + let key = service + .refresh_model_catalog(&key.id, discovery(&key)) + .unwrap(); + let mut key = service + .update_default_variant_overrides( + &key.id, + &ModelType::CustomApi, + vec![default("model", "model")], + ) + .unwrap(); + key.base_url = Some("https://new-route.example/v1".into()); + let saved = service.save_key(key).unwrap(); + let reloaded = service.get_key_by_id(&saved.id).unwrap(); + assert!(reloaded.discovered_default_variants.is_empty()); + assert_eq!(reloaded.default_variants, vec![default("model", "model")]); + assert!(!KeyInfo::from(reloaded) + .default_variants + .iter() + .any(|v| v.base_model == "new-family")); +} + +#[test] +fn invalid_family_delta_does_not_mutate_saved_preferences() { + let dir = tempfile::tempdir().unwrap(); + let service = KeyService::new(Some(dir.path().to_path_buf())); + let key = service + .save_key(ModelKey::new(ModelType::CustomApi)) + .unwrap(); + let key = service + .refresh_model_catalog(&key.id, discovery(&key)) + .unwrap(); + assert!(service + .update_default_variant_overrides( + &key.id, + &ModelType::CustomApi, + vec![default("", "model")] + ) + .is_err()); + assert!(service + .get_key_by_id(&key.id) + .unwrap() + .default_variants + .is_empty()); + assert!(service + .update_default_variant_overrides( + &key.id, + &ModelType::Codex, + vec![default("model", "model")] + ) + .is_err()); +} diff --git a/src-tauri/crates/key-vault/src/key_store/types.rs b/src-tauri/crates/key-vault/src/key_store/types.rs index 03db3148ac..56572917cd 100644 --- a/src-tauri/crates/key-vault/src/key_store/types.rs +++ b/src-tauri/crates/key-vault/src/key_store/types.rs @@ -428,6 +428,10 @@ pub struct ModelKey { /// Old managed profiles and in-flight results belong to their original generation. #[serde(default)] pub credential_generation: u64, + /// Increments when discovery, validation, or a catalog edit commits; + /// rejects late refreshes and stale whole-account saves. + #[serde(default)] + pub model_catalog_generation: u64, /// Exact local Codex login copied by the import; never exposed in KeyInfo. #[serde(default, skip_serializing_if = "Option::is_none")] pub codex_cli_auth_path: Option, @@ -471,6 +475,10 @@ pub struct ModelKey { /// concrete variant model id stored here (e.g. `claude-4.6-opus-high`). #[serde(default)] pub default_variants: Vec, + /// Discovery defaults are separate from user choices. Legacy defaults stay + /// user-owned because their original intent cannot be recovered safely. + #[serde(default)] + pub discovered_default_variants: Vec, #[serde(default)] pub oauth_refresh_failure_count: u32, #[serde(default, with = "optional_flexible_datetime")] @@ -506,7 +514,7 @@ pub struct ModelAlias { pub icon: Option, } -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] pub struct ModelVariant { pub model: String, pub base_model: String, @@ -524,7 +532,7 @@ pub struct ModelVariant { /// A user-chosen default variant for one base model family. `base_model` is /// the family root (e.g. `claude-4.6-opus`); `model` is the concrete variant /// id the runtime should launch when that family is selected. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] pub struct DefaultVariant { pub base_model: String, pub model: String, @@ -546,6 +554,7 @@ impl ModelKey { account_metadata: HashMap::new(), auth_method: AuthMethod::ApiKey, credential_generation: 0, + model_catalog_generation: 0, codex_cli_auth_path: None, codex_pending_source_token_hash: None, oauth_auto_disabled: false, @@ -563,6 +572,7 @@ impl ModelKey { model_aliases: Vec::new(), model_variants: Vec::new(), default_variants: Vec::new(), + discovered_default_variants: Vec::new(), oauth_refresh_failure_count: 0, last_oauth_refresh_failed_at: None, temporary_unavailable_until: None, @@ -581,6 +591,7 @@ impl ModelKey { && self.api_key == other.api_key && self.session_token == other.session_token && self.base_url == other.base_url + && self.protocol == other.protocol && self.env_vars == other.env_vars } diff --git a/src-tauri/src/agent_sessions/cli/session_runner/helpers.rs b/src-tauri/src/agent_sessions/cli/session_runner/helpers.rs index 06437e3dbb..185b60663c 100644 --- a/src-tauri/src/agent_sessions/cli/session_runner/helpers.rs +++ b/src-tauri/src/agent_sessions/cli/session_runner/helpers.rs @@ -13,6 +13,7 @@ use crate::agent_sessions::event_pipeline::commands::{ use crate::api::websocket_handler; use agent_core::bus::broadcast_event; use agent_core::foundation::streaming::CLI_STREAMING_BUFFER; +pub use agent_core::state::session_identity_lock; type RunningSessionsMap = HashMap>; @@ -29,13 +30,6 @@ type SessionControlLocksMap = HashMap>>; static SESSION_CONTROL_LOCKS: std::sync::LazyLock> = std::sync::LazyLock::new(|| Mutex::new(HashMap::new())); -// Provider identity (runtime/account/native UUID) is immutable for the whole -// runner lifetime. Unlike the short control lock, this guard travels with the -// background task through final native publication; a model picker may stage a -// next-turn choice but cannot retarget the active runner's filesystem binding. -static SESSION_IDENTITY_LOCKS: std::sync::LazyLock> = - std::sync::LazyLock::new(|| Mutex::new(HashMap::new())); - pub async fn session_control_lock(session_id: &str) -> Arc> { let mut locks = SESSION_CONTROL_LOCKS.lock().await; locks.retain(|_, lock| lock.strong_count() > 0); @@ -47,17 +41,6 @@ pub async fn session_control_lock(session_id: &str) -> Arc> { lock } -pub async fn session_identity_lock(session_id: &str) -> Arc> { - let mut locks = SESSION_IDENTITY_LOCKS.lock().await; - locks.retain(|_, lock| lock.strong_count() > 0); - if let Some(lock) = locks.get(session_id).and_then(Weak::upgrade) { - return lock; - } - let lock = Arc::new(Mutex::new(())); - locks.insert(session_id.to_string(), Arc::downgrade(&lock)); - lock -} - /// Persist an ActivityChunk to the database and broadcast it via WebSocket. /// /// Delta chunks (`action_type` contains "delta") are routed through the diff --git a/src-tauri/src/agent_sessions/session_directory/patch.rs b/src-tauri/src/agent_sessions/session_directory/patch.rs index c7612afeec..205a3dc9ac 100644 --- a/src-tauri/src/agent_sessions/session_directory/patch.rs +++ b/src-tauri/src/agent_sessions/session_directory/patch.rs @@ -179,12 +179,14 @@ fn locate_session(session_id: &str) -> SqliteResult> { /// error on the NEXT turn. Failing the patch up-front gives the frontend /// rollback path a useful message instead. /// -/// Matching is deliberately lenient (exact, prefix in either direction, or -/// alias) because palette model ids may carry variant suffixes; an empty -/// `enabled_models` list means "no restriction" (e.g. fallback-populated -/// native accounts). -fn validate_account_model_compat(account_id: &str, model: &str) -> Result<(), String> { - let Some(key) = key_vault::key_store::KEY_SERVICE.get_key_by_id(account_id) else { +/// Uses the same effective inventory and variant enablement as the picker; +/// empty enablement never grants access to the entire discovered inventory. +fn validate_account_model_compat( + account_id: &str, + model: &str, + key: Option, +) -> Result<(), String> { + let Some(key) = key else { return Err(format!( "session_patch: account {account_id} not found in key vault" )); @@ -192,17 +194,16 @@ fn validate_account_model_compat(account_id: &str, model: &str) -> Result<(), St if !key.enabled { return Err(format!("session_patch: account {account_id} is disabled")); } - if key.enabled_models.is_empty() { - return Ok(()); - } - let compatible = key.enabled_models.iter().any(|enabled| { - enabled == model || model.starts_with(enabled.as_str()) || enabled.starts_with(model) - }) || key.model_aliases.iter().any(|alias| alias.alias == model); + let info = key_vault::commands::key_info_from_entry(key)?; + let compatible = info + .selectable_model_ids() + .iter() + .any(|enabled| enabled == model); if !compatible { return Err(format!( "session_patch: model {model} is not enabled for account {account_id} \ (enabled: {:?})", - key.enabled_models + info.enabled_models )); } Ok(()) @@ -246,10 +247,13 @@ fn resolve_atomic_mode_axes( ))) } -/// Apply a patch synchronously. Public for `#[tauri::command]` -/// adapter; tests can also call this directly with an in-memory DB -/// once the connection abstraction allows it. -pub fn apply_session_patch(session_id: &str, patch: &SessionPatch) -> Result<(), String> { +/// Synchronous persistence phase. Keep private so transports cannot bypass +/// identity serialization, runtime invalidation, or notifications. +fn apply_session_patch( + session_id: &str, + patch: &SessionPatch, + account_lookup: &impl Fn(&str) -> Result, String>, +) -> Result<(), String> { // Reject the only structurally-invalid combination upfront so the // frontend gets a useful error instead of a silent no-op. if patch.account_id.is_some() && patch.model.is_none() { @@ -258,9 +262,6 @@ pub fn apply_session_patch(session_id: &str, patch: &SessionPatch) -> Result<(), .to_string(), ); } - if let (Some(account_id), Some(model)) = (patch.account_id.as_deref(), patch.model.as_deref()) { - validate_account_model_compat(account_id, model)?; - } if patch.name.is_none() && patch.model.is_none() && patch.agent_exec_mode.is_none() @@ -276,6 +277,21 @@ pub fn apply_session_patch(session_id: &str, patch: &SessionPatch) -> Result<(), .map_err(|err| format!("session_patch lookup failed: {err}"))? .ok_or_else(|| format!("session_patch: session {session_id} not found"))?; + if let Some(model) = patch.model.as_deref() { + if model.trim().is_empty() { + return Err("session_patch: model cannot be empty".to_string()); + } + // A model-only edit retains its account; validate the resulting pair, + // not just the fields that happened to be present in the wire payload. + let account_id = patch + .account_id + .clone() + .or_else(|| read_current_account(session_id)); + if let Some(account_id) = account_id.as_deref() { + validate_account_model_compat(account_id, model, account_lookup(account_id)?)?; + } + } + if let Some(name) = patch.name.as_deref() { let trimmed = name.trim(); if trimmed.is_empty() { @@ -465,6 +481,31 @@ pub async fn session_patch( state: tauri::State<'_, agent_core::state::AgentAppState>, session_id: String, patch: SessionPatch, +) -> Result<(), String> { + patch_session(state.inner(), session_id, patch).await +} + +/// Shared desktop/mobile mutation boundary. Every identity edit must serialize +/// with native publication, persist before invalidating the cached runtime, and +/// emit the same account-switch notification before acknowledging success. +pub async fn patch_session( + state: &agent_core::state::AgentAppState, + session_id: String, + patch: SessionPatch, +) -> Result<(), String> { + patch_session_with_account_lookup(state, session_id, patch, |account_id| { + key_vault::key_store::KEY_SERVICE.get_key_by_id_checked(account_id) + }) + .await +} + +async fn patch_session_with_account_lookup( + state: &agent_core::state::AgentAppState, + session_id: String, + patch: SessionPatch, + account_lookup: impl Fn(&str) -> Result, String> + + Send + + 'static, ) -> Result<(), String> { let identity_changed = patch.model.is_some() || patch.account_id.is_some(); // Model/account identity participates in provider-native publication. @@ -474,7 +515,7 @@ pub async fn session_patch( // committed for the next turn once the current provider boundary settles. let _identity_guard = if identity_changed { Some( - crate::agent_sessions::cli::session_runner::session_identity_lock(&session_id) + agent_core::state::session_identity_lock(&session_id) .await .lock_owned() .await, @@ -482,131 +523,89 @@ pub async fn session_patch( } else { None }; - let switched_to_project = patch.product_mode.as_deref() == Some("project"); - let renamed = patch - .name - .as_deref() - .map(str::trim) - .filter(|name| !name.is_empty()) - .map(str::to_string); - let switched_account = patch.account_id.clone(); - let switched_model = patch.model.clone(); - let patched_session_id = session_id.clone(); - let prev_account = tokio::task::spawn_blocking(move || { - let prev_account = patch - .account_id - .is_some() - .then(|| read_current_account(&session_id)) - .flatten(); - apply_session_patch(&session_id, &patch).map(|()| prev_account) - }) - .await - .map_err(|err| format!("session_patch task join error: {err}"))??; - if identity_changed { - state.invalidate_session(&patched_session_id).await; - } - if switched_to_project { - // Convert to Project (orgtrack/v1 §7.2): entering the Project - // product mode must invalidate Plan mode's snapshot/restore - // state, otherwise the pending-approval restore path would - // bounce a later turn back to the pre-Plan exec mode. - if let Some(session) = state.get_session(&patched_session_id).await { - let had_slot = session.plan_slot_cache.get(&patched_session_id).is_some(); - let _ = session.pre_plan_mode_cache.take(&patched_session_id); - session.plan_slot_cache.clear(&patched_session_id); - if had_slot { - agent_core::bus::broadcast_event( - "agent:exit_plan_mode", - serde_json::json!({ - "sessionId": &patched_session_id, - "source": "convert_to_project", - "nextMode": agent_core::session::AgentExecMode::Build.as_str(), - }), - ); + // Once admitted, persistence and cache invalidation form one lifecycle. + // A disconnected mobile caller can cancel its response future, but cannot + // leave a completed blocking DB write paired with the previous runtime. + let state = state.clone(); + tokio::spawn(async move { + let _identity_guard = _identity_guard; + let identity_session = if identity_changed { + state.get_session(&session_id).await + } else { + None + }; + let runtime_mutation = match identity_session.as_ref() { + Some(session) => Some(session.begin_identity_mutation().await?), + None => None, + }; + let switched_to_project = patch.product_mode.as_deref() == Some("project"); + let renamed = patch + .name + .as_deref() + .map(str::trim) + .filter(|name| !name.is_empty()) + .map(str::to_string); + let switched_account = patch.account_id.clone(); + let switched_model = patch.model.clone(); + let patched_session_id = session_id.clone(); + let prev_account = tokio::task::spawn_blocking(move || { + let prev_account = patch + .account_id + .is_some() + .then(|| read_current_account(&session_id)) + .flatten(); + apply_session_patch(&session_id, &patch, &account_lookup).map(|()| prev_account) + }) + .await + .map_err(|err| format!("session_patch task join error: {err}"))??; + if let Some(mutation) = runtime_mutation { + mutation.invalidate_runtime().await; + } + if switched_to_project { + // Convert to Project (orgtrack/v1 §7.2): entering the Project + // product mode must invalidate Plan mode's snapshot/restore + // state, otherwise the pending-approval restore path would + // bounce a later turn back to the pre-Plan exec mode. + if let Some(session) = state.get_session(&patched_session_id).await { + let had_slot = session.plan_slot_cache.get(&patched_session_id).is_some(); + let _ = session.pre_plan_mode_cache.take(&patched_session_id); + session.plan_slot_cache.clear(&patched_session_id); + if had_slot { + agent_core::bus::broadcast_event( + "agent:exit_plan_mode", + serde_json::json!({ + "sessionId": &patched_session_id, + "source": "convert_to_project", + "nextMode": agent_core::session::AgentExecMode::Build.as_str(), + }), + ); + } } } - } - if let Some(name) = renamed.as_deref() { - agent_core::lifecycle::emit_session_renamed( - state.app_handle.as_ref(), - &patched_session_id, - name, - ); - } - if let Some(to_account) = switched_account.as_deref() { - if prev_account.as_deref() != Some(to_account) { - agent_core::lifecycle::emit_session_account_switched( + if let Some(name) = renamed.as_deref() { + agent_core::lifecycle::emit_session_renamed( state.app_handle.as_ref(), &patched_session_id, - prev_account.as_deref(), - to_account, - switched_model.as_deref(), + name, ); } - } - Ok(()) + if let Some(to_account) = switched_account.as_deref() { + if prev_account.as_deref() != Some(to_account) { + agent_core::lifecycle::emit_session_account_switched( + state.app_handle.as_ref(), + &patched_session_id, + prev_account.as_deref(), + to_account, + switched_model.as_deref(), + ); + } + } + Ok(()) + }) + .await + .map_err(|err| format!("session_patch lifecycle task join error: {err}"))? } #[cfg(test)] -mod tests { - use super::*; - - #[test] - fn project_derives_build_and_ordinary_modes_never_gain_pm_capability() { - assert_eq!( - resolve_atomic_mode_axes(Some("project"), Some("ask")).unwrap(), - Some(("project".to_string(), "build".to_string())) - ); - assert_eq!( - resolve_atomic_mode_axes(Some("build"), Some("build")).unwrap(), - Some(("build".to_string(), "build".to_string())) - ); - assert_eq!( - resolve_atomic_mode_axes(Some("plan"), Some("plan")).unwrap(), - Some(("plan".to_string(), "plan".to_string())) - ); - assert!(resolve_atomic_mode_axes(Some("project-ish"), Some("build")).is_err()); - assert!(resolve_atomic_mode_axes(Some("project"), Some("unrestricted")).is_err()); - } - - #[test] - fn double_option_distinguishes_absent_null_value() { - // Field absent → None → "leave alone" - let absent: SessionPatch = serde_json::from_str("{}").unwrap(); - assert!(absent.draft_text.is_none()); - assert!(absent.reply_target_event_id.is_none()); - - // Field is JSON null → Some(None) → "clear" - let nulled: SessionPatch = - serde_json::from_str(r#"{"draftText": null, "replyTargetEventId": null}"#).unwrap(); - assert_eq!(nulled.draft_text, Some(None)); - assert_eq!(nulled.reply_target_event_id, Some(None)); - - // Field is a string → Some(Some(_)) → "set" - let set: SessionPatch = - serde_json::from_str(r#"{"draftText": "hello", "replyTargetEventId": "evt_42"}"#) - .unwrap(); - assert_eq!(set.draft_text, Some(Some("hello".to_string()))); - assert_eq!(set.reply_target_event_id, Some(Some("evt_42".to_string()))); - } - - #[test] - fn empty_patch_is_rejected() { - let patch = SessionPatch::default(); - let err = apply_session_patch("nonexistent", &patch).unwrap_err(); - assert!(err.contains("at least one field"), "got: {err}"); - } - - #[test] - fn account_without_model_is_rejected() { - let patch = SessionPatch { - account_id: Some("acc_1".to_string()), - ..SessionPatch::default() - }; - let err = apply_session_patch("nonexistent", &patch).unwrap_err(); - assert!( - err.contains("account_id provided without model"), - "got: {err}" - ); - } -} +#[path = "patch_tests.rs"] +mod tests; diff --git a/src-tauri/src/agent_sessions/session_directory/patch_tests.rs b/src-tauri/src/agent_sessions/session_directory/patch_tests.rs new file mode 100644 index 0000000000..3d630a574b --- /dev/null +++ b/src-tauri/src/agent_sessions/session_directory/patch_tests.rs @@ -0,0 +1,370 @@ +use super::*; +use agent_core::definitions::{resolved::ResolvedAgent, AgentDefinition}; +use agent_core::state::{AgentAppState, AgentSession, SessionRuntime}; +use key_vault::key_store::{ModelKey, ModelType}; +use std::sync::Arc; + +const FIRST_MODEL: &str = "e2e-fake-provider-a"; +const SECOND_MODEL: &str = "e2e-fake-provider-b"; + +fn lookup(account: &str) -> Result, String> { + if account == "missing" { + return Ok(None); + } + let mut key = ModelKey::new(ModelType::OpenaiApi); + key.id = account.into(); + key.available_models = vec![FIRST_MODEL.into(), SECOND_MODEL.into()]; + key.enabled_models = key.available_models.clone(); + if account == "disabled" { + key.enabled_models.clear(); + } + Ok(Some(key)) +} + +async fn fixture(state: &AgentAppState, sid: &str, root: &std::path::Path) -> Arc { + let conn = get_connection().unwrap(); + conn.execute( + "INSERT INTO agent_sessions (session_id,name,session_type,status,model,account_id,workspace_path,created_at,updated_at) VALUES (?1,'Model fixture','agent','idle',?2,'account-a',?3,datetime('now'),datetime('now'))", + params![sid, FIRST_MODEL, root.to_string_lossy()], + ).unwrap(); + let definition = AgentDefinition { + selected_model_id: Some(FIRST_MODEL.into()), + ..Default::default() + }; + let resolved = ResolvedAgent::resolve(&definition, None, &Default::default()).unwrap(); + let handle = state + .register_session(AgentSession::new(sid.into(), definition)) + .await; + handle + .set_runtime(Arc::new(SessionRuntime { + provider: Arc::new(agent_core::providers::e2e_fake::E2eFakeProvider), + tool_registry: Arc::new(agent_core::tools::registry::ToolRegistry::new()), + policy: Arc::new(agent_core::tools::policy::ResolvedToolPolicy::permissive()), + model: FIRST_MODEL.into(), + account_id: Some("account-a".into()), + native_harness_type: None, + workspace_state: Arc::new(parking_lot::RwLock::new( + agent_core::session::workspace::SessionWorkspace::new(root.into()), + )), + mcp_auto_approved: vec![], + resolved, + integrations_snapshot: Default::default(), + overrides: Default::default(), + agent_soul: None, + sovereign_prompt: false, + policy_context_activator: None, + agent_org_context: None, + agent_org_current_member_id: None, + agent_definition_id: None, + })) + .await + .unwrap(); + handle +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn identity_patch_commits_pair_invalidates_runtime_and_next_mobile_send_uses_it() { + crate::test_utils::install_crypto_provider_for_tests(); + let sandbox = crate::test_utils::test_env::sandbox(); + let state = AgentAppState::new(); + for (sid, model) in [ + ("model-patch-same", FIRST_MODEL), + ("model-patch-both", SECOND_MODEL), + ] { + let handle = fixture(&state, sid, sandbox.path()).await; + let active_runtime = handle.get_runtime().await.unwrap(); + patch_session_with_account_lookup( + &state, + sid.into(), + SessionPatch { + model: Some(model.into()), + account_id: Some("account-b".into()), + ..Default::default() + }, + lookup, + ) + .await + .unwrap(); + let stored = session_persistence::get_session(sid).unwrap().unwrap(); + assert_eq!(stored.model.as_deref(), Some(model)); + assert_eq!(stored.account_id.as_deref(), Some("account-b")); + assert!(handle.get_runtime().await.is_none()); + // A running reply's captured provider remains unchanged. + assert_eq!(active_runtime.account_id.as_deref(), Some("account-a")); + agent_core::state::commands::session::message::send_message_impl_for_mobile_remote( + &state, + sid.into(), + "continue".into(), + Some(format!("turn-{sid}")), + Some(model.into()), + None, + ) + .await + .unwrap(); + let next = handle.get_runtime().await.unwrap(); + assert_eq!(next.account_id.as_deref(), Some("account-b")); + assert_eq!(next.model, model); + let stored = session_persistence::get_session(sid).unwrap().unwrap(); + assert_eq!(stored.account_id.as_deref(), Some("account-b")); + // Let the local fake-provider turn settle before dropping its sandbox. + tokio::time::timeout(std::time::Duration::from_secs(5), async { + while handle.scheduler.is_processing() || handle.scheduler.pending_count() > 0 { + tokio::task::yield_now().await; + } + }) + .await + .unwrap(); + state.remove_session(sid).await; + } + state + .begin_shutdown(std::time::Duration::from_secs(1)) + .await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn failed_patch_preserves_persisted_pair_and_runtime_and_can_retry() { + let sandbox = crate::test_utils::test_env::sandbox(); + let state = AgentAppState::new(); + let sid = "model-patch-invalid"; + let handle = fixture(&state, sid, sandbox.path()).await; + let runtime = handle.get_runtime().await.unwrap(); + for account in ["missing", "disabled"] { + assert!(patch_session_with_account_lookup( + &state, + sid.into(), + SessionPatch { + model: Some(SECOND_MODEL.into()), + account_id: Some(account.into()), + ..Default::default() + }, + lookup + ) + .await + .is_err()); + assert!(Arc::ptr_eq(&runtime, &handle.get_runtime().await.unwrap())); + let stored = session_persistence::get_session(sid).unwrap().unwrap(); + assert_eq!(stored.model.as_deref(), Some(FIRST_MODEL)); + assert_eq!(stored.account_id.as_deref(), Some("account-a")); + } + assert!(patch_session_with_account_lookup( + &state, + sid.into(), + SessionPatch { + model: Some("unlisted-model".into()), + ..Default::default() + }, + lookup + ) + .await + .is_err()); + patch_session_with_account_lookup( + &state, + sid.into(), + SessionPatch { + model: Some(SECOND_MODEL.into()), + account_id: Some("account-b".into()), + ..Default::default() + }, + lookup, + ) + .await + .unwrap(); + assert!(handle.get_runtime().await.is_none()); + state + .begin_shutdown(std::time::Duration::from_secs(1)) + .await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn identity_patch_waits_for_prior_admission_then_wins_persistence() { + let sandbox = crate::test_utils::test_env::sandbox(); + let state = Arc::new(AgentAppState::new()); + let sid = "model-patch-admission"; + let handle = fixture(&state, sid, sandbox.path()).await; + let guard = agent_core::state::session_identity_lock(sid) + .await + .lock_owned() + .await; + let patch_state = state.clone(); + let patch = tokio::spawn(async move { + patch_session_with_account_lookup( + &patch_state, + sid.into(), + SessionPatch { + model: Some(SECOND_MODEL.into()), + account_id: Some("account-b".into()), + ..Default::default() + }, + lookup, + ) + .await + }); + tokio::task::yield_now().await; + assert!(!patch.is_finished()); + assert!(handle.get_runtime().await.is_some()); + session_persistence::update_model_and_account(sid, FIRST_MODEL, Some("account-a")).unwrap(); + drop(guard); + patch.await.unwrap().unwrap(); + let stored = session_persistence::get_session(sid).unwrap().unwrap(); + assert_eq!(stored.model.as_deref(), Some(SECOND_MODEL)); + assert_eq!(stored.account_id.as_deref(), Some("account-b")); + assert!(handle.get_runtime().await.is_none()); + state + .begin_shutdown(std::time::Duration::from_secs(1)) + .await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn persistence_failure_keeps_old_identity_and_runtime_until_retry_succeeds() { + let sandbox = crate::test_utils::test_env::sandbox(); + let state = AgentAppState::new(); + let sid = "model-patch-storage-failure"; + let handle = fixture(&state, sid, sandbox.path()).await; + let runtime = handle.get_runtime().await.unwrap(); + get_connection().unwrap().execute_batch( + "CREATE TRIGGER reject_identity_patch BEFORE UPDATE OF model, account_id ON agent_sessions + WHEN NEW.session_id = 'model-patch-storage-failure' + BEGIN SELECT RAISE(ABORT, 'simulated identity storage failure'); END;", + ).unwrap(); + let patch = SessionPatch { + model: Some(SECOND_MODEL.into()), + account_id: Some("account-b".into()), + ..Default::default() + }; + let error = patch_session_with_account_lookup(&state, sid.into(), patch.clone(), lookup) + .await + .unwrap_err(); + assert!(error.contains("simulated identity storage failure")); + assert!(Arc::ptr_eq(&runtime, &handle.get_runtime().await.unwrap())); + let stored = session_persistence::get_session(sid).unwrap().unwrap(); + assert_eq!(stored.model.as_deref(), Some(FIRST_MODEL)); + assert_eq!(stored.account_id.as_deref(), Some("account-a")); + get_connection() + .unwrap() + .execute_batch("DROP TRIGGER reject_identity_patch;") + .unwrap(); + patch_session_with_account_lookup(&state, sid.into(), patch, lookup) + .await + .unwrap(); + assert!(handle.get_runtime().await.is_none()); + assert_eq!( + session_persistence::get_session(sid) + .unwrap() + .unwrap() + .account_id + .as_deref(), + Some("account-b") + ); + state + .begin_shutdown(std::time::Duration::from_secs(1)) + .await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn cancelled_mobile_response_still_finishes_admitted_identity_mutation() { + let sandbox = crate::test_utils::test_env::sandbox(); + let state = Arc::new(AgentAppState::new()); + let sid = "model-patch-disconnected"; + let handle = fixture(&state, sid, sandbox.path()).await; + let (started, admitted) = tokio::sync::oneshot::channel(); + let started = std::sync::Mutex::new(Some(started)); + let (release, resume) = std::sync::mpsc::channel(); + let resume = std::sync::Mutex::new(resume); + let task_state = state.clone(); + let caller = tokio::spawn(async move { + patch_session_with_account_lookup( + &task_state, + sid.into(), + SessionPatch { + model: Some(SECOND_MODEL.into()), + account_id: Some("account-b".into()), + ..Default::default() + }, + move |account| { + started.lock().unwrap().take().unwrap().send(()).unwrap(); + resume.lock().unwrap().recv().unwrap(); + lookup(account) + }, + ) + .await + }); + admitted.await.unwrap(); + caller.abort(); + assert!(caller.await.unwrap_err().is_cancelled()); + release.send(()).unwrap(); + // The admitted lifecycle still owns the identity lock until the DB write, + // invalidation and notifications have all completed after disconnect. + let _settled = tokio::time::timeout(std::time::Duration::from_secs(5), async { + agent_core::state::session_identity_lock(sid) + .await + .lock_owned() + .await + }) + .await + .unwrap(); + let stored = session_persistence::get_session(sid).unwrap().unwrap(); + assert_eq!(stored.model.as_deref(), Some(SECOND_MODEL)); + assert_eq!(stored.account_id.as_deref(), Some("account-b")); + assert!(handle.get_runtime().await.is_none()); + state + .begin_shutdown(std::time::Duration::from_secs(1)) + .await; +} + +#[test] +fn project_derives_build_and_ordinary_modes_never_gain_pm_capability() { + assert_eq!( + resolve_atomic_mode_axes(Some("project"), Some("ask")).unwrap(), + Some(("project".to_string(), "build".to_string())) + ); + assert_eq!( + resolve_atomic_mode_axes(Some("build"), Some("build")).unwrap(), + Some(("build".to_string(), "build".to_string())) + ); + assert_eq!( + resolve_atomic_mode_axes(Some("plan"), Some("plan")).unwrap(), + Some(("plan".to_string(), "plan".to_string())) + ); + assert!(resolve_atomic_mode_axes(Some("project-ish"), Some("build")).is_err()); + assert!(resolve_atomic_mode_axes(Some("project"), Some("unrestricted")).is_err()); +} + +#[test] +fn double_option_distinguishes_absent_null_value() { + // Field absent → None → "leave alone" + let absent: SessionPatch = serde_json::from_str("{}").unwrap(); + assert!(absent.draft_text.is_none()); + assert!(absent.reply_target_event_id.is_none()); + + // Field is JSON null → Some(None) → "clear" + let nulled: SessionPatch = + serde_json::from_str(r#"{"draftText": null, "replyTargetEventId": null}"#).unwrap(); + assert_eq!(nulled.draft_text, Some(None)); + assert_eq!(nulled.reply_target_event_id, Some(None)); + + // Field is a string → Some(Some(_)) → "set" + let set: SessionPatch = + serde_json::from_str(r#"{"draftText": "hello", "replyTargetEventId": "evt_42"}"#).unwrap(); + assert_eq!(set.draft_text, Some(Some("hello".to_string()))); + assert_eq!(set.reply_target_event_id, Some(Some("evt_42".to_string()))); +} + +#[test] +fn empty_patch_is_rejected() { + let patch = SessionPatch::default(); + let err = apply_session_patch("nonexistent", &patch, &|_| Ok(None)).unwrap_err(); + assert!(err.contains("at least one field"), "got: {err}"); +} + +#[test] +fn account_without_model_is_rejected() { + let patch = SessionPatch { + account_id: Some("acc_1".to_string()), + ..SessionPatch::default() + }; + let err = apply_session_patch("nonexistent", &patch, &|_| Ok(None)).unwrap_err(); + assert!( + err.contains("account_id provided without model"), + "got: {err}" + ); +} diff --git a/src-tauri/src/api/mobile_bridge/adapters/model.rs b/src-tauri/src/api/mobile_bridge/adapters/model.rs index 5d550d86da..404bd99a19 100644 --- a/src-tauri/src/api/mobile_bridge/adapters/model.rs +++ b/src-tauri/src/api/mobile_bridge/adapters/model.rs @@ -2,11 +2,14 @@ use serde_json::{json, Value}; use std::collections::HashSet; +use tauri::Manager; use crate::agent_sessions::cli::persistence as cli_persistence; -use crate::agent_sessions::session_directory::patch::{apply_session_patch, SessionPatch}; +use crate::agent_sessions::session_directory::patch::{patch_session, SessionPatch}; use agent_core::session::persistence as session_persistence; -use key_vault::key_store::{HealthStatus, ModelKey, ModelType, KEY_SERVICE}; +use key_vault::commands::registry::is_cli_provider_compatible; +use key_vault::commands::{key_info_from_entry, KeyInfo}; +use key_vault::key_store::KEY_SERVICE; use super::session::mobile_session_execution; use super::session::MobileSessionExecution; @@ -86,105 +89,78 @@ fn load_session_model_state(session_id: &str) -> Result bool { - if !entry.enabled { - return false; - } - matches!( - entry.health_status, - HealthStatus::Valid | HealthStatus::Degraded | HealthStatus::Unknown - ) -} - -fn has_api_key(entry: &ModelKey) -> bool { - match entry.model_type { - ModelType::CursorCli => entry.api_key.as_deref().is_some_and(|api_key| { - let trimmed = api_key.trim(); - trimmed.len() >= 20 && (trimmed.starts_with("key_") || trimmed.starts_with("crsr_")) - }), - _ => entry - .api_key - .as_deref() - .is_some_and(|secret| !secret.trim().is_empty()), - } -} - -fn has_session_token(entry: &ModelKey) -> bool { - entry - .session_token - .as_deref() - .is_some_and(|token| !token.trim().is_empty()) +fn key_is_usable(entry: &KeyInfo) -> bool { + entry.enabled + && matches!( + entry.health_status.as_str(), + "valid" | "degraded" | "unknown" + ) } -fn supports_rust_agents(entry: &ModelKey) -> bool { - let has_api_key = has_api_key(entry); - let has_session_token = has_session_token(entry); - let can_use_native_harness = - matches!(entry.model_type, ModelType::CursorCli) && has_session_token; - if can_use_native_harness { - return true; - } - let has_usable_key_material = has_api_key || has_session_token; - match entry.model_type { - ModelType::CursorCli | ModelType::OrgiiOrchestrator => false, - ModelType::ClaudeCode - | ModelType::Codex - | ModelType::Copilot - | ModelType::Kiro - | ModelType::KimiCli - | ModelType::OpenCode => has_usable_key_material, - _ => has_api_key, - } -} - -fn models_for_key(entry: &ModelKey) -> Vec { - if !entry.enabled_models.is_empty() { - return entry.enabled_models.clone(); - } - if !entry.available_models.is_empty() { - return entry.available_models.clone(); - } - Vec::new() -} - -fn account_label_for(entry: &ModelKey) -> String { +fn account_label_for(entry: &KeyInfo) -> String { entry .name .clone() .filter(|name| !name.trim().is_empty()) - .unwrap_or_else(|| entry.model_type.as_str().to_string()) + .unwrap_or_else(|| entry.agent_type.clone()) } -fn collect_model_options_for_session( - session_id: &str, +fn supports_session( + entry: &KeyInfo, + execution: MobileSessionExecution, state: &MobileSessionModelState, -) -> Result, RpcError> { - let keys = KEY_SERVICE - .list_keys_checked() - .map_err(|err| RpcError::new(RpcErrorCode::InvalidRequest, err))?; +) -> bool { + match execution { + MobileSessionExecution::ManagedCli => { + state.cli_agent_type.as_deref().is_some_and(|agent| { + if agent == entry.agent_type { + entry.can_launch_cli + } else { + entry.has_api_key && is_cli_provider_compatible(agent, &entry.agent_type) + } + }) + } + MobileSessionExecution::NativeAgent => entry.supports_rust_agents, + MobileSessionExecution::ImportedHistory => false, + } +} +/// A bounded projection of the same account inventory/enablement used on desktop. +/// The current value is retained for presentation, even when no longer selectable. +fn collect_model_options( + keys: &[KeyInfo], + execution: MobileSessionExecution, + state: &MobileSessionModelState, +) -> Vec { let mut options = Vec::new(); let mut seen = HashSet::new(); - let execution = mobile_session_execution(session_id); - for entry in keys.iter().filter(|entry| key_is_usable(entry)) { - let include = match execution { - MobileSessionExecution::ManagedCli => state - .cli_agent_type - .as_deref() - .and_then(ModelType::from_str) - .is_some_and(|agent| agent == entry.model_type), - MobileSessionExecution::NativeAgent => supports_rust_agents(entry), - MobileSessionExecution::ImportedHistory => false, - }; - if !include { - continue; - } + // Reserve the current selection before filling the bounded catalog so it + // cannot create a 257th row that the mobile wire validator would reject. + if let Some(current_model) = state.model.as_deref() { + let current_account = state.account_id.as_deref().unwrap_or(""); + seen.insert((current_account.to_string(), current_model.to_string())); + options.push(MobileModelOption { + id: current_model.to_string(), + account_id: current_account.to_string(), + account_label: keys + .iter() + .find(|entry| entry.id == current_account) + .map(account_label_for) + .unwrap_or_else(|| "Current".to_string()), + }); + } + for entry in keys + .iter() + .filter(|entry| key_is_usable(entry) && supports_session(entry, execution, state)) + { let account_label = account_label_for(entry); - for model_id in models_for_key(entry) { - let dedupe_key = format!("{}::{}", entry.id, model_id); - if !seen.insert(dedupe_key) { + for model_id in entry.selectable_model_ids() { + if options.len() >= MAX_MOBILE_MODEL_OPTIONS { + return options; + } + if !seen.insert((entry.id.clone(), model_id.clone())) { continue; } options.push(MobileModelOption { @@ -192,42 +168,34 @@ fn collect_model_options_for_session( account_id: entry.id.clone(), account_label: account_label.clone(), }); - if options.len() >= MAX_MOBILE_MODEL_OPTIONS { - break; - } - } - if options.len() >= MAX_MOBILE_MODEL_OPTIONS { - break; - } - } - - if let Some(current_model) = state.model.as_deref() { - let current_account = state.account_id.as_deref().unwrap_or(""); - let dedupe_key = format!("{current_account}::{current_model}"); - if seen.insert(dedupe_key) { - let account_label = keys - .iter() - .find(|entry| entry.id == current_account) - .map(account_label_for) - .unwrap_or_else(|| "Current".to_string()); - options.insert( - 0, - MobileModelOption { - id: current_model.to_string(), - account_id: current_account.to_string(), - account_label, - }, - ); } } + options +} - Ok(options) +fn collect_model_options_for_session( + session_id: &str, + state: &MobileSessionModelState, +) -> Result, RpcError> { + let keys = KEY_SERVICE + .list_keys_checked() + .and_then(|keys| { + keys.into_iter() + .map(key_info_from_entry) + .collect::, _>>() + }) + .map_err(|err| RpcError::new(RpcErrorCode::InvalidRequest, err))?; + Ok(collect_model_options( + &keys, + mobile_session_execution(session_id), + state, + )) } /// Return the session's current model configuration for the mobile picker. pub async fn session_config(params: &Value) -> Result { let session_id = parse_session_id(params)?; - let state = load_session_model_state(&session_id)?; + let state = load_session_model_state_async(session_id.clone()).await?; Ok(json!({ "sessionId": session_id, @@ -261,7 +229,7 @@ pub async fn session_patch(params: &Value) -> Result { .filter(|value| !value.is_empty()) .map(str::to_string); - let state = load_session_model_state(&session_id)?; + let state = load_session_model_state_async(session_id.clone()).await?; if !state.model_editable { return Err(RpcError::new( RpcErrorCode::InvalidRequest, @@ -269,14 +237,19 @@ pub async fn session_patch(params: &Value) -> Result { )); } - apply_session_patch( - &session_id, - &SessionPatch { + let handle = crate::api::get_app_handle() + .ok_or_else(|| RpcError::new(RpcErrorCode::InvalidRequest, "desktop agent not ready"))?; + let app_state = handle.state::(); + patch_session( + app_state.inner(), + session_id.clone(), + SessionPatch { model: Some(model.clone()), account_id, ..Default::default() }, ) + .await .map_err(|err| RpcError::new(RpcErrorCode::InvalidRequest, err))?; Ok(json!({ @@ -288,12 +261,15 @@ pub async fn session_patch(params: &Value) -> Result { /// List selectable models for a session, sourced from the desktop KeyVault. pub async fn models_list(params: &Value) -> Result { let session_id = parse_session_id(params)?; - let state = load_session_model_state(&session_id)?; - if !state.model_editable { - return Ok(json!({ "models": [] })); - } - - let options = collect_model_options_for_session(&session_id, &state)?; + let options = tokio::task::spawn_blocking(move || { + let state = load_session_model_state(&session_id)?; + if !state.model_editable { + return Ok(Vec::new()); + } + collect_model_options_for_session(&session_id, &state) + }) + .await + .map_err(|err| RpcError::new(RpcErrorCode::InvalidRequest, err.to_string()))??; let models = options .into_iter() .map(|option| { @@ -308,31 +284,14 @@ pub async fn models_list(params: &Value) -> Result { Ok(json!({ "models": models })) } -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn parse_session_id_rejects_empty() { - let err = parse_session_id(&json!({ "sessionId": " " })).unwrap_err(); - assert_eq!(err.code, RpcErrorCode::InvalidParams); - } - - #[test] - fn supports_rust_agents_allows_byok_api_keys() { - let mut entry = ModelKey::new(ModelType::AnthropicApi); - entry.api_key = Some("sk-ant-test-key-1234567890".to_string()); - assert!(supports_rust_agents(&entry)); - } - - #[test] - fn models_for_key_prefers_enabled_models() { - let mut entry = ModelKey::new(ModelType::AnthropicApi); - entry.enabled_models = vec!["claude-sonnet-4-5".to_string()]; - entry.available_models = vec!["claude-opus-4-5".to_string()]; - assert_eq!( - models_for_key(&entry), - vec!["claude-sonnet-4-5".to_string()] - ); - } +async fn load_session_model_state_async( + session_id: String, +) -> Result { + tokio::task::spawn_blocking(move || load_session_model_state(&session_id)) + .await + .map_err(|err| RpcError::new(RpcErrorCode::InvalidRequest, err.to_string()))? } + +#[cfg(test)] +#[path = "model_tests.rs"] +mod tests; diff --git a/src-tauri/src/api/mobile_bridge/adapters/model_tests.rs b/src-tauri/src/api/mobile_bridge/adapters/model_tests.rs new file mode 100644 index 0000000000..a3a924b859 --- /dev/null +++ b/src-tauri/src/api/mobile_bridge/adapters/model_tests.rs @@ -0,0 +1,118 @@ +use super::*; +use key_vault::key_store::{ModelKey, ModelType}; + +fn state(model: Option<&str>, account_id: Option<&str>) -> MobileSessionModelState { + MobileSessionModelState { + model: model.map(str::to_string), + account_id: account_id.map(str::to_string), + key_source: "own_key".into(), + cli_agent_type: None, + model_editable: true, + } +} + +fn account(provider: ModelType, id: &str, models: &[&str]) -> KeyInfo { + let mut key = ModelKey::new(provider); + key.id = id.into(); + key.api_key = Some("fixture-key".into()); + key.available_models = models.iter().map(|model| model.to_string()).collect(); + key.enabled_models = key.available_models.clone(); + KeyInfo::from(key) +} + +#[test] +fn empty_enabled_inventory_never_becomes_selectable() { + let mut key = account(ModelType::AnthropicApi, "account", &["model-a"]); + key.enabled_models.clear(); + let options = collect_model_options( + &[key], + MobileSessionExecution::NativeAgent, + &state(None, None), + ); + assert!(options.is_empty()); +} + +#[test] +fn disabled_and_unhealthy_accounts_are_not_selectable() { + let mut disabled = account(ModelType::AnthropicApi, "disabled", &["model-a"]); + disabled.enabled = false; + let mut invalid = account(ModelType::AnthropicApi, "invalid", &["model-a"]); + invalid.health_status = "invalid".into(); + assert!(collect_model_options( + &[disabled, invalid], + MobileSessionExecution::NativeAgent, + &state(None, None), + ) + .is_empty()); +} + +#[test] +fn cli_catalog_uses_registry_provider_compatibility_and_launch_capability() { + let mut selection = state(None, None); + selection.cli_agent_type = Some("codex".into()); + let mut native = account(ModelType::Codex, "native", &["gpt-5.5"]); + native.can_launch_cli = false; + let keys = [ + native, + account(ModelType::OpenaiApi, "compatible", &["gpt-5.5"]), + account( + ModelType::AnthropicApi, + "incompatible", + &["claude-sonnet-4-5"], + ), + ]; + let options = collect_model_options(&keys, MobileSessionExecution::ManagedCli, &selection); + assert_eq!(options.len(), 1); + assert_eq!(options[0].account_id, "compatible"); +} + +#[test] +fn current_selection_is_retained_once_inside_wire_limit() { + let models = (0..300) + .map(|i| format!("custom-model-{i}")) + .collect::>(); + let model_refs = models.iter().map(String::as_str).collect::>(); + let keys = [account(ModelType::AnthropicApi, "account", &model_refs)]; + for current_model in ["no-longer-listed", "custom-model-250"] { + let options = collect_model_options( + &keys, + MobileSessionExecution::NativeAgent, + &state(Some(current_model), Some("account")), + ); + assert_eq!(options.len(), MAX_MOBILE_MODEL_OPTIONS); + assert_eq!(options[0].id, current_model); + assert_eq!( + options + .iter() + .filter(|option| option.id == current_model) + .count(), + 1 + ); + } +} + +#[tokio::test] +async fn invalid_and_imported_patch_requests_do_not_reach_mutation() { + let invalid = session_patch(&json!({"sessionId": "sde:fixture", "patch": {"model": " "}})) + .await + .unwrap_err(); + assert_eq!(invalid.code, RpcErrorCode::InvalidParams); + // Imported history is read-only at the adapter before any application state + // or persistence mutation is requested. + let read_only = + session_patch(&json!({"sessionId": "codexapp-fixture", "patch": {"model": "model-b"}})) + .await + .unwrap_err(); + assert_eq!(read_only.code, RpcErrorCode::InvalidRequest); + assert!(read_only.message.contains("cannot be changed")); +} + +#[test] +fn parse_session_id_rejects_empty() { + assert_eq!( + parse_session_id(&json!({"sessionId": " "})) + .unwrap_err() + .code, + RpcErrorCode::InvalidParams + ); +} diff --git a/src/api/services/keyValidation.ts b/src/api/services/keyValidation.ts index 6f0af26f42..98f65abeea 100644 --- a/src/api/services/keyValidation.ts +++ b/src/api/services/keyValidation.ts @@ -350,6 +350,26 @@ export async function updateKeyHealth( }); } +/** Atomically replace discovery metadata without replaying user preferences. */ +export async function refreshKeyModelCatalog( + keyId: string, + catalogRefresh: { + expectedCredentialGeneration: number; + expectedCatalogGeneration: number; + availableModels: string[]; + modelVariants: ModelVariantInfo[] | null; + defaultVariants: DefaultVariantInfo[] | null; + modelContextLengths: ModelContextLengths; + } +): Promise { + return rpc.validation.updateKeyHealth({ + keyId, + // Catalog writes preserve health from the authoritative stored entry. + healthStatus: "unknown", + catalogRefresh, + }); +} + /** Polish a chat draft through the configured local MiniCPM vLLM account. */ export async function promptPolish( text: string, diff --git a/src/api/tauri/rpc/schemas/validationProcedures.ts b/src/api/tauri/rpc/schemas/validationProcedures.ts index b174210491..25f21c5eb6 100644 --- a/src/api/tauri/rpc/schemas/validationProcedures.ts +++ b/src/api/tauri/rpc/schemas/validationProcedures.ts @@ -207,6 +207,15 @@ export const DeleteKeyByIdInput = z.object({ keyId: z.string(), }); +export const ModelCatalogRefreshSchema = z.object({ + expectedCredentialGeneration: z.number().int().nonnegative(), + expectedCatalogGeneration: z.number().int().nonnegative(), + availableModels: z.array(z.string()), + modelVariants: z.array(ModelVariantInfoSchema).nullable(), + defaultVariants: z.array(DefaultVariantInfoSchema).nullable(), + modelContextLengths: ModelContextLengthsSchema, +}); + export const UpdateKeyHealthInput = z.object({ keyId: z.string(), healthStatus: HealthStatusSchema, @@ -215,6 +224,7 @@ export const UpdateKeyHealthInput = z.object({ enabledModels: z.array(z.string()).nullable().optional(), quotaInfo: z.record(z.string(), z.unknown()).nullable().optional(), modelContextLengths: ModelContextLengthsSchema.nullable().optional(), + catalogRefresh: ModelCatalogRefreshSchema.nullable().optional(), }); export const PromptPolishRequestSchema = z.object({ diff --git a/src/api/tauri/rpc/schemas/validationValueObjects.ts b/src/api/tauri/rpc/schemas/validationValueObjects.ts index f12442a2aa..330b31a972 100644 --- a/src/api/tauri/rpc/schemas/validationValueObjects.ts +++ b/src/api/tauri/rpc/schemas/validationValueObjects.ts @@ -133,6 +133,8 @@ export const KeyInfoSchema = z.object({ }); export const FullKeyResponseSchema = z.object({ + credential_generation: z.number().int().nonnegative().default(0), + model_catalog_generation: z.number().int().nonnegative().default(0), id: z.string(), name: z.string().nullable(), agent_type: ModelTypeSchema, @@ -165,6 +167,7 @@ export const SaveKeyRequestSchema = z.object({ model_aliases: z.array(ModelAliasInfoSchema).optional(), model_variants: z.array(ModelVariantInfoSchema).optional(), default_variants: z.array(DefaultVariantInfoSchema).optional(), + default_variant_overrides: z.array(DefaultVariantInfoSchema).optional(), quota_info: z.record(z.string(), z.unknown()).optional(), has_local_key: z.boolean().optional(), is_listed: z.boolean().optional(), diff --git a/src/components/ModelSelectorPill/ModelSelectorPillView.tsx b/src/components/ModelSelectorPill/ModelSelectorPillView.tsx index 51cf4e55af..adf1bd8d65 100644 --- a/src/components/ModelSelectorPill/ModelSelectorPillView.tsx +++ b/src/components/ModelSelectorPill/ModelSelectorPillView.tsx @@ -433,7 +433,7 @@ const ModelSelectorPillView = forwardRef< {levelLabel && ( {levelLabel} @@ -451,7 +451,7 @@ const ModelSelectorPillView = forwardRef< ariaExpanded={hasModelSelection ? open : undefined} ariaLabel={`${ariaLabel ?? defaultLabel}: ${combinedLabel}${variant?.fast ? " · Fast" : ""}`} dataTestId={dataTestId} - className={`shrink-0 ${triggerClassName ?? ""} ${className ?? ""}`} + className={`max-w-full shrink-0 ${triggerClassName ?? ""} ${className ?? ""}`} leadingFlush={triggerLeadingFlush} paddingX={paddingX} onClick={hasModelSelection ? openMenu : onClick} @@ -466,8 +466,8 @@ const ModelSelectorPillView = forwardRef< span]:opacity-50" : ""} ${triggerClassName ?? ""}`.trim()} + className={`max-w-full shrink-0 flex-wrap text-[13px] ${className ?? ""}`} + segmentClassName={`h-[28px] max-w-full ${disabled ? "cursor-not-allowed [&>span]:opacity-50" : ""} ${triggerClassName ?? ""}`.trim()} /> ); diff --git a/src/engines/ChatPanel/InputArea/components/ModelPill.memberOwnership.test.ts b/src/engines/ChatPanel/InputArea/components/ModelPill.memberOwnership.test.ts index dec1602652..2285ca41df 100644 --- a/src/engines/ChatPanel/InputArea/components/ModelPill.memberOwnership.test.ts +++ b/src/engines/ChatPanel/InputArea/components/ModelPill.memberOwnership.test.ts @@ -4,11 +4,14 @@ import React, { act } from "react"; import { type Root, createRoot } from "react-dom/client"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import { Message } from "@src/components/Message"; +import { creatorDefaultModelSelectionAtom } from "@src/store/session/creatorDefaultModelAtom"; import { modelSelectorAtom } from "@src/store/ui/modelSelectorAtom"; import ModelPill from "./ModelPill"; const fixture = vi.hoisted(() => ({ + sessionId: "member-session", rootApplyModelPick: vi.fn(), setMemberModel: vi.fn(), memberSession: { @@ -33,7 +36,7 @@ vi.mock("@src/engines/ChatPanel/ConversationExecutionBindingContext", () => ({ useConversationExecutionBinding: () => null, })); vi.mock("@src/engines/SessionCore/hooks/session", () => ({ - useSessionId: () => ({ sessionId: "member-session" }), + useSessionId: () => ({ sessionId: fixture.sessionId }), })); vi.mock("@src/hooks/models/useValidatedLastPair", () => ({ useValidatedLastPair: () => null, @@ -64,7 +67,7 @@ vi.mock("@src/store/ui/chatPanel/displayPrefsAtoms", async () => { vi.mock("@src/components/AnyIcon", () => ({ default: () => null })); vi.mock("@src/components/ModelIcon", () => ({ default: () => null })); vi.mock("@src/components/Message", () => ({ - Message: { info: vi.fn(), warning: vi.fn() }, + Message: { info: vi.fn(), warning: vi.fn(), error: vi.fn() }, })); vi.mock("@src/components/SelectorPill", () => ({ default: () => null })); vi.mock("@src/components/ModelSelectorPill", async () => { @@ -126,6 +129,7 @@ describe("ModelPill direct Member ownership", () => { beforeEach(() => { Object.assign(globalThis, { IS_REACT_ACT_ENVIRONMENT: true }); fixture.rootApplyModelPick.mockReset(); + fixture.sessionId = "member-session"; fixture.setMemberModel.mockReset().mockResolvedValue(undefined); store = createStore(); store.set(modelSelectorAtom, { isOpen: false }); @@ -134,7 +138,11 @@ describe("ModelPill direct Member ownership", () => { root = createRoot(container); act(() => root.render( - React.createElement(Provider, { store }, React.createElement(ModelPill)) + React.createElement( + Provider, + { store }, + React.createElement(ModelPill, { key: fixture.sessionId }) + ) ) ); }); @@ -174,4 +182,92 @@ describe("ModelPill direct Member ownership", () => { ); expect(fixture.rootApplyModelPick).not.toHaveBeenCalled(); }); + + it("updates the creator default only after the session accepts the pick", async () => { + let finish!: () => void; + fixture.setMemberModel.mockImplementationOnce( + () => + new Promise((resolve) => { + finish = resolve; + }) + ); + act(() => + container + .querySelector( + '[data-testid="chat-model-pill-model"]' + ) + ?.click() + ); + act(() => + container + .querySelector('[data-testid="member-model-c"]') + ?.click() + ); + expect(store.get(creatorDefaultModelSelectionAtom)).toBeNull(); + await act(async () => finish()); + expect(store.get(creatorDefaultModelSelectionAtom)).toMatchObject({ + model: "gpt-member-c", + selectedAccountId: "member-account-c", + }); + }); + + it("keeps the previous default on failure and permits a retry", async () => { + fixture.setMemberModel.mockRejectedValueOnce(new Error("Patch refused")); + act(() => + container + .querySelector( + '[data-testid="chat-model-pill-model"]' + ) + ?.click() + ); + await act(async () => + container + .querySelector('[data-testid="member-model-c"]') + ?.click() + ); + expect(store.get(creatorDefaultModelSelectionAtom)).toBeNull(); + expect(Message.error).toHaveBeenCalledWith("Patch refused"); + await act(async () => + container + .querySelector('[data-testid="member-model-c"]') + ?.click() + ); + expect(store.get(creatorDefaultModelSelectionAtom)).toMatchObject({ + model: "gpt-member-c", + }); + }); + + it("does not publish a late default after navigating to another session", async () => { + let finish!: () => void; + fixture.setMemberModel.mockImplementationOnce( + () => + new Promise((resolve) => { + finish = resolve; + }) + ); + act(() => + container + .querySelector( + '[data-testid="chat-model-pill-model"]' + ) + ?.click() + ); + act(() => + container + .querySelector('[data-testid="member-model-c"]') + ?.click() + ); + fixture.sessionId = "another-session"; + act(() => + root.render( + React.createElement( + Provider, + { store }, + React.createElement(ModelPill, { key: fixture.sessionId }) + ) + ) + ); + await act(async () => finish()); + expect(store.get(creatorDefaultModelSelectionAtom)).toBeNull(); + }); }); diff --git a/src/engines/ChatPanel/InputArea/components/ModelPill.tsx b/src/engines/ChatPanel/InputArea/components/ModelPill.tsx index bcf558007b..b690a17af5 100644 --- a/src/engines/ChatPanel/InputArea/components/ModelPill.tsx +++ b/src/engines/ChatPanel/InputArea/components/ModelPill.tsx @@ -17,7 +17,14 @@ * default atom only. Used by the SessionCreator preview. */ import { useAtom, useAtomValue, useSetAtom } from "jotai"; -import React, { memo, useCallback, useMemo, useRef, useState } from "react"; +import React, { + memo, + useCallback, + useLayoutEffect, + useMemo, + useRef, + useState, +} from "react"; import { useTranslation } from "react-i18next"; import type { CliAgentType } from "@src/api/tauri/rpc/schemas/validation"; @@ -48,6 +55,7 @@ import { creatorDefaultModelSelectionAtom, extractModelPair, } from "@src/store/session/creatorDefaultModelAtom"; +import { dispatchCategoryAtom } from "@src/store/session/creatorStateAtom"; import { findRecentByCredentialSource, recentModelEntriesAtom, @@ -97,6 +105,7 @@ const ModelPillComponent: React.FC = () => { // strings, etc. — fields not stored on the session row). const creatorDefaultLastModel = useValidatedLastPair(); const setCreatorDefaultModel = useSetAtom(creatorDefaultModelSelectionAtom); + const creatorDispatchCategory = useAtomValue(dispatchCategoryAtom); const { sessionId } = useSessionId(); const [pendingRuntimeOwner, setPendingRuntimeOwner] = useState(sessionId); @@ -107,6 +116,13 @@ const ModelPillComponent: React.FC = () => { setPendingRuntimeOwner(sessionId); setPendingRuntimePick(null); } + const selectionOwner = useRef(0); + useLayoutEffect(() => { + selectionOwner.current += 1; + return () => { + selectionOwner.current += 1; + }; + }, [sessionId]); const isInSession = Boolean(sessionId); const session = useAtomValue(sessionByIdAtom(sessionId ?? "")); const recentModelEntries = useAtomValue(recentModelEntriesAtom); @@ -231,19 +247,20 @@ const ModelPillComponent: React.FC = () => { }, [lastModel]); const handleConfigChange = useCallback( - (config: AdvancedConfig) => { + async (config: AdvancedConfig): Promise => { + const generation = ++selectionOwner.current; // Team-conversation composer: the pick belongs to the remembered // runner setup, never to the imported row (whose model field is // deliberately empty and whose patches a family refresh wipes). // Runtime has its own standard New Session picker. This picker only // changes the model/account source for that selected runtime. if (conversationBinding) { - if ( - conversationBinding.applyModelPick(config, pendingRuntimeSelection) - ) { - clearPendingRuntimeSelection(); - } - return; + const accepted = conversationBinding.applyModelPick( + config, + pendingRuntimeSelection + ); + if (accepted) clearPendingRuntimeSelection(); + return accepted; } // In-session: keySource / cliAgentType / tier are session-create // immutables (mis-billing risk + zombie CLI processes if mutated; @@ -287,12 +304,11 @@ const ModelPillComponent: React.FC = () => { credentialSourceDiffers ) { Message.warning(t("sessions:modelPill.immutableInSession")); - return; + return false; } } const pair = extractModelPair(config); - setCreatorDefaultModel(pair); // For in-session model swaps, persist `(model, accountId)` to // the session row via session_patch. For market sessions the @@ -315,9 +331,26 @@ const ModelPillComponent: React.FC = () => { if (runtimeStatus === "running" && accountChanges) { Message.info(t("sessions:modelPill.appliesNextTurn")); } - void setSessionModel(wireModel, wireAccount); + try { + await setSessionModel(wireModel, wireAccount); + } catch (error) { + if (generation === selectionOwner.current) { + Message.error( + error instanceof Error ? error.message : String(error) + ); + } + return false; + } + } else { + return false; } } + if (generation !== selectionOwner.current) return false; + setCreatorDefaultModel( + pair, + paletteCategoryOverride ?? creatorDispatchCategory + ); + return true; }, [ setCreatorDefaultModel, @@ -328,6 +361,8 @@ const ModelPillComponent: React.FC = () => { conversationBinding, clearPendingRuntimeSelection, pendingRuntimeSelection, + paletteCategoryOverride, + creatorDispatchCategory, t, ] ); @@ -368,7 +403,9 @@ const ModelPillComponent: React.FC = () => { model: nextModelId, }; - handleConfigChange(updatedConfig); + handleConfigChange(updatedConfig).catch((error) => { + Message.error(error instanceof Error ? error.message : String(error)); + }); }, [advancedConfig, handleConfigChange, lastModel] ); diff --git a/src/engines/ChatPanel/InputArea/components/useContextUsageInfo.session.test.ts b/src/engines/ChatPanel/InputArea/components/useContextUsageInfo.session.test.ts new file mode 100644 index 0000000000..6f97446814 --- /dev/null +++ b/src/engines/ChatPanel/InputArea/components/useContextUsageInfo.session.test.ts @@ -0,0 +1,89 @@ +// @vitest-environment jsdom +import { Provider, atom, createStore } from "jotai"; +import { act, createElement } from "react"; +import { createRoot } from "react-dom/client"; +import { expect, it, vi } from "vitest"; + +import { useContextUsageInfo } from "./useContextUsageInfo"; + +const fixture = vi.hoisted(() => ({ sessionId: "native-session" })); +const sessionAtom = atom({ model: "session-model", accountId: "session-key" }); +const usageAtom = atom<{ maxTokens: number; usedTokens: number } | null>(null); +vi.mock("react-i18next", () => ({ + useTranslation: () => ({ t: (key: string) => key }), +})); +vi.mock("@src/engines/SessionCore/hooks/session", () => ({ + useSessionId: () => ({ sessionId: fixture.sessionId }), +})); +vi.mock("@src/store/session", () => ({ sessionByIdAtom: () => sessionAtom })); +vi.mock("@src/api/tauri/externalHistory", () => ({ + getImportedHistorySourceBySessionId: () => null, +})); +vi.mock("@src/util/session/sessionDispatch", () => ({ + isCliSession: () => false, +})); +vi.mock("@src/hooks/models/useValidatedLastPair", () => ({ + useValidatedLastPair: () => ({ + model: "creator-model", + selectedAccountId: "creator-key", + }), +})); +vi.mock("@src/hooks/keyVault", () => ({ + useKeyVault: () => ({ + accounts: [ + { + id: "session-key", + modelVariants: [ + { + model: "session-model", + base_model: "session-model", + context_window: 64_000, + }, + ], + }, + { + id: "creator-key", + modelVariants: [ + { + model: "creator-model", + base_model: "creator-model", + context_window: 1_000_000, + }, + ], + }, + ], + }), +})); +vi.mock("@src/store/session/cliSessionStatusAtom", async () => { + const { atom: makeAtom } = await import("jotai"); + return { + sessionContextTokensAtom: makeAtom(1000), + // Factories avoid reading the test binding before module evaluation. + sessionContextUsageAtom: makeAtom((get) => get(usageAtom)), + }; +}); + +it("uses the active session account until runtime telemetry supplies its actual window", async () => { + Object.assign(globalThis, { IS_REACT_ACT_ENVIRONMENT: true }); + const store = createStore(); + const container = document.createElement("div"); + const root = createRoot(container); + function Probe() { + return createElement("span", null, useContextUsageInfo().maxTokens); + } + try { + await act(async () => + root.render(createElement(Provider, { store }, createElement(Probe))) + ); + expect(container.textContent).toBe("64000"); + await act(async () => + store.set(usageAtom, { maxTokens: 128_000, usedTokens: 1000 }) + ); + expect(container.textContent).toBe("128000"); + await act(async () => store.set(usageAtom, null)); + expect(container.textContent).toBe("64000"); + } finally { + await act(async () => root.unmount()); + Reflect.deleteProperty(globalThis, "IS_REACT_ACT_ENVIRONMENT"); + } +}); diff --git a/src/engines/ChatPanel/InputArea/components/useContextUsageInfo.ts b/src/engines/ChatPanel/InputArea/components/useContextUsageInfo.ts index 3fc43b3a83..d098ac175f 100644 --- a/src/engines/ChatPanel/InputArea/components/useContextUsageInfo.ts +++ b/src/engines/ChatPanel/InputArea/components/useContextUsageInfo.ts @@ -6,6 +6,7 @@ import { getImportedHistorySourceBySessionId } from "@src/api/tauri/externalHist import { useSessionId } from "@src/engines/SessionCore/hooks/session"; import { useKeyVault } from "@src/hooks/keyVault"; import { useValidatedLastPair } from "@src/hooks/models/useValidatedLastPair"; +import { sessionByIdAtom } from "@src/store/session"; import type { ContextUsageSnapshot } from "@src/store/session/cliSessionStatusAtom"; import { sessionContextTokensAtom, @@ -16,6 +17,7 @@ import { type ResolvedModelVariantFields, getModelVariantBaseModel, } from "@src/util/modelVariants"; +import { selectionFromSession } from "@src/util/session/selectionFromSession"; import { isCliSession } from "@src/util/session/sessionDispatch"; type ContextWindowVariant = Pick< @@ -124,7 +126,11 @@ export function useContextUsageInfo(): ContextUsageInfo { (!!sessionId && isCliSession(sessionId)); const sessionTokens = useAtomValue(sessionContextTokensAtom); const contextUsage = useAtomValue(sessionContextUsageAtom); - const lastModel = useValidatedLastPair(); + const creatorDefault = useValidatedLastPair(); + const session = useAtomValue(sessionByIdAtom(sessionId ?? "")); + const lastModel = sessionId + ? selectionFromSession(session, null) + : creatorDefault; const { accounts } = useKeyVault({ autoLoad: true }); const modelName = lastModel?.model || lastModel?.listingModel || ""; diff --git a/src/hooks/keyVault/defaultVariantSaveCoordinator.test.ts b/src/hooks/keyVault/defaultVariantSaveCoordinator.test.ts new file mode 100644 index 0000000000..dd1668ae68 --- /dev/null +++ b/src/hooks/keyVault/defaultVariantSaveCoordinator.test.ts @@ -0,0 +1,197 @@ +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import type { KeyInfo, SaveKeyRequest } from "@src/api/types/keys"; + +import { saveDefaultVariantOverrides } from "./defaultVariantSaveCoordinator"; +import { + getSharedLocalKeys, + publishSharedLocalKeys, +} from "./sharedLocalKeyStore"; + +const mocks = vi.hoisted(() => ({ + saveKey: vi.fn<(request: SaveKeyRequest) => Promise>(), +})); +vi.mock("@src/api/services/keyValidation", () => ({ + saveKey: mocks.saveKey, + listKeys: vi.fn(), +})); + +const choice = (model: string, base_model = "family") => ({ + base_model, + model, +}); +function key(defaults = [choice("original")]): KeyInfo { + return { + id: "account", + agent_type: "codex", + name: "original name", + default_variants: defaults, + } as KeyInfo; +} +function deferred() { + let resolve!: (key: KeyInfo) => void; + let reject!: (error: Error) => void; + const promise = new Promise((yes, no) => { + resolve = yes; + reject = no; + }); + return { promise, resolve, reject }; +} +function submit(model: string, family = "family") { + return saveDefaultVariantOverrides({ + id: "account", + agent_type: "codex", + default_variant_overrides: [choice(model, family)], + }); +} +const current = () => + Object.fromEntries( + (getSharedLocalKeys()[0]?.default_variants ?? []).map((item) => [ + item.base_model, + item.model, + ]) + ); +beforeEach(() => { + vi.clearAllMocks(); + publishSharedLocalKeys([key()]); +}); + +describe("shared default family write coordinator", () => { + it("two failed picks roll back to confirmed state, not the earlier optimistic pick", async () => { + const a = deferred(), + b = deferred(); + mocks.saveKey.mockReturnValueOnce(a.promise).mockReturnValueOnce(b.promise); + const first = submit("a").catch(() => null); + const second = submit("b").catch(() => null); + expect(current()).toEqual({ family: "b" }); + expect(mocks.saveKey).toHaveBeenCalledTimes(1); + a.reject(new Error("a failed")); + await first; + expect(current()).toEqual({ family: "b" }); + expect(mocks.saveKey).toHaveBeenCalledTimes(2); + b.reject(new Error("b failed")); + await second; + expect(current()).toEqual({ family: "original" }); + }); + + it("an older reply cannot cover a newer pick or overwrite unrelated account fields", async () => { + const a = deferred(), + b = deferred(); + mocks.saveKey.mockReturnValueOnce(a.promise).mockReturnValueOnce(b.promise); + const first = submit("a"); + const second = submit("b").catch(() => null); + publishSharedLocalKeys([ + { ...getSharedLocalKeys()[0], name: "renamed elsewhere" }, + ]); + a.resolve(key([choice("a")])); + await first; + expect(current()).toEqual({ family: "b" }); + expect(getSharedLocalKeys()[0].name).toBe("renamed elsewhere"); + b.reject(new Error("b failed")); + await second; + expect(current()).toEqual({ family: "a" }); + }); + + it("separates families and preserves independent source updates", async () => { + publishSharedLocalKeys([ + key([choice("original"), choice("provider-b", "other")]), + ]); + const a = deferred(), + b = deferred(); + mocks.saveKey.mockReturnValueOnce(a.promise).mockReturnValueOnce(b.promise); + const first = submit("a").catch(() => null); + const second = submit("b", "other"); + publishSharedLocalKeys([ + { + ...getSharedLocalKeys()[0], + default_variants: [ + ...getSharedLocalKeys()[0].default_variants!, + choice("external", "third"), + ], + }, + ]); + a.reject(new Error("a failed")); + await first; + b.resolve(key([choice("original"), choice("b", "other")])); + await second; + expect(current()).toEqual({ + family: "original", + other: "b", + third: "external", + }); + expect(mocks.saveKey).toHaveBeenNthCalledWith(2, { + id: "account", + agent_type: "codex", + default_variant_overrides: [choice("b", "other")], + }); + }); + + it("an external choice between failed picks becomes the confirmed rollback baseline", async () => { + const a = deferred(), + b = deferred(); + mocks.saveKey.mockReturnValueOnce(a.promise).mockReturnValueOnce(b.promise); + const first = submit("a").catch(() => null); + publishSharedLocalKeys([key([choice("external")])]); + const second = submit("b").catch(() => null); + a.reject(new Error("a failed")); + await first; + b.reject(new Error("b failed")); + await second; + expect(current()).toEqual({ family: "external" }); + }); + + it("bounds pending work and releases the lane so retry starts from confirmed state", async () => { + const a = deferred(); + mocks.saveKey + .mockReturnValueOnce(a.promise) + .mockRejectedValue(new Error("failed")); + const requests = Array.from({ length: 64 }, (_, i) => + submit(String(i)).catch(() => null) + ); + await expect(submit("overflow")).rejects.toThrow("Too many"); + a.reject(new Error("failed")); + await Promise.all(requests); + expect(current()).toEqual({ family: "original" }); + // An idle lane would wrongly mask this independent update. + publishSharedLocalKeys([key([choice("after-idle")])]); + expect(current()).toEqual({ family: "after-idle" }); + mocks.saveKey.mockResolvedValueOnce(key([choice("retry")])); + await submit("retry"); + expect(current()).toEqual({ family: "retry" }); + }); + + it("a slow account does not block saves for another account", async () => { + publishSharedLocalKeys([key(), { ...key(), id: "other-account" }]); + const slow = deferred(); + mocks.saveKey.mockReturnValueOnce(slow.promise).mockResolvedValueOnce({ + ...key([choice("other-pick")]), + id: "other-account", + }); + const first = submit("slow-pick"); + await saveDefaultVariantOverrides({ + id: "other-account", + agent_type: "codex", + default_variant_overrides: [choice("other-pick")], + }); + expect(mocks.saveKey).toHaveBeenCalledTimes(2); + expect(getSharedLocalKeys()[1].default_variants).toEqual([ + choice("other-pick"), + ]); + expect(current()).toEqual({ family: "slow-pick" }); + slow.resolve(key([choice("slow-pick")])); + await first; + }); + + it("deletion stops queued writes and a late response cannot recreate the account", async () => { + const a = deferred(); + mocks.saveKey.mockReturnValueOnce(a.promise); + const first = submit("a"); + const second = submit("b").catch(() => null); + publishSharedLocalKeys([]); + a.resolve(key([choice("a")])); + await first; + await second; + expect(getSharedLocalKeys()).toEqual([]); + expect(mocks.saveKey).toHaveBeenCalledTimes(1); + }); +}); diff --git a/src/hooks/keyVault/defaultVariantSaveCoordinator.ts b/src/hooks/keyVault/defaultVariantSaveCoordinator.ts new file mode 100644 index 0000000000..2b81306b61 --- /dev/null +++ b/src/hooks/keyVault/defaultVariantSaveCoordinator.ts @@ -0,0 +1,246 @@ +import { saveKey } from "@src/api/services/keyValidation"; +import type { + DefaultVariantInfo, + KeyInfo, + SaveKeyRequest, +} from "@src/api/types/keys"; + +import { + getSharedLocalKeys, + subscribeSharedLocalKeys, + updateSharedLocalKeys, +} from "./sharedLocalKeyStore"; + +interface Operation { + choices: DefaultVariantInfo[]; + epochs: Map; + resolve: (key: KeyInfo) => void; + reject: (error: unknown) => void; +} +interface AccountLane { + id: string; + modelType: KeyInfo["agent_type"]; + confirmed: Map; + projected: Map; + owned: Set; + epochs: Map; + pending: Operation[]; + unsubscribe?: () => void; +} + +const MAX_PENDING_PER_ACCOUNT = 64; +const MAX_ACTIVE_ACCOUNTS = 64; +// Only active saves retain a lane/subscription. There are no idle timers. +const lanes = new Map(); +const choicesOf = (key: KeyInfo) => + new Map( + (key.default_variants ?? []).map((choice) => [ + choice.base_model, + choice.model, + ]) + ); +const read = (lane: Pick) => + getSharedLocalKeys().find( + (key) => key.id === lane.id && key.agent_type === lane.modelType + ); + +function assign( + map: Map, + family: string, + model: string | undefined +) { + if (model === undefined) map.delete(family); + else map.set(family, model); +} + +function observe(lane: AccountLane, current: Map) { + for (const family of lane.owned) { + const model = current.get(family); + if ( + model === lane.projected.get(family) || + model === lane.confirmed.get(family) + ) + continue; + if ( + lane.pending.some( + (operation) => + operation.epochs.get(family) === lane.epochs.get(family) && + operation.choices.some( + (choice) => choice.base_model === family && choice.model === model + ) + ) + ) + continue; + // An independent edit supersedes older intent for this family only. + assign(lane.confirmed, family, model); + lane.epochs.set(family, (lane.epochs.get(family) ?? 0) + 1); + lane.owned.delete(family); + } +} + +function reconcile(lane: AccountLane) { + const current = read(lane); + if (!current) return; // Never recreate a deleted/replaced account. + const visible = choicesOf(current); + observe(lane, visible); + const projected = new Map(visible); + for (const family of lane.owned) + assign(projected, family, lane.confirmed.get(family)); + for (const operation of lane.pending) { + for (const choice of operation.choices) { + if ( + lane.owned.has(choice.base_model) && + operation.epochs.get(choice.base_model) === + lane.epochs.get(choice.base_model) + ) { + projected.set(choice.base_model, choice.model); + } + } + } + lane.projected = projected; + if ( + projected.size === visible.size && + [...projected].every(([family, model]) => visible.get(family) === model) + ) + return; + updateSharedLocalKeys((keys) => + keys.map((key) => + key.id === lane.id && key.agent_type === lane.modelType + ? { + ...key, + default_variants: [...projected].map(([base_model, model]) => ({ + base_model, + model, + })), + } + : key + ) + ); +} + +async function drain(lane: AccountLane) { + try { + while (lane.pending.length) { + const operation = lane.pending[0]; + try { + const before = read(lane); + if (!before) throw new Error("Account no longer exists"); + const active = operation.choices.filter( + (choice) => + operation.epochs.get(choice.base_model) === + lane.epochs.get(choice.base_model) + ); + const saved = active.length + ? await saveKey({ + id: lane.id, + agent_type: lane.modelType, + default_variant_overrides: active, + }) + : before; + // Only this operation's families become confirmed. A returned full + // row must not replace newer picks or unrelated health/catalog edits. + const defaults = choicesOf(saved); + for (const choice of active) { + if ( + operation.epochs.get(choice.base_model) === + lane.epochs.get(choice.base_model) + ) { + assign( + lane.confirmed, + choice.base_model, + defaults.get(choice.base_model) + ); + } + } + lane.pending.shift(); + reconcile(lane); + operation.resolve(read(lane) ?? saved); + } catch (error) { + lane.pending.shift(); + reconcile(lane); + operation.reject(error); + } + } + } finally { + lane.unsubscribe?.(); + lanes.delete(lane.id); + } +} + +/** All default-pick entry points share this ordered, family-scoped writer. */ +export function saveDefaultVariantOverrides( + request: SaveKeyRequest +): Promise { + if ( + !request.id || + !request.default_variant_overrides || + Object.entries(request).some( + ([key, value]) => + value !== undefined && + !["id", "agent_type", "default_variant_overrides"].includes(key) + ) + ) { + return Promise.reject( + new Error( + "Default variant overrides require a separate account-scoped request" + ) + ); + } + const before = read({ id: request.id, modelType: request.agent_type }); + if (!before) return Promise.reject(new Error("Account no longer exists")); + if (!request.default_variant_overrides.length) return Promise.resolve(before); + let lane = lanes.get(request.id); + if ( + (lane && lane.pending.length >= MAX_PENDING_PER_ACCOUNT) || + (!lane && lanes.size >= MAX_ACTIVE_ACCOUNTS) + ) { + return Promise.reject(new Error("Too many pending default model changes")); + } + const start = !lane; + if (!lane) { + lane = { + id: request.id, + modelType: request.agent_type, + confirmed: choicesOf(before), + projected: choicesOf(before), + owned: new Set(), + epochs: new Map(), + pending: [], + }; + lanes.set(request.id, lane); + } else { + if (lane.modelType !== request.agent_type) + return Promise.reject(new Error("Account provider changed")); + observe(lane, choicesOf(before)); + } + const epochs = new Map(); + const choices = request.default_variant_overrides.map((choice) => ({ + ...choice, + })); + for (const choice of choices) { + if (!lane.owned.has(choice.base_model)) { + assign( + lane.confirmed, + choice.base_model, + choicesOf(before).get(choice.base_model) + ); + } + lane.owned.add(choice.base_model); + const epoch = lane.epochs.get(choice.base_model) ?? 0; + lane.epochs.set(choice.base_model, epoch); + epochs.set(choice.base_model, epoch); + } + const promise = new Promise((resolve, reject) => { + lane.pending.push({ choices, epochs, resolve, reject }); + }); + if (start) { + const active = lane; + lane.unsubscribe = subscribeSharedLocalKeys(() => reconcile(active)); + } + reconcile(lane); + if (start) + drain(lane).catch((error) => { + for (const operation of lane.pending.splice(0)) operation.reject(error); + }); + return promise; +} diff --git a/src/hooks/keyVault/useLocalKeys.test.ts b/src/hooks/keyVault/useLocalKeys.test.ts index d58738b6ff..4a92bce38f 100644 --- a/src/hooks/keyVault/useLocalKeys.test.ts +++ b/src/hooks/keyVault/useLocalKeys.test.ts @@ -11,12 +11,13 @@ import React, { act } from "react"; import { createRoot } from "react-dom/client"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; -import type { KeyInfo } from "@src/api/services/keyValidation"; +import type { KeyInfo, SaveKeyRequest } from "@src/api/services/keyValidation"; import { useLocalKeys } from "./useLocalKeys"; const mocks = vi.hoisted(() => ({ listKeys: vi.fn<() => Promise>(), + saveKey: vi.fn<(request: SaveKeyRequest) => Promise>(), })); vi.mock("@src/api/services/keyValidation", () => ({ @@ -26,7 +27,7 @@ vi.mock("@src/api/services/keyValidation", () => ({ getFullKey: vi.fn(), getKey: vi.fn(), refreshKeyQuota: vi.fn(), - saveKey: vi.fn(), + saveKey: mocks.saveKey, updateKeyHealth: vi.fn(), validateKey: vi.fn(), })); @@ -133,3 +134,106 @@ describe("useLocalKeys loading", () => { expect(result.current.loading).toBe(false); }); }); + +describe("useLocalKeys default family deltas", () => { + it("updates the display immediately while sending only the edited family", async () => { + const key = { + ...keyRecord("delta-key"), + default_variants: [ + { base_model: "family-a", model: "provider-a" }, + { base_model: "family-b", model: "provider-b" }, + ], + }; + mocks.listKeys.mockResolvedValueOnce([key]); + const result = renderLocalKeys(); + await act(async () => { + await result.current.refreshAgents(true); + }); + const save = deferred(); + mocks.saveKey.mockReturnValueOnce(save.promise); + let pending!: Promise; + act(() => { + pending = result.current.saveKey({ + id: key.id, + agent_type: key.agent_type, + default_variant_overrides: [ + { base_model: "family-a", model: "chosen-a" }, + ], + }); + }); + expect(result.current.allKeys[0].default_variants).toEqual([ + { base_model: "family-a", model: "chosen-a" }, + { base_model: "family-b", model: "provider-b" }, + ]); + expect(mocks.saveKey).toHaveBeenLastCalledWith({ + id: key.id, + agent_type: key.agent_type, + default_variant_overrides: [ + { base_model: "family-a", model: "chosen-a" }, + ], + }); + await act(async () => { + save.resolve({ + ...key, + default_variants: [ + { base_model: "family-a", model: "chosen-a" }, + { base_model: "family-b", model: "provider-b" }, + ], + }); + await pending; + }); + }); + + it("a rejected family pick preserves independently refreshed defaults and accounts", async () => { + const key = { + ...keyRecord("delta-key"), + default_variants: [ + { base_model: "family-a", model: "provider-a" }, + { base_model: "family-b", model: "provider-b" }, + ], + }; + mocks.listKeys.mockResolvedValueOnce([key]); + const result = renderLocalKeys(); + await act(async () => { + await result.current.refreshAgents(true); + }); + let reject!: (error: Error) => void; + mocks.saveKey.mockReturnValueOnce( + new Promise((_resolve, no) => { + reject = no; + }) + ); + let pending!: Promise; + act(() => { + pending = result.current.saveKey({ + id: key.id, + agent_type: key.agent_type, + default_variant_overrides: [ + { base_model: "family-a", model: "failed-a" }, + ], + }); + }); + mocks.listKeys.mockResolvedValueOnce([ + { + ...key, + default_variants: [ + { base_model: "family-a", model: "failed-a" }, + { base_model: "family-b", model: "newer-b" }, + ], + }, + keyRecord("new-account"), + ]); + await act(async () => { + await result.current.refreshAgents(true); + }); + await act(async () => { + reject(new Error("failed")); + await pending; + }); + expect(result.current.allKeys).toHaveLength(2); + expect(result.current.allKeys[0].default_variants).toEqual([ + { base_model: "family-a", model: "provider-a" }, + { base_model: "family-b", model: "newer-b" }, + ]); + }); +}); diff --git a/src/hooks/keyVault/useLocalKeys.ts b/src/hooks/keyVault/useLocalKeys.ts index 3c9b8ec435..2d1b48df11 100644 --- a/src/hooks/keyVault/useLocalKeys.ts +++ b/src/hooks/keyVault/useLocalKeys.ts @@ -26,6 +26,7 @@ import type { } from "@src/api/services/keyValidation"; import { createLogger } from "@src/hooks/logger"; +import { saveDefaultVariantOverrides } from "./defaultVariantSaveCoordinator"; import { runSharedQuotaRefresh } from "./quotaRefreshCoordinator"; import { areSharedLocalKeysLoaded, @@ -200,6 +201,13 @@ export function useLocalKeys( */ const saveKeyFn = useCallback( async (request: SaveKeyRequest): Promise => { + if (request.default_variant_overrides) { + try { + return await saveDefaultVariantOverrides(request); + } catch { + return null; + } + } const previousKeys = getSharedLocalKeys(); let appliedOptimisticUpdate = false; diff --git a/src/hooks/models/accountModelCatalog.test.ts b/src/hooks/models/accountModelCatalog.test.ts new file mode 100644 index 0000000000..3b97da6d35 --- /dev/null +++ b/src/hooks/models/accountModelCatalog.test.ts @@ -0,0 +1,110 @@ +import { describe, expect, it } from "vitest"; + +import { + CLI_AGENT, + NATIVE_HARNESS_TYPE, +} from "@src/api/tauri/rpc/schemas/validation"; +import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; +import type { AgentRegistry } from "@src/store/session/agentRegistryAtom"; + +import { + getModelPickerAccounts, + groupCatalogModels, + resolveAccountModelVariant, + selectableAccountModelIds, +} from "./accountModelCatalog"; +import { withNativeHarnessModels } from "./nativeHarnessAccountModels"; + +function account(overrides: Partial = {}): KeyVaultAccount { + return { + id: "key", + name: "Account", + modelType: "codex", + enabled: true, + hasLocalKey: true, + isListed: false, + hasKey: true, + hasApiKey: false, + hasSessionToken: true, + status: "ready", + authMethod: "oauth", + availableModels: [], + enabledModels: ["catalog-family"], + modelVariants: [ + { + model: "deployment-a", + base_model: "catalog-family", + reasoning: "low", + fast: false, + }, + { + model: "deployment-b", + base_model: "catalog-family", + reasoning: "high", + fast: true, + }, + ], + ...overrides, + }; +} + +describe("account model catalog", () => { + it("derives selectable IDs and grouping from a variants-only catalog", () => { + const source = account(); + expect(selectableAccountModelIds(source)).toEqual([ + "deployment-a", + "deployment-b", + ]); + expect( + groupCatalogModels(selectableAccountModelIds(source), [source]) + ).toEqual([ + { + label: "catalog-family", + sortVersion: -1, + models: ["deployment-a", "deployment-b"], + }, + ]); + expect(resolveAccountModelVariant(source, "deployment-b")).toEqual( + source.modelVariants?.[1] + ); + }); + + it.each([ + { enabled: false }, + { hasKey: false }, + { status: "error" as const }, + { enabledModels: [] }, + ])("never offers an unavailable account/model: %j", (overrides) => { + expect(selectableAccountModelIds(account(overrides))).toEqual([]); + }); + + it.each([CLI_AGENT.CODEX, CLI_AGENT.CLAUDE_CODE, CLI_AGENT.CURSOR] as const)( + "preserves %s health through native adaptation", + (modelType) => { + const source = account({ + modelType, + status: "error", + canUseNativeHarness: true, + nativeHarnessType: NATIVE_HARNESS_TYPE.CURSOR, + }); + const adapted = withNativeHarnessModels([source], "rust_agent")[0]; + expect(adapted.status).toBe("error"); + expect(selectableAccountModelIds(adapted)).toEqual([]); + } + ); + + it("uses the same account capability for Agent Org and palette callers", () => { + const registry = { + agents: [], + apiProviders: [], + } as unknown as AgentRegistry; + const unsupported = account({ + id: "unsupported", + supportsRustAgents: false, + }); + const supported = account({ id: "supported", supportsRustAgents: true }); + expect( + getModelPickerAccounts(registry, [unsupported, supported], "rust_agent") + ).toEqual([supported]); + }); +}); diff --git a/src/hooks/models/accountModelCatalog.ts b/src/hooks/models/accountModelCatalog.ts new file mode 100644 index 0000000000..5dd75c7612 --- /dev/null +++ b/src/hooks/models/accountModelCatalog.ts @@ -0,0 +1,147 @@ +import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; +import type { AgentRegistry } from "@src/store/session/agentRegistryAtom"; +import { type ModelGroup, groupModels } from "@src/util/modelGrouping"; +import { + type ResolvedModelVariantFields, + resolveModelVariantFields, +} from "@src/util/modelVariants"; +import { getModelEffortBaseModel } from "@src/util/selectableModelVariants"; + +import { withNativeHarnessModels } from "./nativeHarnessAccountModels"; +import { + getCliCompatibleAccounts, + isSourceCompatibleWithAgent, +} from "./useAgentCompatibility"; + +/** A launchable credential is independent of the executor's model adapter. */ +export function isSelectableModelAccount(account: KeyVaultAccount): boolean { + return account.enabled && account.status === "ready" && account.hasKey; +} + +/** Catalog membership includes synthesized variants, even without available IDs. */ +export function accountModelIds( + account: Pick +): string[] { + return [ + ...new Set( + [ + ...(account.availableModels ?? []), + ...(account.modelVariants ?? []).map((variant) => variant.model), + ].filter(Boolean) + ), + ]; +} + +/** An explicit empty enabled set stays empty; variant rows inherit their base. */ +export function accountHasModel( + account: Pick< + KeyVaultAccount, + "availableModels" | "enabledModels" | "enabled" | "modelVariants" + >, + modelId: string +): boolean { + if (!account.enabled) return false; + const enabled = new Set(account.enabledModels ?? []); + return ( + enabled.has(modelId) || + (account.modelVariants ?? []).some( + (variant) => variant.model === modelId && enabled.has(variant.base_model) + ) + ); +} + +export function selectableAccountModelIds(account: KeyVaultAccount): string[] { + return isSelectableModelAccount(account) + ? accountModelIds(account).filter((modelId) => + accountHasModel(account, modelId) + ) + : []; +} + +export function resolveAccountModelVariant( + account: KeyVaultAccount | undefined, + modelId: string +) { + return resolveModelVariantFields( + modelId, + account?.modelVariants?.find((variant) => variant.model === modelId) + ); +} + +/** Model-first browsing has no chosen account yet; source rows resolve locally. */ +export function resolveCatalogModelVariant( + accounts: readonly KeyVaultAccount[], + modelId: string +) { + const account = accounts.find( + (candidate) => + isSelectableModelAccount(candidate) && + accountHasModel(candidate, modelId) && + candidate.modelVariants?.some((variant) => variant.model === modelId) + ); + return resolveAccountModelVariant(account, modelId); +} + +/** Keep existing display grouping, but use catalog base IDs when supplied. */ +export function groupCatalogModels( + modelIds: string[], + accounts: readonly KeyVaultAccount[], + agentType?: string +): ModelGroup[] { + // This index is local to the projection; it adds no retained cache or loader. + const metadata = new Map(); + for (const account of accounts) { + if (!isSelectableModelAccount(account)) continue; + const enabled = new Set(account.enabledModels ?? []); + for (const variant of account.modelVariants ?? []) { + if ( + !metadata.has(variant.model) && + (enabled.has(variant.model) || enabled.has(variant.base_model)) + ) { + metadata.set(variant.model, variant); + } + } + } + const catalogBases = new Set( + [...metadata.values()].map((variant) => variant.base_model) + ); + const groups = new Map(); + for (const modelId of modelIds) { + const variant = metadata.get(modelId); + const base = variant?.base_model ?? getModelEffortBaseModel(modelId); + const [group] = groupModels([variant?.base_model ?? modelId], agentType); + if (!group) continue; + const key = catalogBases.has(base) + ? `catalog:${base}` + : `label:${group.label}`; + const existing = groups.get(key); + if (existing) existing.models.push(modelId); + else groups.set(key, { ...group, models: [modelId] }); + } + return [...groups.values()].sort( + (left, right) => right.sortVersion - left.sortVersion + ); +} + +/** Shared executor compatibility for palette, saved pair validation and Agent Org. */ +export function getModelPickerAccounts( + registry: AgentRegistry, + accounts: KeyVaultAccount[], + dispatchCategory: string | null, + cliAgentType?: string | null +): KeyVaultAccount[] { + if (dispatchCategory === "cli_agent" && cliAgentType) { + return getCliCompatibleAccounts(registry, cliAgentType, accounts); + } + if (dispatchCategory !== "rust_agent") return accounts; + return withNativeHarnessModels(accounts, dispatchCategory).filter( + (account) => + account.supportsRustAgents ?? + isSourceCompatibleWithAgent( + registry, + "rust_agent", + undefined, + account.modelType + ) + ); +} diff --git a/src/hooks/models/nativeHarnessAccountModels.ts b/src/hooks/models/nativeHarnessAccountModels.ts index 17325dde9b..8e15709f0b 100644 --- a/src/hooks/models/nativeHarnessAccountModels.ts +++ b/src/hooks/models/nativeHarnessAccountModels.ts @@ -65,7 +65,6 @@ export function withCursorNativeModels( return { ...account, - status: "ready", canUseNativeHarness: true, nativeHarnessType: account.nativeHarnessType ?? NATIVE_HARNESS_TYPE.CURSOR, availableModels: Array.from(availableModels), @@ -76,13 +75,13 @@ export function withCursorNativeModels( export function withClaudeCodeOAuthModels( account: KeyVaultAccount ): KeyVaultAccount { - return { ...account, status: "ready" }; + return account; } export function withCodexOAuthModels( account: KeyVaultAccount ): KeyVaultAccount { - return { ...account, status: "ready" }; + return account; } export function withNativeHarnessModels( diff --git a/src/hooks/models/resolveModelDisplaySelection.test.ts b/src/hooks/models/resolveModelDisplaySelection.test.ts new file mode 100644 index 0000000000..fdab934236 --- /dev/null +++ b/src/hooks/models/resolveModelDisplaySelection.test.ts @@ -0,0 +1,67 @@ +import { describe, expect, it } from "vitest"; + +import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; +import type { LastModelSelection } from "@src/store/session/creatorDefaultModelAtom"; + +import { resolveModelDisplaySelection } from "./resolveModelDisplaySelection"; + +const account: KeyVaultAccount = { + id: "key", + name: "Account", + modelType: "codex", + enabled: true, + hasLocalKey: true, + isListed: false, + hasKey: true, + hasApiKey: false, + hasSessionToken: true, + status: "ready", + availableModels: ["family"], + enabledModels: ["family"], + modelVariants: [ + { + model: "deployment-a", + base_model: "family", + reasoning: "low", + fast: false, + }, + { + model: "deployment-b", + base_model: "family", + reasoning: "high", + fast: false, + }, + ], + defaultVariants: [{ base_model: "family", model: "deployment-b" }], +}; +const selection: LastModelSelection = { + model: "family", + selectedAccountId: "key", + keySource: "own_key", +}; + +describe("model display selection catalog metadata", () => { + it("resolves an active bare family through variant-only account defaults", () => { + expect( + resolveModelDisplaySelection(selection, [account], true)?.model + ).toBe("deployment-b"); + }); + it("keeps an explicitly selected opaque variant rather than applying the account default", () => { + const explicit = { ...selection, model: "deployment-a" }; + expect(resolveModelDisplaySelection(explicit, [account], true)).toBe( + explicit + ); + }); + it("keeps historical sessions and explicitly disabled models unchanged", () => { + expect(resolveModelDisplaySelection(selection, [account], false)).toBe( + selection + ); + expect( + resolveModelDisplaySelection( + selection, + [{ ...account, enabledModels: [] }], + true + ) + ).toBe(selection); + }); +}); diff --git a/src/hooks/models/resolveModelDisplaySelection.ts b/src/hooks/models/resolveModelDisplaySelection.ts index 72673aa81e..7b5ce64509 100644 --- a/src/hooks/models/resolveModelDisplaySelection.ts +++ b/src/hooks/models/resolveModelDisplaySelection.ts @@ -3,10 +3,11 @@ import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; import { accountHasModel } from "@src/hooks/models/useModelAccountLookup"; import type { LastModelSelection } from "@src/store/session/creatorDefaultModelAtom"; import { resolveDefaultVariant } from "@src/util/defaultModelVariant"; + import { - parseModelVariant, - resolveModelVariantFields, -} from "@src/util/modelVariants"; + resolveAccountModelVariant, + selectableAccountModelIds, +} from "./accountModelCatalog"; /** * Resolves the effective model id shown in chat input pills. @@ -30,10 +31,6 @@ export function resolveModelDisplaySelection( } if (!isActiveSession) return selection; - // The session's stored model already encodes a user-chosen variant/effort; - // treat it as authoritative and do not overwrite it with the account default. - if (parseModelVariant(selection.model)) return selection; - const selectedAccount = accounts.find((account) => { if (selection.selectedAccountId) { return account.id === selection.selectedAccountId; @@ -51,11 +48,21 @@ export function resolveModelDisplaySelection( }); if (!selectedAccount) return selection; - const baseModel = resolveModelVariantFields(selection.model).base_model; - const accountModelIds = (selectedAccount.availableModels ?? []).filter( + const variant = resolveAccountModelVariant(selectedAccount, selection.model); + if ( + variant.reasoning || + variant.fast || + variant.thinking || + variant.base_model !== selection.model + ) + return selection; + + const baseModel = variant.base_model; + const accountModelIds = selectableAccountModelIds(selectedAccount).filter( (modelId) => accountHasModel(selectedAccount, modelId) && - resolveModelVariantFields(modelId).base_model === baseModel + resolveAccountModelVariant(selectedAccount, modelId).base_model === + baseModel ); if (accountModelIds.length === 0) return selection; @@ -65,7 +72,7 @@ export function resolveModelDisplaySelection( accountModelIds.includes(variant.model) )?.model; const variantInfos = accountModelIds.map((modelId) => - resolveModelVariantFields(modelId) + resolveAccountModelVariant(selectedAccount, modelId) ); const effectiveModel = resolveDefaultVariant( baseModel, diff --git a/src/hooks/models/useModelAccountLookup.test.ts b/src/hooks/models/useModelAccountLookup.test.ts index d7240234df..b20cebfb8a 100644 --- a/src/hooks/models/useModelAccountLookup.test.ts +++ b/src/hooks/models/useModelAccountLookup.test.ts @@ -13,6 +13,7 @@ function claudeAccount( modelType: "claude_code", status: "ready", enabled: true, + hasKey: true, availableModels: ["claude-opus-4-8"], enabledModels: ["claude-opus-4-8"], // Backend-synthesized effort ladder: variant ids exist ONLY here, diff --git a/src/hooks/models/useModelAccountLookup.ts b/src/hooks/models/useModelAccountLookup.ts index 451009f865..e78bfdbc74 100644 --- a/src/hooks/models/useModelAccountLookup.ts +++ b/src/hooks/models/useModelAccountLookup.ts @@ -2,46 +2,10 @@ import { useMemo } from "react"; import { type KeyVaultAccount, useKeyVault } from "@src/hooks/keyVault"; +import { selectableAccountModelIds } from "./accountModelCatalog"; import type { ModelAccountInfo } from "./types"; -/** - * Returns true if `account` has `modelId` enabled. - * - * Two ways an id counts as enabled: - * - it is in `enabledModels` directly, or - * - it is a variant row from `modelVariants` (e.g. a backend-synthesized - * effort rung like `claude-opus-4-8-high`) whose BASE model is enabled — - * variant ids never appear in enabledModels themselves, so gating on - * enabledModels alone hides every synthesized effort ladder from the - * picker's variant-edit affordance. - */ -export function accountHasModel( - account: Pick< - KeyVaultAccount, - "availableModels" | "enabledModels" | "enabled" | "modelVariants" - >, - modelId: string -): boolean { - if (!account.enabled) return false; - const enabled = new Set(account.enabledModels ?? []); - if (enabled.has(modelId)) return true; - return (account.modelVariants ?? []).some( - (variant) => variant.model === modelId && enabled.has(variant.base_model) - ); -} - -/** - * Every model id an account exposes: `availableModels` plus variant ids - * from `modelVariants` (deduped). The variant ids must enter the model - * universe or `groupByModel` family expansion can never offer them. - */ -export function accountModelIds(account: KeyVaultAccount): string[] { - const ids = new Set(account.availableModels ?? []); - for (const variant of account.modelVariants ?? []) { - if (variant.model) ids.add(variant.model); - } - return [...ids]; -} +export { accountHasModel, accountModelIds } from "./accountModelCatalog"; /** * Pure utility: build a lookup from model ID → account info @@ -52,9 +16,7 @@ export function buildAccountLookup( ): Map { const lookup = new Map(); for (const account of accounts) { - if (account.status !== "ready") continue; - for (const modelId of accountModelIds(account)) { - if (!modelId || !accountHasModel(account, modelId)) continue; + for (const modelId of selectableAccountModelIds(account)) { const existing = lookup.get(modelId); if (existing) { existing.totalKeys += 1; diff --git a/src/hooks/models/useModelEffortSegment.ts b/src/hooks/models/useModelEffortSegment.ts index 1815dc14a1..89bb8d5b27 100644 --- a/src/hooks/models/useModelEffortSegment.ts +++ b/src/hooks/models/useModelEffortSegment.ts @@ -7,11 +7,11 @@ import { buildGroupByModel } from "@src/scaffold/GlobalSpotlight/palettes/Unifie import type { LastModelSelection } from "@src/store/session/creatorDefaultModelAtom"; import { formatReasoningLevel, - parseModelVariant, - resolveModelVariantFields, + toModelReasoningLevel, } from "@src/util/modelVariants"; import { buildVariantEditOptions } from "@src/util/variantEditOptions"; +import { resolveAccountModelVariant } from "./accountModelCatalog"; import { resolveModelDisplaySelection } from "./resolveModelDisplaySelection"; import { accountHasModel, @@ -61,7 +61,7 @@ export function useModelEffortSegment({ }; } - const groupByModel = buildGroupByModel(accountLookup.keys()); + const groupByModel = buildGroupByModel(accountLookup.keys(), accounts); const family = groupByModel.get(modelId) ?? [modelId]; const selectedAccount = accounts.find((account) => { @@ -95,12 +95,20 @@ export function useModelEffortSegment({ }; }, [accountLookup, accounts, displaySelection, isHosted, modelId, onApply]); + const variantMetadata = useMemo(() => { + const account = accounts.find((entry) => entry.id === selectedAccountId); + return ( + groupModelIds.length > 0 ? groupModelIds : modelId ? [modelId] : [] + ).map((model) => resolveAccountModelVariant(account, model)); + }, [accounts, groupModelIds, modelId, selectedAccountId]); + const variantOptions = useMemo( () => buildVariantEditOptions( - groupModelIds.length > 0 ? groupModelIds : modelId ? [modelId] : [] + groupModelIds.length > 0 ? groupModelIds : modelId ? [modelId] : [], + variantMetadata ), - [groupModelIds, modelId] + [groupModelIds, modelId, variantMetadata] ); const effectiveModelId = modelId @@ -109,13 +117,15 @@ export function useModelEffortSegment({ ) ?? modelId) : undefined; const variant = effectiveModelId - ? parseModelVariant(effectiveModelId) + ? variantMetadata.find((entry) => entry.model === effectiveModelId) : undefined; const effortLabel = useMemo(() => { const parts: string[] = []; if (variant?.reasoning) { - parts.push(formatReasoningLevel(variant.reasoning)); + parts.push( + formatReasoningLevel(toModelReasoningLevel(variant.reasoning)) + ); } if (variant?.fast) { parts.push("Fast"); @@ -135,17 +145,18 @@ export function useModelEffortSegment({ const account = accounts.find((entry) => entry.id === selectedAccountId); if (!account) return; - const baseModel = resolveModelVariantFields(nextModelId).base_model; - const nextDefaults = (account.defaultVariants ?? []).filter( - (entry) => entry.base_model !== baseModel - ); - nextDefaults.push({ base_model: baseModel, model: nextModelId }); + const baseModel = resolveAccountModelVariant( + account, + nextModelId + ).base_model; - void saveKey({ + saveKey({ id: account.id, agent_type: account.modelType, - default_variants: nextDefaults, - }); + default_variant_overrides: [ + { base_model: baseModel, model: nextModelId }, + ], + }).catch(() => undefined); }, [accounts, saveKey, selectedAccountId] ); diff --git a/src/hooks/models/useValidatedLastPair.ts b/src/hooks/models/useValidatedLastPair.ts index 1e5386754f..b01541bd0f 100644 --- a/src/hooks/models/useValidatedLastPair.ts +++ b/src/hooks/models/useValidatedLastPair.ts @@ -26,11 +26,8 @@ import { } from "@src/features/MarketConnect/marketProfiles"; import { parseAppliedMarketSelection } from "@src/features/MarketConnect/marketSelection"; import { useKeyVault } from "@src/hooks/keyVault"; -import { withNativeHarnessModels } from "@src/hooks/models/nativeHarnessAccountModels"; -import { - getCliCompatibleAccounts, - useAgentCompatibility, -} from "@src/hooks/models/useAgentCompatibility"; +import { getModelPickerAccounts } from "@src/hooks/models/accountModelCatalog"; +import { useAgentCompatibility } from "@src/hooks/models/useAgentCompatibility"; import { type LastModelSelection, creatorDefaultModelPairAtom, @@ -58,12 +55,16 @@ export function useValidatedLastPair(): LastModelSelection | null { const { accounts: allAccounts } = useKeyVault({ autoLoad: true }); - const accounts = useMemo(() => { - if (dispatchCategory === "cli_agent" && cliAgentType) { - return getCliCompatibleAccounts(registry, cliAgentType, allAccounts); - } - return withNativeHarnessModels(allAccounts, dispatchCategory); - }, [dispatchCategory, cliAgentType, allAccounts, registry]); + const accounts = useMemo( + () => + getModelPickerAccounts( + registry, + allAccounts, + dispatchCategory, + cliAgentType + ), + [dispatchCategory, cliAgentType, allAccounts, registry] + ); // Only fetch ORGII pool config when the stored pair actually needs it: // hosted_key sessions or ORGII tier model IDs (orgii:*). Own-key pairs diff --git a/src/hooks/models/withNativeHarnessModels.test.ts b/src/hooks/models/withNativeHarnessModels.test.ts index 077403d2ba..e8fc30a3a1 100644 --- a/src/hooks/models/withNativeHarnessModels.test.ts +++ b/src/hooks/models/withNativeHarnessModels.test.ts @@ -76,12 +76,12 @@ describe("withCodexOAuthModels — model list population", () => { expect(enriched.availableModels ?? []).toContain("custom-codex-model"); }); - it("sets status to ready", () => { + it("preserves account health", () => { const account = baseAccount({ - status: "loading" as KeyVaultAccount["status"], + status: "error" as KeyVaultAccount["status"], }); const enriched = withCodexOAuthModels(account); - expect(enriched.status).toBe("ready"); + expect(enriched.status).toBe("error"); }); }); diff --git a/src/hooks/session/__tests__/sessionPatchQueue.test.ts b/src/hooks/session/__tests__/sessionPatchQueue.test.ts new file mode 100644 index 0000000000..56b892136f --- /dev/null +++ b/src/hooks/session/__tests__/sessionPatchQueue.test.ts @@ -0,0 +1,229 @@ +import { describe, expect, it, vi } from "vitest"; + +import type { Session } from "@src/store/session/sessionAtom/types"; + +import { createSessionPatchQueue } from "../sessionPatchQueue"; + +function deferred() { + let resolve!: () => void; + let reject!: (error: Error) => void; + const promise = new Promise((yes, no) => { + resolve = yes; + reject = no; + }); + return { promise, resolve, reject }; +} + +function setup() { + let current: Session | undefined = { + session_id: "session", + created_at: "2026-09-23", + updated_at: "2026-09-23", + status: "idle", + model: "original", + accountId: "a", + }; + const writes: ReturnType[] = []; + const persist = vi.fn(() => { + const operation = deferred(); + writes.push(operation); + return operation.promise; + }); + const watchers = new Set<() => void>(); + const patch = createSessionPatchQueue({ + read: () => current, + publish: (row) => { + current = row; + watchers.forEach((changed) => changed()); + }, + persist, + subscribe: (_id, changed) => { + watchers.add(changed); + return () => { + watchers.delete(changed); + }; + }, + }); + return { + patch, + persist, + writes, + watchers, + get current() { + return current; + }, + set current(row) { + current = row; + watchers.forEach((changed) => changed()); + }, + }; +} + +describe("ordered session mutation boundary", () => { + it("keeps rapid model/account picks optimistic but persists them in order", async () => { + const state = setup(); + const first = state.patch("session", { model: "one", accountId: "b" }); + const second = state.patch("session", { model: "two", accountId: "c" }); + expect(state.current).toMatchObject({ model: "two", accountId: "c" }); + expect(state.persist).toHaveBeenCalledTimes(1); + state.writes[0].resolve(); + await first; + expect(state.persist).toHaveBeenNthCalledWith(2, "session", { + model: "two", + accountId: "c", + }); + expect(state.current?.model).toBe("two"); + state.writes[1].resolve(); + await second; + expect(state.current).toMatchObject({ model: "two", accountId: "c" }); + }); + + it("two failures restore confirmed identity without undoing an unrelated update", async () => { + const state = setup(); + const first = state.patch("session", { model: "one", accountId: "b" }); + const firstRejected = expect(first).rejects.toThrow("first"); + const second = state.patch("session", { model: "two", accountId: "c" }); + const secondRejected = expect(second).rejects.toThrow("second"); + state.current = { ...state.current!, name: "renamed elsewhere" }; + state.writes[0].reject(new Error("first")); + await firstRejected; + expect(state.current?.model).toBe("two"); + state.writes[1].reject(new Error("second")); + await secondRejected; + expect(state.current).toMatchObject({ + model: "original", + accountId: "a", + name: "renamed elsewhere", + }); + const retry = state.patch("session", { model: "retry", accountId: "d" }); + state.writes[2].resolve(); + await retry; + expect(state.current).toMatchObject({ model: "retry", accountId: "d" }); + }); + + it("rolls a later failure back to the previous successful pair", async () => { + const state = setup(); + const first = state.patch("session", { model: "one", accountId: "b" }); + const second = state.patch("session", { model: "two", accountId: "c" }); + const rejected = expect(second).rejects.toThrow("rejected"); + state.writes[0].resolve(); + await first; + state.writes[1].reject(new Error("rejected")); + await rejected; + expect(state.current).toMatchObject({ model: "one", accountId: "b" }); + }); + + it("preserves an external selection observed between two rejected local picks", async () => { + const state = setup(); + const first = state.patch("session", { model: "one", accountId: "b" }); + const firstRejected = expect(first).rejects.toThrow("first"); + state.current = { ...state.current!, model: "external", accountId: "x" }; + const second = state.patch("session", { model: "two", accountId: "c" }); + const secondRejected = expect(second).rejects.toThrow("second"); + state.writes[0].reject(new Error("first")); + await firstRejected; + expect(state.current).toMatchObject({ model: "two", accountId: "c" }); + state.writes[1].reject(new Error("second")); + await secondRejected; + expect(state.current).toMatchObject({ model: "external", accountId: "x" }); + expect(state.watchers.size).toBe(0); + }); + + it("replays the latest pick immediately over an earlier RPC account-switch echo", async () => { + const state = setup(); + const first = state.patch("session", { model: "one", accountId: "b" }); + const second = state.patch("session", { model: "two", accountId: "c" }); + state.current = { ...state.current!, model: "one", accountId: "b" }; + expect(state.current).toMatchObject({ model: "two", accountId: "c" }); + state.writes[0].resolve(); + await first; + expect(state.current).toMatchObject({ model: "two", accountId: "c" }); + state.writes[1].resolve(); + await second; + expect(state.current).toMatchObject({ model: "two", accountId: "c" }); + expect(state.watchers.size).toBe(0); + }); + + it("never splits an external model/account pair that shares a pending model", async () => { + const state = setup(); + const pending = state.patch("session", { model: "one", accountId: "b" }); + const rejected = expect(pending).rejects.toThrow("failed"); + state.current = { ...state.current!, model: "one", accountId: "external" }; + state.writes[0].reject(new Error("failed")); + await rejected; + expect(state.current).toMatchObject({ + model: "one", + accountId: "external", + }); + expect(state.watchers.size).toBe(0); + }); + + it("keeps an external product/exec mode pair atomic on failure", async () => { + const state = setup(); + state.current = { + ...state.current!, + productMode: "ask", + agentExecMode: "ask", + }; + const pending = state.patch("session", { + productMode: "project", + agentExecMode: "build", + }); + const rejected = expect(pending).rejects.toThrow("failed"); + state.current = { + ...state.current!, + productMode: "build", + agentExecMode: "build", + }; + state.writes[0].reject(new Error("failed")); + await rejected; + expect(state.current).toMatchObject({ + productMode: "build", + agentExecMode: "build", + }); + }); + + it("does not resurrect a deleted session or overwrite a newer external identity", async () => { + const state = setup(); + const first = state.patch("session", { model: "one" }); + state.current = { ...state.current!, model: "external" }; + state.writes[0].resolve(); + await first; + expect(state.current?.model).toBe("external"); + const second = state.patch("session", { model: "two" }); + state.current = undefined; + state.writes[1].resolve(); + await second; + expect(state.current).toBeUndefined(); + }); + + it("preserves clear-draft semantics and avoids writes for unchanged fields", async () => { + const state = setup(); + await state.patch("session", { model: "original" }); + expect(state.persist).not.toHaveBeenCalled(); + state.current = { ...state.current!, draftText: "draft" }; + const clear = state.patch("session", { draftText: null }); + expect(state.current?.draftText).toBeUndefined(); + expect(state.persist).toHaveBeenCalledWith("session", { draftText: null }); + state.writes[0].resolve(); + await clear; + }); + + it("bounds pending work and releases the subscription after draining", async () => { + const state = setup(); + const pending = Array.from({ length: 64 }, (_, index) => + state.patch("session", { model: `model-${index}`, accountId: "a" }) + ); + await expect(state.patch("session", { model: "overflow" })).rejects.toThrow( + "Too many pending session changes" + ); + expect(state.current?.model).toBe("model-63"); + expect(state.watchers.size).toBe(1); + for (let index = 0; index < pending.length; index += 1) { + state.writes[index].resolve(); + await pending[index]; + } + expect(state.persist).toHaveBeenCalledTimes(64); + expect(state.watchers.size).toBe(0); + }); +}); diff --git a/src/hooks/session/sessionPatchQueue.ts b/src/hooks/session/sessionPatchQueue.ts new file mode 100644 index 0000000000..e22bab6688 --- /dev/null +++ b/src/hooks/session/sessionPatchQueue.ts @@ -0,0 +1,187 @@ +import type { Session } from "@src/store/session/sessionAtom/types"; + +export interface SessionPatchOptions { + name?: string; + model?: string; + accountId?: string; + agentExecMode?: string; + productMode?: string; + draftText?: string | null; + replyTargetEventId?: string | null; + pinned?: boolean; +} + +type PatchFields = Partial>; +// These pairs are domain identities. Observing or rolling back only half of a +// pair can create a combination that neither the user nor the server selected. +const fieldGroups: readonly (readonly (keyof SessionPatchOptions)[])[] = [ + ["model", "accountId"], + ["productMode", "agentExecMode"], + ["name"], + ["draftText"], + ["replyTargetEventId"], + ["pinned"], +]; +interface PendingPatch { + options: SessionPatchOptions; + fields: PatchFields; + resolve: () => void; + reject: (error: unknown) => void; +} +interface PendingSession { + confirmed: Session; + projected: Session; + owned: Set; + pending: PendingPatch[]; + unsubscribe?: () => void; +} + +function patchFields(options: SessionPatchOptions): PatchFields { + return Object.fromEntries( + Object.entries(options) + .filter(([, value]) => value !== undefined) + .map(([key, value]) => [ + key, + value === null || + ((key === "draftText" || key === "replyTargetEventId") && value === "") + ? undefined + : value, + ]) + ) as PatchFields; +} + +/** + * One ordered write stream per session. Pending intent stays optimistic, but + * rejected writes are replayed from the last confirmed fields, never from a + * newer optimistic snapshot. The map contains in-flight work only. + */ +export function createSessionPatchQueue(owner: { + read: (id: string) => Session | undefined; + publish: (session: Session) => void; + persist: (id: string, patch: SessionPatchOptions) => Promise; + subscribe?: (id: string, changed: () => void) => () => void; +}) { + const sessions = new Map(); + + function observe(state: PendingSession, current: Session) { + for (const group of fieldGroups) { + if (!group.some((key) => state.owned.has(key))) continue; + const matches = (row: Session) => + group.every((key) => current[key] === row[key]); + if (matches(state.projected)) continue; + // RPC account-switch events arrive before the RPC resolves. Their + // earlier accepted pair must not hide a later optimistic pick. + let candidate = state.confirmed; + if (matches(candidate)) continue; + const expectedEcho = state.pending.some((operation) => { + candidate = { ...candidate, ...operation.fields }; + return matches(candidate); + }); + if (expectedEcho) continue; + for (const key of group) { + state.confirmed = { ...state.confirmed, [key]: current[key] }; + state.owned.delete(key); + } + } + } + + function reconcile(id: string, state: PendingSession) { + const current = owner.read(id); + if (!current) return; // Never resurrect a deleted session. + observe(state, current); + const projected = state.pending.reduce( + (row, operation) => ({ ...row, ...operation.fields }), + state.confirmed + ); + const fields = Object.fromEntries( + [...state.owned].map((key) => [key, projected[key]]) + ); + state.projected = { ...current, ...fields }; + if ( + Object.entries(fields).some( + ([key, value]) => current[key as keyof Session] !== value + ) + ) { + owner.publish(state.projected); + } + } + + async function drain(id: string, state: PendingSession) { + while (state.pending.length > 0) { + const operation = state.pending[0]; + try { + if (!owner.read(id)) throw new Error(`Session ${id} no longer exists`); + if (!state.confirmed.importedFrom) { + await owner.persist(id, operation.options); + } + state.confirmed = { ...state.confirmed, ...operation.fields }; + state.pending.shift(); + reconcile(id, state); + operation.resolve(); + } catch (error) { + state.pending.shift(); + reconcile(id, state); + operation.reject(error); + } + } + state.unsubscribe?.(); + sessions.delete(id); + } + + return (id: string, options: SessionPatchOptions): Promise => { + const before = owner.read(id); + if (!before) + return Promise.reject(new Error(`Session ${id} not in local store`)); + const fields = patchFields(options); + let state = sessions.get(id); + if ( + !state && + Object.entries(fields).every( + ([key, value]) => before[key as keyof Session] === value + ) + ) + return Promise.resolve(); + if (state && state.pending.length >= 64) { + return Promise.reject(new Error("Too many pending session changes")); + } + const start = !state; + if (!state) { + state = { + confirmed: before, + projected: before, + owned: new Set(), + pending: [], + }; + sessions.set(id, state); + } else { + observe(state, before); + } + const promise = new Promise((resolve, reject) => { + state.pending.push({ options, fields, resolve, reject }); + }); + for (const group of fieldGroups) { + if ( + !group.some((key) => Object.prototype.hasOwnProperty.call(fields, key)) + ) + continue; + for (const key of group) { + if (!state.owned.has(key)) { + state.confirmed = { ...state.confirmed, [key]: before[key] }; + } + state.owned.add(key); + } + } + state.projected = { ...before, ...fields }; + if (start) { + const active = state; + state.unsubscribe = owner.subscribe?.(id, () => reconcile(id, active)); + } + owner.publish(state.projected); + if (start) + drain(id, state).catch((error) => { + for (const operation of state.pending.splice(0)) + operation.reject(error); + }); + return promise; + }; +} diff --git a/src/hooks/session/useSessionPatch.ts b/src/hooks/session/useSessionPatch.ts index 2e47839602..416195b929 100644 --- a/src/hooks/session/useSessionPatch.ts +++ b/src/hooks/session/useSessionPatch.ts @@ -33,89 +33,53 @@ import { useAtomValue } from "jotai"; import { useCallback, useEffect, useRef, useState } from "react"; import { rpc } from "@src/api/tauri/rpc"; +import { sessionByIdAtom, upsertSession } from "@src/store/session"; import { - type Session, - sessionByIdAtom, - upsertSession, -} from "@src/store/session"; + getInstrumentedStore, + isStoreInitialized, +} from "@src/util/core/state/instrumentedStore"; -interface PatchOptions { - name?: string; - model?: string; - accountId?: string; - agentExecMode?: string; - /** - * Persistent product mode (orgtrack/v1 §5.2): build|plan|ask|project. - * Validated as a closed enum on the Rust side; agent sessions only. - */ - productMode?: string; - /** - * Three-state per-session draft text (P3): - * undefined → leave column alone - * null → clear the draft (composer was emptied / message sent) - * string → set the draft to this value - * Mirrors the Rust `Option>` deserialize on - * `SessionPatch::draft_text`. - */ - draftText?: string | null; - /** Three-state reply target event id (P3). Same semantics as `draftText`. */ - replyTargetEventId?: string | null; - /** Pin toggle (P5). Absent = leave alone. */ - pinned?: boolean; -} +import { + type SessionPatchOptions, + createSessionPatchQueue, +} from "./sessionPatchQueue"; + +type PatchOptions = SessionPatchOptions; interface PatchState { isPatching: boolean; error: string | null; } -function normalizedOptionalText( - value: string | null | undefined -): string | undefined { - return value == null || value === "" ? undefined : value; -} - -function patchWouldChangeSession( - before: Session, - options: PatchOptions -): boolean { - if (options.name !== undefined && before.name !== options.name) return true; - if (options.model !== undefined && before.model !== options.model) - return true; - if (options.accountId !== undefined && before.accountId !== options.accountId) - return true; - if ( - options.agentExecMode !== undefined && - before.agentExecMode !== options.agentExecMode - ) - return true; - if ( - options.productMode !== undefined && - before.productMode !== options.productMode - ) - return true; - if ( - options.draftText !== undefined && - before.draftText !== normalizedOptionalText(options.draftText) - ) - return true; - if ( - options.replyTargetEventId !== undefined && - before.replyTargetEventId !== - normalizedOptionalText(options.replyTargetEventId) - ) - return true; - if (options.pinned !== undefined && before.pinned !== options.pinned) - return true; - return false; +const patchQueues = new WeakMap< + ReturnType, + ReturnType +>(); + +function sessionPatchQueue() { + const store = getInstrumentedStore(); + let queue = patchQueues.get(store); + if (!queue) { + queue = createSessionPatchQueue({ + read: (id) => + isStoreInitialized() && getInstrumentedStore() === store + ? store.get(sessionByIdAtom(id)) + : undefined, + publish: (session) => { + if (isStoreInitialized() && getInstrumentedStore() === store) + upsertSession(session); + }, + subscribe: (id, changed) => store.sub(sessionByIdAtom(id), changed), + persist: async (sessionId, patch) => { + await rpc.sessionAggregate.patch({ sessionId, patch }); + }, + }); + patchQueues.set(store, queue); + } + return queue; } -/** - * Low-level patch primitive: optimistic write → RPC → rollback on error. - * - * Returns a stable function `(sessionId, patch) => Promise` plus - * the in-flight / error state for the most recent call. - */ +/** Shared ordered persistence; only the latest call controls hook feedback. */ function usePatchSession(): { patch: (sessionId: string, patch: PatchOptions) => Promise; } & PatchState { @@ -123,102 +87,34 @@ function usePatchSession(): { isPatching: false, error: null, }); + const generation = useRef(0); + useEffect( + () => () => { + generation.current += 1; + }, + [] + ); const patch = useCallback( - async (sessionId: string, options: PatchOptions): Promise => { + async (sessionId: string, options: PatchOptions) => { + const current = ++generation.current; setState({ isPatching: true, error: null }); - // Snapshot the prior values BEFORE the optimistic write so we - // can restore them on error. Reading via the instrumented store - // avoids a stale-closure issue if the same hook instance fires - // back-to-back patches for the same session. - const { getInstrumentedStore } = - await import("@src/util/core/state/instrumentedStore"); - const store = getInstrumentedStore(); - const before = store.get(sessionByIdAtom(sessionId)) as - | Session - | undefined; - - if (!before) { - // Session not in the local store — bail out before we send a - // patch the backend would just reject with "not found". This - // typically means the session was deleted while the user had - // a stale pill open. - const message = `useSessionPatch: session ${sessionId} not in local store`; - setState({ isPatching: false, error: message }); - throw new Error(message); - } - - if (!patchWouldChangeSession(before, options)) { - setState({ isPatching: false, error: null }); - return; - } - - const optimistic: Session = { ...before }; - if (options.name !== undefined) optimistic.name = options.name; - if (options.model !== undefined) optimistic.model = options.model; - if (options.accountId !== undefined) - optimistic.accountId = options.accountId; - if (options.agentExecMode !== undefined) - optimistic.agentExecMode = options.agentExecMode; - if (options.productMode !== undefined) - optimistic.productMode = options.productMode; - // Three-state fields: `null` clears (write `undefined` into the - // optimistic session, since the Session type uses `undefined` for - // "no value"); a string sets; a property left absent on `options` - // means "don't touch it". - if (options.draftText !== undefined) - optimistic.draftText = options.draftText ?? undefined; - if (options.replyTargetEventId !== undefined) - optimistic.replyTargetEventId = options.replyTargetEventId ?? undefined; - if (options.pinned !== undefined) optimistic.pinned = options.pinned; - upsertSession(optimistic); - - // Imported teammate copies (`Session.importedFrom`) exist ONLY in the - // TS store — there is no agent_sessions/code_sessions row for the - // backend to patch, so the RPC would reject with "session not found" - // (surfacing as a full-screen App error from e.g. the composer's - // debounced draft save or the model picker). The optimistic local - // write above IS the persistence these rows get; skip the RPC. - if (before.importedFrom) { - setState({ isPatching: false, error: null }); - return; - } - try { - await rpc.sessionAggregate.patch({ - sessionId, - patch: { - name: options.name, - model: options.model, - accountId: options.accountId, - agentExecMode: options.agentExecMode, - productMode: options.productMode, - // Forward the tri-state values verbatim. zod's - // `.nullable().optional()` lines up with the Rust double- - // Option deserialize: undefined skips, null clears, string - // sets. - draftText: options.draftText, - replyTargetEventId: options.replyTargetEventId, - pinned: options.pinned, - }, - }); - setState({ isPatching: false, error: null }); - } catch (err) { - // Roll back to the snapshot taken above. We re-write the full - // prior session record (not just the touched fields) so a - // partial backend success — which the Rust handler currently - // can't produce, but a future split write could — wouldn't - // leave the UI in a hybrid state. - upsertSession(before); - const message = - err instanceof Error ? err.message : String(err ?? "patch failed"); - setState({ isPatching: false, error: message }); - throw err; + await sessionPatchQueue()(sessionId, options); + if (current === generation.current) + setState({ isPatching: false, error: null }); + } catch (error) { + if (current === generation.current) { + setState({ + isPatching: false, + error: error instanceof Error ? error.message : String(error), + }); + } + throw error; } }, [] ); - return { patch, ...state }; } @@ -228,7 +124,7 @@ function usePatchSession(): { * Returns the current values (from `sessionByIdAtom`) plus a * `setModel` function that performs an atomic backend patch. * - * Pass `accountId: null` to leave it unchanged when only the model + * Omit `accountId` to leave it unchanged when only the model * name changes (e.g. switching between two Anthropic models on the * same key). */ diff --git a/src/modules/MainApp/AgentOrgs/config/shared/ModelPicker.tsx b/src/modules/MainApp/AgentOrgs/config/shared/ModelPicker.tsx index 7dd02ab8b4..0bc2ac1e04 100644 --- a/src/modules/MainApp/AgentOrgs/config/shared/ModelPicker.tsx +++ b/src/modules/MainApp/AgentOrgs/config/shared/ModelPicker.tsx @@ -13,10 +13,10 @@ import ModelIcon from "@src/components/ModelIcon"; import Select, { type SelectOptionGroup } from "@src/components/Select"; import { buildAccountLookup, - getRustCompatibleAccounts, useAgentCompatibility, useModelAccountLookup, } from "@src/hooks/models"; +import { getModelPickerAccounts } from "@src/hooks/models/accountModelCatalog"; import { formatModelNameFull } from "@src/util/formatModelName"; interface ModelPickerProps { @@ -43,7 +43,7 @@ const ModelPicker: React.FC = ({ const { registry } = useAgentCompatibility(); const accounts = useMemo( - () => getRustCompatibleAccounts(registry, allAccounts), + () => getModelPickerAccounts(registry, allAccounts, "rust_agent"), [registry, allAccounts] ); diff --git a/src/modules/MainApp/Integrations/KeyVault/Table/AccountsTable.tsx b/src/modules/MainApp/Integrations/KeyVault/Table/AccountsTable.tsx index d8003cda3d..6f8c816222 100644 --- a/src/modules/MainApp/Integrations/KeyVault/Table/AccountsTable.tsx +++ b/src/modules/MainApp/Integrations/KeyVault/Table/AccountsTable.tsx @@ -251,7 +251,6 @@ export const AccountsTable: React.FC = ({ return saveKey({ id: account.id, agent_type: account.modelType, - available_models: account.availableModels ?? [], enabled_models: [...enabledModels], }); }) diff --git a/src/modules/MainApp/Integrations/KeyVault/Table/useDefaultVariantSaves.test.ts b/src/modules/MainApp/Integrations/KeyVault/Table/useDefaultVariantSaves.test.ts new file mode 100644 index 0000000000..37aba610a3 --- /dev/null +++ b/src/modules/MainApp/Integrations/KeyVault/Table/useDefaultVariantSaves.test.ts @@ -0,0 +1,226 @@ +// @vitest-environment jsdom +import React, { act, useEffect, useSyncExternalStore } from "react"; +import { createRoot } from "react-dom/client"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; + +import type { KeyInfo, SaveKeyRequest } from "@src/api/types/keys"; +import type { KeyVaultAccount } from "@src/hooks/keyVault"; +import { saveDefaultVariantOverrides } from "@src/hooks/keyVault/defaultVariantSaveCoordinator"; +import { + getSharedLocalKeys, + publishSharedLocalKeys, + subscribeSharedLocalKeys, +} from "@src/hooks/keyVault/sharedLocalKeyStore"; + +import { useDefaultVariantSaves } from "./useDefaultVariantSaves"; + +const fixture = vi.hoisted(() => ({ + saveKey: vi.fn<(request: SaveKeyRequest) => Promise>(), +})); +vi.mock("@src/api/services/keyValidation", () => ({ + saveKey: fixture.saveKey, +})); + +const account: KeyVaultAccount = { + id: "account", + name: "Fixture", + modelType: "codex", + status: "ready", + enabled: true, + hasLocalKey: true, + isListed: false, + hasKey: true, + hasApiKey: false, + hasSessionToken: true, + defaultVariants: [ + { base_model: "family-a", model: "a-provider-default" }, + { base_model: "family-b", model: "b-provider-default" }, + ], +}; +const cleanups: Array<() => void> = []; +function setup() { + const onRefresh = vi.fn().mockResolvedValue(undefined); + const resultRef: { current?: ReturnType } = {}; + function Probe() { + const keys = useSyncExternalStore( + subscribeSharedLocalKeys, + getSharedLocalKeys + ); + const result = useDefaultVariantSaves({ + accounts: keys.map((key) => ({ + ...account, + defaultVariants: key.default_variants, + })), + onRefresh, + }); + useEffect(() => { + resultRef.current = result; + }, [result]); + return null; + } + const root = createRoot(document.createElement("div")); + act(() => root.render(React.createElement(Probe))); + let mounted = true; + const unmount = () => { + if (mounted) { + act(() => root.unmount()); + mounted = false; + } + }; + cleanups.push(unmount); + if (!resultRef.current) throw new Error("Hook did not render"); + return { hook: resultRef.current, onRefresh, unmount }; +} +beforeEach(() => { + Object.assign(globalThis, { IS_REACT_ACT_ENVIRONMENT: true }); + vi.useFakeTimers(); + vi.clearAllMocks(); + publishSharedLocalKeys([ + { + id: "account", + agent_type: "codex", + default_variants: account.defaultVariants, + } as KeyInfo, + ]); + fixture.saveKey.mockImplementation( + async (request) => + ({ + id: "account", + agent_type: "codex", + default_variants: request.default_variant_overrides, + }) as KeyInfo + ); +}); +afterEach(() => { + cleanups.splice(0).forEach((cleanup) => cleanup()); + vi.useRealTimers(); +}); + +describe("family-scoped default variant persistence", () => { + it("does not serialize unedited provider defaults as user overrides", async () => { + const { hook } = setup(); + act(() => + hook.updateDefaultVariant("account", "family-a", "a-user-choice") + ); + await act(async () => { + await vi.advanceTimersByTimeAsync(300); + }); + expect(fixture.saveKey).toHaveBeenCalledExactlyOnceWith({ + id: "account", + agent_type: "codex", + default_variant_overrides: [ + { base_model: "family-a", model: "a-user-choice" }, + ], + }); + expect(getSharedLocalKeys()[0].default_variants).toContainEqual({ + base_model: "family-a", + model: "a-user-choice", + }); + }); + + it("coalesces repeated picks and includes only explicitly edited families", async () => { + const { hook } = setup(); + act(() => { + hook.updateDefaultVariant("account", "family-a", "earlier"); + hook.updateDefaultVariant("account", "family-a", "latest"); + hook.updateDefaultVariant("account", "family-c", "new-choice"); + }); + await act(async () => { + await vi.advanceTimersByTimeAsync(300); + }); + expect(fixture.saveKey).toHaveBeenCalledExactlyOnceWith({ + id: "account", + agent_type: "codex", + default_variant_overrides: [ + { base_model: "family-a", model: "latest" }, + { base_model: "family-c", model: "new-choice" }, + ], + }); + }); + + it("reloads authoritative data after failure and permits a fresh family-only retry", async () => { + fixture.saveKey.mockRejectedValueOnce(new Error("write failed")); + const { hook, onRefresh } = setup(); + act(() => + hook.updateDefaultVariant("account", "family-a", "failed-choice") + ); + await act(async () => { + await vi.advanceTimersByTimeAsync(300); + }); + expect(onRefresh).toHaveBeenCalledTimes(1); + expect(getSharedLocalKeys()[0].default_variants).toEqual( + account.defaultVariants + ); + act(() => hook.updateDefaultVariant("account", "family-b", "retry-choice")); + await act(async () => { + await vi.advanceTimersByTimeAsync(300); + }); + expect(fixture.saveKey).toHaveBeenLastCalledWith({ + id: "account", + agent_type: "codex", + default_variant_overrides: [ + { base_model: "family-b", model: "retry-choice" }, + ], + }); + }); + + it("unmount flush joins the same in-flight writer and survives table disposal", async () => { + let resolveFirst!: (key: KeyInfo) => void; + let resolveSecond!: (key: KeyInfo) => void; + fixture.saveKey + .mockImplementationOnce( + () => + new Promise((resolve) => { + resolveFirst = resolve; + }) + ) + .mockImplementationOnce( + () => + new Promise((resolve) => { + resolveSecond = resolve; + }) + ); + const { hook, unmount } = setup(); + let first!: Promise; + act(() => { + first = saveDefaultVariantOverrides({ + id: "account", + agent_type: "codex", + default_variant_overrides: [{ base_model: "family-a", model: "first" }], + }); + hook.updateDefaultVariant("account", "family-a", "last-before-close"); + }); + // No debounce elapsed; leaving the table hands its last click to the + // same coordinator already serving the other picker surface. + unmount(); + expect(vi.getTimerCount()).toBe(0); + expect(fixture.saveKey).toHaveBeenCalledTimes(1); + expect(getSharedLocalKeys()[0].default_variants).toContainEqual({ + base_model: "family-a", + model: "last-before-close", + }); + resolveFirst({ + id: "account", + agent_type: "codex", + default_variants: [{ base_model: "family-a", model: "first" }], + } as KeyInfo); + await first; + expect(fixture.saveKey).toHaveBeenCalledTimes(2); + expect(getSharedLocalKeys()[0].default_variants).toContainEqual({ + base_model: "family-a", + model: "last-before-close", + }); + resolveSecond({ + id: "account", + agent_type: "codex", + default_variants: [ + { base_model: "family-a", model: "last-before-close" }, + ], + } as KeyInfo); + await vi.advanceTimersByTimeAsync(0); + expect(getSharedLocalKeys()[0].default_variants).toContainEqual({ + base_model: "family-a", + model: "last-before-close", + }); + }); +}); diff --git a/src/modules/MainApp/Integrations/KeyVault/Table/useDefaultVariantSaves.ts b/src/modules/MainApp/Integrations/KeyVault/Table/useDefaultVariantSaves.ts index bba5314236..b0cd57d9d1 100644 --- a/src/modules/MainApp/Integrations/KeyVault/Table/useDefaultVariantSaves.ts +++ b/src/modules/MainApp/Integrations/KeyVault/Table/useDefaultVariantSaves.ts @@ -9,13 +9,11 @@ */ import { useCallback, useEffect, useRef, useState } from "react"; -import { saveKey } from "@src/api/services/keyValidation"; import type { KeyVaultAccount } from "@src/hooks/keyVault"; -import { upsertSharedLocalKey } from "@src/hooks/keyVault/sharedLocalKeyStore"; +import { saveDefaultVariantOverrides } from "@src/hooks/keyVault/defaultVariantSaveCoordinator"; import { type DefaultVariantOverrides, - applyDefaultVariantOverrides, defaultVariantOverridesSettled, } from "./defaultVariantOverrides"; @@ -49,7 +47,8 @@ export function useDefaultVariantSaves({ >(new Map()); const optimisticRef = useRef>(new Map()); const timerRef = useRef | null>(null); - const queueRef = useRef>(new Set()); + const queueRef = useRef>(new Map()); + const mountedRef = useRef(true); const flushQueue = useCallback(() => { if (timerRef.current) { @@ -57,42 +56,40 @@ export function useDefaultVariantSaves({ timerRef.current = null; } - const queued = [...queueRef.current]; - if (queued.length === 0) return; - queueRef.current = new Set(); - + const queued = queueRef.current; + if (queued.size === 0) return; + queueRef.current = new Map(); const accountById = new Map( accounts.map((account) => [account.id, account]) ); - void Promise.all( - queued.map((accountId) => { - const account = accountById.get(accountId); - const overrides = optimisticRef.current.get(accountId); - if (!account || !overrides) return Promise.resolve(undefined); - return saveKey({ - id: account.id, - agent_type: account.modelType, - default_variants: applyDefaultVariantOverrides( - account.defaultVariants, - overrides - ), - }); - }) - ) - // `saveKey` answers with the stored record, so publishing it settles the - // override. Re-listing every key would tell us nothing new. - .then((savedKeys) => { - for (const saved of savedKeys) { - if (saved) upsertSharedLocalKey(saved); - } - }) - .catch(() => { - const empty = new Map(); - optimisticRef.current = empty; - setOptimisticDefaultVariants(empty); - // The write failed, so the store is the only trustworthy source left. + for (const [accountId, overrides] of queued) { + const account = accountById.get(accountId); + if (!account) continue; + // The shared coordinator owns optimistic state from this point onward, + // including writes handed off while this table is unmounting. + void saveDefaultVariantOverrides({ + id: account.id, + agent_type: account.modelType, + default_variant_overrides: [...overrides].map( + ([base_model, model]) => ({ base_model, model }) + ), + }).catch(() => { void onRefresh?.(); }); + } + // Drop only the local debounce overlay that was handed off. A newer click + // must retain its own overlay until its next debounce flush. + const next = new Map(optimisticRef.current); + for (const [accountId, submitted] of queued) { + const remaining = new Map(next.get(accountId)); + for (const [family, model] of submitted) { + if (remaining.get(family) === model) remaining.delete(family); + } + if (remaining.size) next.set(accountId, remaining); + else next.delete(accountId); + } + optimisticRef.current = next; + if (mountedRef.current) setOptimisticDefaultVariants(next); }, [accounts, onRefresh]); // Keep the latest flush impl in a ref so the unmount cleanup can fire a @@ -102,16 +99,17 @@ export function useDefaultVariantSaves({ flushQueueRef.current = flushQueue; }, [flushQueue]); - useEffect( - () => () => { + useEffect(() => { + mountedRef.current = true; + return () => { + mountedRef.current = false; if (timerRef.current) { clearTimeout(timerRef.current); timerRef.current = null; } flushQueueRef.current(); - }, - [] - ); + }; + }, []); const updateDefaultVariant = useCallback( (accountId: string, baseModel: string, model: string) => { @@ -123,7 +121,10 @@ export function useDefaultVariantSaves({ optimisticRef.current = next; setOptimisticDefaultVariants(next); - queueRef.current.add(accountId); + const queued = + queueRef.current.get(accountId) ?? new Map(); + queued.set(baseModel, model); + queueRef.current.set(accountId, queued); if (timerRef.current) clearTimeout(timerRef.current); timerRef.current = setTimeout(() => { flushQueueRef.current(); diff --git a/src/modules/MainApp/Integrations/KeyVault/hooks/refreshAccountModels.test.ts b/src/modules/MainApp/Integrations/KeyVault/hooks/refreshAccountModels.test.ts new file mode 100644 index 0000000000..f2cd735591 --- /dev/null +++ b/src/modules/MainApp/Integrations/KeyVault/hooks/refreshAccountModels.test.ts @@ -0,0 +1,194 @@ +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { + getFullKey, + getOAuthModelCatalog, + refreshKeyModelCatalog, + refreshOauthToken, + updateKeyHealth, +} from "@src/api/services/keyValidation"; +import type { KeyVaultAccount } from "@src/hooks/keyVault"; + +import { + refreshAccountModels, + refreshAllAccountModels, +} from "./refreshAccountModels"; + +vi.mock("@src/api/services/keyValidation", () => ({ + getFullKey: vi.fn(), + getOAuthModelCatalog: vi.fn(), + refreshKeyModelCatalog: vi.fn(), + refreshOauthToken: vi.fn(), + updateKeyHealth: vi.fn(), + getCursorNativeModels: vi.fn(), + validateKey: vi.fn(), +})); +const account: KeyVaultAccount = { + id: "account", + name: "Fixture account", + status: "ready", + hasLocalKey: true, + isListed: false, + hasKey: true, + hasApiKey: false, + hasSessionToken: true, + enabled: true, + modelType: "codex", + authMethod: "oauth", + availableModels: ["stale-ui-model"], + enabledModels: [], +}; +const snapshot = { + id: "account", + agent_type: "codex", + name: null, + api_key: null, + session_token: "fixture-token", + base_url: null, + env_vars: {}, + account_metadata: {}, + available_models: ["model"], + auth_method: "oauth", + credential_generation: 3, + model_catalog_generation: 7, +} as const; +const catalog = { + models: ["model", "next"], + modelContextLengths: { model: 10000 }, + defaultEnabledModels: ["next"], + modelVariants: [ + { + model: "opaque-effort", + base_model: "model", + reasoning: "high", + fast: false, + }, + ], + defaultVariants: [{ base_model: "model", model: "opaque-effort" }], + source: "live", +} as const; + +beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(getFullKey).mockResolvedValue({ + ...snapshot, + available_models: [...snapshot.available_models], + auth_method: "oauth", + }); + vi.mocked(getOAuthModelCatalog).mockResolvedValue({ + ...catalog, + models: [...catalog.models], + defaultEnabledModels: [...catalog.defaultEnabledModels], + modelVariants: [...catalog.modelVariants], + defaultVariants: [...catalog.defaultVariants], + }); + vi.mocked(refreshKeyModelCatalog).mockResolvedValue({ + available_models: ["model", "next", "manual"], + } as Awaited>); +}); + +describe("refreshAccountModels", () => { + it("commits the full discovered metadata and uses persisted readback without replaying UI enablement", async () => { + await expect(refreshAccountModels(account)).resolves.toEqual({ + models: ["model", "next", "manual"], + previousModels: ["model"], + }); + expect(refreshKeyModelCatalog).toHaveBeenCalledExactlyOnceWith("account", { + expectedCredentialGeneration: 3, + expectedCatalogGeneration: 7, + availableModels: ["model", "next"], + modelVariants: catalog.modelVariants, + defaultVariants: catalog.defaultVariants, + modelContextLengths: { model: 10000 }, + }); + expect(updateKeyHealth).not.toHaveBeenCalled(); + }); + + it("single-flights simultaneous entry points and releases state after success", async () => { + const first = refreshAccountModels(account); + const second = refreshAccountModels(account); + expect(first).toBe(second); + await Promise.all([first, second]); + expect(getOAuthModelCatalog).toHaveBeenCalledTimes(1); + await refreshAccountModels(account); + expect(getOAuthModelCatalog).toHaveBeenCalledTimes(2); + }); + + it("preserves the previous catalog on fallback and empty discovery, then allows retry", async () => { + vi.mocked(getOAuthModelCatalog).mockResolvedValueOnce({ + ...catalog, + models: [], + defaultEnabledModels: [], + modelVariants: [], + defaultVariants: [], + }); + await expect(refreshAccountModels(account)).rejects.toThrow( + "empty model list" + ); + expect(refreshKeyModelCatalog).not.toHaveBeenCalled(); + await refreshAccountModels(account); + expect(refreshKeyModelCatalog).toHaveBeenCalledTimes(1); + }); + + it("does not replace last-known-good data with a bootstrap fallback", async () => { + vi.mocked(getOAuthModelCatalog).mockResolvedValueOnce({ + ...catalog, + source: "fallback", + models: [...catalog.models], + defaultEnabledModels: [...catalog.defaultEnabledModels], + modelVariants: [...catalog.modelVariants], + defaultVariants: [...catalog.defaultVariants], + }); + await expect(refreshAccountModels(account)).rejects.toMatchObject({ + kind: "transient", + }); + expect(refreshKeyModelCatalog).not.toHaveBeenCalled(); + }); + + it("re-reads the credential snapshot after one OAuth retry", async () => { + vi.mocked(getOAuthModelCatalog).mockRejectedValueOnce( + new Error("HTTP 401") + ); + vi.mocked(getFullKey).mockResolvedValueOnce({ + ...snapshot, + available_models: [...snapshot.available_models], + auth_method: "oauth", + }); + vi.mocked(getFullKey).mockResolvedValueOnce({ + ...snapshot, + available_models: [...snapshot.available_models], + auth_method: "oauth", + credential_generation: 4, + }); + await refreshAccountModels(account); + expect(refreshOauthToken).toHaveBeenCalledExactlyOnceWith("account"); + expect(refreshKeyModelCatalog).toHaveBeenCalledWith( + "account", + expect.objectContaining({ expectedCredentialGeneration: 4 }) + ); + }); + + it("does not turn a retry's network failure into an unguarded invalid-account write", async () => { + vi.mocked(getOAuthModelCatalog) + .mockRejectedValueOnce(new Error("HTTP 401")) + .mockRejectedValueOnce(new Error("network unavailable")); + await expect(refreshAccountModels(account)).rejects.toMatchObject({ + kind: "transient", + }); + expect(refreshKeyModelCatalog).not.toHaveBeenCalled(); + expect(updateKeyHealth).not.toHaveBeenCalled(); + await refreshAccountModels(account); + expect(refreshKeyModelCatalog).toHaveBeenCalledTimes(1); + }); + + it("reports a stale commit as failure and preserves other accounts' successes", async () => { + vi.mocked(refreshKeyModelCatalog).mockRejectedValueOnce( + new Error("Account changed") + ); + const summary = await refreshAllAccountModels([ + account, + { ...account, id: "other" }, + ]); + expect(summary).toEqual({ total: 2, failed: 1, added: 2, removed: 0 }); + }); +}); diff --git a/src/modules/MainApp/Integrations/KeyVault/hooks/refreshAccountModels.ts b/src/modules/MainApp/Integrations/KeyVault/hooks/refreshAccountModels.ts index 535ce452c3..49e9837295 100644 --- a/src/modules/MainApp/Integrations/KeyVault/hooks/refreshAccountModels.ts +++ b/src/modules/MainApp/Integrations/KeyVault/hooks/refreshAccountModels.ts @@ -16,21 +16,22 @@ * performs reactively. * * On success, writes the discovered model list back to the key store via - * updateKeyHealth (preserving the existing healthStatus and enabledModels — - * new models default to "addable", never auto-enabled). On hard failure - * (refresh also rejected, list call still failing), flips healthStatus to - * "invalid" so the row reflects that the user needs to re-add the account. + * refreshKeyModelCatalog (preserving the latest health and enabled choices — + * new models default to "addable", never auto-enabled). OAuth refresh health + * remains owned by the backend's credential-generation-checked write path. */ import { + type FullKeyResponse, type ModelContextLengths, getCursorNativeModels, getFullKey, getOAuthModelCatalog, + refreshKeyModelCatalog, refreshOauthToken, - updateKeyHealth, validateKey, } from "@src/api/services/keyValidation"; import { CLI_AGENT } from "@src/api/tauri/rpc/schemas/validation"; +import type { DefaultVariantInfo, ModelVariantInfo } from "@src/api/types/keys"; import type { KeyVaultAccount } from "@src/hooks/keyVault"; /** @@ -65,7 +66,8 @@ function isOAuthAccount(account: KeyVaultAccount): boolean { interface FetchedAccountModels { models: string[]; modelContextLengths: ModelContextLengths; - defaultEnabledModels?: string[]; + modelVariants?: ModelVariantInfo[]; + defaultVariants?: DefaultVariantInfo[]; } async function fetchOAuthCatalogForAccount( @@ -92,21 +94,32 @@ async function fetchOAuthCatalogForAccount( return { models: catalog.models, modelContextLengths: catalog.modelContextLengths, - defaultEnabledModels: catalog.defaultEnabledModels, + modelVariants: catalog.modelVariants, + defaultVariants: catalog.defaultVariants, }; } +interface AccountModelDiscovery extends FetchedAccountModels { + snapshot: FullKeyResponse; +} + async function fetchModelsForAccount( account: KeyVaultAccount -): Promise { - const fullKey = await getFullKey(account.modelType, account.id); - if (!fullKey) { +): Promise { + const snapshot = await getFullKey(account.modelType, account.id); + if (!snapshot) { throw new RefreshModelsError( `Key not found for account ${account.id}`, "transient" ); } + return { ...(await discoverModelsForAccount(account, snapshot)), snapshot }; +} +async function discoverModelsForAccount( + account: KeyVaultAccount, + fullKey: FullKeyResponse +): Promise { switch (account.modelType) { case CLI_AGENT.CURSOR: { const token = fullKey.session_token; @@ -184,12 +197,10 @@ export interface RefreshAccountModelsResult { previousModels: string[]; } -export async function refreshAccountModels( +async function performAccountModelsRefresh( account: KeyVaultAccount ): Promise { - const previousHealth = account.healthStatus ?? "valid"; - const previousModels = account.availableModels ?? []; - let fetched: FetchedAccountModels; + let fetched: AccountModelDiscovery; try { fetched = await fetchModelsForAccount(account); @@ -201,32 +212,21 @@ export async function refreshAccountModels( try { await refreshOauthToken(account.id); } catch (refreshErr) { - // Refresh itself rejected — refresh_token is dead or revoked. Mark - // the account invalid so the row visibly degrades; user needs to - // re-add the account. - await updateKeyHealth( - account.id, - "invalid", - refreshErr instanceof Error ? refreshErr.message : String(refreshErr) - ); + // OAuth refresh already records health at its guarded backend boundary. + // A late catalog request must not mark replacement credentials invalid. throw new RefreshModelsError( refreshErr instanceof Error ? refreshErr.message : String(refreshErr), - "auth_expired" + isUnauthorizedError(refreshErr) ? "auth_expired" : "transient" ); } try { fetched = await fetchModelsForAccount(account); } catch (retryErr) { - await updateKeyHealth( - account.id, - "invalid", - retryErr instanceof Error ? retryErr.message : String(retryErr) - ); throw retryErr instanceof RefreshModelsError ? retryErr : new RefreshModelsError( retryErr instanceof Error ? retryErr.message : String(retryErr), - "auth_expired" + isUnauthorizedError(retryErr) ? "auth_expired" : "transient" ); } } else { @@ -246,34 +246,38 @@ export async function refreshAccountModels( ); } - const refreshedEnabledModels = (() => { - if (!isOAuthAccount(account)) { - return undefined; - } - if ( - account.modelType !== CLI_AGENT.CLAUDE_CODE && - account.modelType !== CLI_AGENT.CODEX - ) { - return undefined; - } - const enabled = new Set(account.enabledModels ?? []); - for (const modelId of fetched.defaultEnabledModels ?? []) { - if (fetched.models.includes(modelId)) enabled.add(modelId); - } - return [...enabled]; - })(); + const saved = await refreshKeyModelCatalog(account.id, { + expectedCredentialGeneration: fetched.snapshot.credential_generation, + expectedCatalogGeneration: fetched.snapshot.model_catalog_generation, + availableModels: fetched.models, + modelVariants: fetched.modelVariants ?? null, + defaultVariants: fetched.defaultVariants ?? null, + modelContextLengths: fetched.modelContextLengths, + }); + if (!saved) { + throw new RefreshModelsError("Account no longer exists", "transient"); + } + return { + models: saved.available_models, + previousModels: fetched.snapshot.available_models, + }; +} - await updateKeyHealth( - account.id, - previousHealth, - undefined, - fetched.models, - refreshedEnabledModels, - undefined, - fetched.modelContextLengths - ); +// Retained only while a user-requested refresh is active; both entry points +// share discovery and release the promise on success or failure. +const pendingRefreshes = new Map>(); - return { models: fetched.models, previousModels }; +export function refreshAccountModels( + account: KeyVaultAccount +): Promise { + const pending = pendingRefreshes.get(account.id); + if (pending) return pending; + const refresh = performAccountModelsRefresh(account).finally(() => { + if (pendingRefreshes.get(account.id) === refresh) + pendingRefreshes.delete(account.id); + }); + pendingRefreshes.set(account.id, refresh); + return refresh; } export interface RefreshAllAccountModelsSummary { diff --git a/src/modules/MainApp/Integrations/KeyVault/shared/ModelTable/ModelVariantInlineCard.test.ts b/src/modules/MainApp/Integrations/KeyVault/shared/ModelTable/ModelVariantInlineCard.test.ts index 523ddcd511..2c4a32e701 100644 --- a/src/modules/MainApp/Integrations/KeyVault/shared/ModelTable/ModelVariantInlineCard.test.ts +++ b/src/modules/MainApp/Integrations/KeyVault/shared/ModelTable/ModelVariantInlineCard.test.ts @@ -28,16 +28,24 @@ vi.mock("@src/components/ModelPropertiesDropdown/EffortSlider", () => ({ describe("ModelVariantInlineCard default persistence", () => { it.each([ - ["o4-mini", "o4"], - ["claude-opus-4-7", "claude-opus-4-7"], + ["o4-mini", "o4-mini", undefined], + ["o4-mini", "o4", "o4"], + ["claude-opus-4-7", "claude-opus-4-7", undefined], ])( - "preserves the %s family key after filtering bare choices", - (base, key) => { + "preserves %s under family key %s after filtering bare choices", + (base, key, catalogBase) => { const onChange = vi.fn(); renderToStaticMarkup( React.createElement(ModelVariantInlineCard, { variants: [base, `${base}-low`, `${base}-medium`, `${base}-high`].map( - (model) => resolveModelVariantFields(model) + (model) => { + const variant = resolveModelVariantFields(model); + // Explicit catalog metadata owns the grouping key. Without it, + // size suffixes remain part of the model's effort family. + return catalogBase + ? { ...variant, base_model: catalogBase } + : variant; + } ), embedded: true, defaultVariantByBaseModel: new Map([[key, `${base}-high`]]), diff --git a/src/modules/MainApp/Integrations/KeyVault/shared/ModelTable/ModelVariantInlineCard.tsx b/src/modules/MainApp/Integrations/KeyVault/shared/ModelTable/ModelVariantInlineCard.tsx index 3e1df31c47..dbc5728867 100644 --- a/src/modules/MainApp/Integrations/KeyVault/shared/ModelTable/ModelVariantInlineCard.tsx +++ b/src/modules/MainApp/Integrations/KeyVault/shared/ModelTable/ModelVariantInlineCard.tsx @@ -73,9 +73,8 @@ export default function ModelVariantInlineCard({ // A card always renders one model family, so the whole family collapses to a // single saved selection. We pick the shortest `base_model` string from the // complete family as the persistence key: filtering selectable efforts must - // not change that key (the o4-mini bare record owns the existing o4 key), - // and the unsuffixed and parsed spellings of a model must resolve to one - // entry. + // not change that key. Explicit catalog grouping wins; without metadata, + // size suffixes such as o4-mini remain part of the effort family key. const canonicalBaseModel = variants.length > 0 ? variants diff --git a/src/modules/MainApp/Integrations/hooks/modelToggle.test.ts b/src/modules/MainApp/Integrations/hooks/modelToggle.test.ts new file mode 100644 index 0000000000..2fa1e9d0d1 --- /dev/null +++ b/src/modules/MainApp/Integrations/hooks/modelToggle.test.ts @@ -0,0 +1,36 @@ +import { describe, expect, it, vi } from "vitest"; + +import { saveKey } from "@src/api/services/keyValidation"; +import type { KeyVaultAccount } from "@src/hooks/keyVault"; + +import { toggleModelForAccounts } from "./modelToggle"; + +vi.mock("@src/api/services/keyValidation", () => ({ saveKey: vi.fn() })); + +describe("model enablement write", () => { + it("changes enablement without replaying a stale available-model catalog", async () => { + vi.mocked(saveKey).mockResolvedValue({} as never); + const refresh = vi.fn().mockResolvedValue(undefined); + const account = { + id: "account", + modelType: "custom_api", + availableModels: ["model-a", "model-b"], + enabledModels: ["model-a"], + } as KeyVaultAccount; + + await toggleModelForAccounts( + "model-a", + "custom_api", + false, + [account], + refresh + ); + + expect(saveKey).toHaveBeenCalledWith({ + id: "account", + agent_type: "custom_api", + enabled_models: [], + }); + expect(refresh).toHaveBeenCalledOnce(); + }); +}); diff --git a/src/modules/MainApp/Integrations/hooks/modelToggle.ts b/src/modules/MainApp/Integrations/hooks/modelToggle.ts index 2890b32abb..bb9399bf55 100644 --- a/src/modules/MainApp/Integrations/hooks/modelToggle.ts +++ b/src/modules/MainApp/Integrations/hooks/modelToggle.ts @@ -27,7 +27,6 @@ export async function toggleModelForAccounts( return saveKey({ id: acc.id, agent_type: acc.modelType, - available_models: acc.availableModels ?? [], enabled_models: [...currentEnabled], }); }) diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/VariantPill.tsx b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/VariantPill.tsx index aa63870c46..a5a67d6391 100644 --- a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/VariantPill.tsx +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/VariantPill.tsx @@ -19,8 +19,10 @@ import ModelPropertiesDropdown from "@src/components/ModelPropertiesDropdown"; import { BrainIcon, HugeiconsIcon, Pen01Icon } from "@src/icons"; import { separateEffortPillAtom } from "@src/store/session/separateEffortPillAtom"; import { + type ResolvedModelVariantFields, formatReasoningLevel, - parseModelVariant, + resolveModelVariantFields, + toModelReasoningLevel, } from "@src/util/modelVariants"; import { buildVariantEditOptions } from "@src/util/variantEditOptions"; @@ -35,6 +37,7 @@ interface VariantPillProps { * When omitted, the pill is non-editable (legacy callers). */ groupModelIds?: readonly string[]; + variantMetadata?: readonly ResolvedModelVariantFields[]; /** * Called when the user changes the variant. Receives * the resolved model id. Caller persists it via the relevant @@ -46,16 +49,20 @@ interface VariantPillProps { export const VariantPill: React.FC = ({ modelId, groupModelIds, + variantMetadata, onApply, }) => { const variantOptions = React.useMemo( - () => buildVariantEditOptions(groupModelIds ?? [modelId]), - [groupModelIds, modelId] + () => buildVariantEditOptions(groupModelIds ?? [modelId], variantMetadata), + [groupModelIds, modelId, variantMetadata] ); const effectiveModelId = variantOptions.resolveVariantId(variantOptions.parseSelection(modelId)) ?? modelId; - const variant = parseModelVariant(effectiveModelId); + const variant = resolveModelVariantFields( + effectiveModelId, + variantMetadata?.find((entry) => entry.model === effectiveModelId) + ); const pillClasses = "relative z-10 inline-flex h-[24px] shrink-0 items-center gap-0.5 rounded-full border border-transparent bg-transparent px-2 text-[11px] font-semibold text-text-2 transition-colors group-hover/model-row:border-border-3 group-hover/model-row:bg-bg-1 group-focus-within/model-row:border-border-3 group-focus-within/model-row:bg-bg-1"; @@ -82,7 +89,7 @@ export const VariantPill: React.FC = ({ const editable = onApply !== undefined && (groupModelIds?.length ?? 0) > 1; const parts: string[] = []; if (variant?.reasoning) { - parts.push(formatReasoningLevel(variant.reasoning)); + parts.push(formatReasoningLevel(toModelReasoningLevel(variant.reasoning))); } if (variant?.fast) { parts.push("Fast"); diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/__tests__/sourceItems.test.ts b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/__tests__/sourceItems.test.ts new file mode 100644 index 0000000000..ab47843c5c --- /dev/null +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/__tests__/sourceItems.test.ts @@ -0,0 +1,104 @@ +import { describe, expect, it, vi } from "vitest"; + +import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; +import { buildAccountLookup } from "@src/hooks/models/useModelAccountLookup"; + +import { buildKeyModelItems, selectableKeyAccounts } from "../keyFirstItems"; +import { buildAllModelItems } from "../modelSelectionItems"; +import { buildSourceItems, buildSourceOptions } from "../sourceItems"; + +function account(overrides: Partial = {}): KeyVaultAccount { + return { + id: "key", + name: "Account", + modelType: "codex", + enabled: true, + hasLocalKey: true, + isListed: false, + hasKey: true, + hasApiKey: false, + hasSessionToken: true, + status: "ready", + availableModels: [], + enabledModels: ["catalog-family"], + modelVariants: [ + { + model: "deployment-a", + base_model: "catalog-family", + reasoning: "low", + fast: false, + }, + { + model: "deployment-b", + base_model: "catalog-family", + reasoning: "high", + fast: false, + }, + ], + defaultVariants: [{ base_model: "catalog-family", model: "deployment-b" }], + ...overrides, + }; +} + +describe("model and source catalog parity", () => { + it("offers a variants-only model with the same source and default in both column orders", () => { + const source = account(); + const lookup = buildAccountLookup([source]); + const options = buildSourceOptions([...lookup.keys()], [source], false); + expect(options.map((option) => option.accountId)).toEqual([source.id]); + const allRows = buildAllModelItems({ + accountLookup: lookup, + accounts: [source], + handleModelSelect: vi.fn(), + modelAliasVersion: 0, + resolveGroupLaunchModel: (models) => models[0], + }); + expect(allRows).toHaveLength(1); + expect(allRows[0].data?.groupModelIds).toEqual( + expect.arrayContaining(["deployment-a", "deployment-b"]) + ); + const onCommit = vi.fn(); + const [keyRow] = buildKeyModelItems({ + account: source, + onCommit, + persistDefaultVariantForAccount: vi.fn(), + }); + keyRow.action?.(); + expect(onCommit).toHaveBeenCalledWith(source, "deployment-b"); + const [sourceRow] = buildSourceItems({ + sourceOptions: options, + selectedModelId: "deployment-a", + selectedGroupModelIds: [...lookup.keys()], + accounts: [source], + handleSourceSelect: vi.fn(), + persistDefaultVariantForAccount: vi.fn(), + }); + expect(sourceRow.data?.rightContent).toMatchObject({ + props: { modelId: "deployment-b", variantMetadata: source.modelVariants }, + }); + }); + + it.each([ + { enabled: false }, + { hasKey: false }, + { status: "error" as const }, + { enabledModels: [] }, + ])( + "does not expose an unusable account in either column: %j", + (overrides) => { + const source = account(overrides); + expect(buildAccountLookup([source]).size).toBe(0); + expect(buildSourceOptions(["deployment-a"], [source], true)).toEqual([]); + expect(selectableKeyAccounts([source], true)).toEqual([]); + } + ); + + it("preserves model-less CLI launch, without treating all-disabled models as no catalog", () => { + const noCatalog = account({ modelVariants: [], enabledModels: [] }); + expect(buildSourceOptions([], [noCatalog], true)).toHaveLength(1); + expect(selectableKeyAccounts([noCatalog], true)).toEqual([noCatalog]); + expect( + buildSourceOptions([], [account({ enabledModels: [] })], true) + ).toEqual([]); + }); +}); diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/__tests__/useUnifiedModelPaletteSelection.test.ts b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/__tests__/useUnifiedModelPaletteSelection.test.ts index 21aaaff1be..ba19a338d3 100644 --- a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/__tests__/useUnifiedModelPaletteSelection.test.ts +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/__tests__/useUnifiedModelPaletteSelection.test.ts @@ -263,3 +263,93 @@ it("keeps ordinary Account Key selection usable while signed out of Cloud", asyn mounted.dispose(); } }); + +it("dismisses an ordinary pick immediately but records it only after persistence", async () => { + signedInStore(); + const mounted = await mountPicker(); + let accept!: (accepted: boolean) => void; + mounted.onConfigChange.mockImplementationOnce( + () => + new Promise((resolve) => { + accept = resolve; + }) + ); + try { + mounted.picker.handleSourceSelect( + { + id: "key", + label: "My key", + type: "own_key", + modelType: "openai_api", + accountId: "key", + }, + "gpt" + ); + expect(mounted.onClose).toHaveBeenCalledOnce(); + expect(mounted.recordRecent).not.toHaveBeenCalled(); + await act(async () => accept(true)); + expect(mounted.recordRecent).toHaveBeenCalledOnce(); + } finally { + mounted.dispose(); + } +}); + +it.each([false, true])( + "never records a refused or rejected pick (reject=%s)", + async (reject) => { + signedInStore(); + const mounted = await mountPicker(); + mounted.onConfigChange.mockImplementationOnce(() => + reject ? Promise.reject(new Error("refused")) : false + ); + try { + await act(async () => + mounted.picker.handleSourceSelect( + { + id: "key", + label: "My key", + type: "own_key", + modelType: "openai_api", + accountId: "key", + }, + "gpt" + ) + ); + expect(mounted.recordRecent).not.toHaveBeenCalled(); + } finally { + mounted.dispose(); + } + } +); + +it("does not apply a prepared Market pick after a newer own-key selection", async () => { + signedInStore(); + let finish!: (value: { credentialSource: string }) => void; + vi.mocked(prepareMarketProfileSource).mockImplementationOnce( + () => + new Promise((resolve) => { + finish = resolve; + }) + ); + const mounted = await mountPicker(); + try { + mounted.picker.handleMarketModelSelect(source, "gpt"); + mounted.picker.handleSourceSelect( + { + id: "key", + label: "My key", + type: "own_key", + modelType: "openai_api", + accountId: "key", + }, + "new-model" + ); + await act(async () => finish({ credentialSource: "market:late" })); + expect(mounted.onConfigChange).toHaveBeenCalledOnce(); + expect(mounted.recordRecent).toHaveBeenCalledWith( + expect.objectContaining({ modelId: "new-model", accountId: "key" }) + ); + } finally { + mounted.dispose(); + } +}); diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/keyFirstItems.tsx b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/keyFirstItems.tsx index 6a1172e5c2..3b9e6c9e4c 100644 --- a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/keyFirstItems.tsx +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/keyFirstItems.tsx @@ -11,11 +11,14 @@ import React from "react"; import ModelIcon from "@src/components/ModelIcon"; import type { MarketProfileSource } from "@src/features/MarketConnect/marketProfiles"; import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; -import { getModelAliasDisplayName } from "@src/hooks/models/modelAliasRegistry"; import { - accountHasModel, accountModelIds, -} from "@src/hooks/models/useModelAccountLookup"; + groupCatalogModels, + isSelectableModelAccount, + resolveAccountModelVariant, + selectableAccountModelIds, +} from "@src/hooks/models/accountModelCatalog"; +import { getModelAliasDisplayName } from "@src/hooks/models/modelAliasRegistry"; import { resolveDefaultVariant } from "@src/util/defaultModelVariant"; import { compareModelsByVersion, @@ -38,9 +41,7 @@ export const MARKET_PROFILE_TEST_ID = "unified-model-market-profile-option"; /** Model ids the account can actually launch (enabled, incl. variant rungs). */ export function enabledAccountModelIds(account: KeyVaultAccount): string[] { - return accountModelIds(account).filter((modelId) => - accountHasModel(account, modelId) - ); + return selectableAccountModelIds(account); } /** Keys that appear in the left column of key-first mode. */ @@ -49,10 +50,13 @@ export function selectableKeyAccounts( isCliAgent: boolean ): KeyVaultAccount[] { return accounts.filter((account) => { - if (account.status !== "ready" || !account.hasKey) return false; + if (!isSelectableModelAccount(account)) return false; // A CLI-agent key with no model listing is still launchable — the // session falls back to the agent's own default model — so keep it. - return isCliAgent || enabledAccountModelIds(account).length > 0; + return ( + (isCliAgent && accountModelIds(account).length === 0) || + enabledAccountModelIds(account).length > 0 + ); }); } @@ -73,7 +77,11 @@ export function buildKeyItems({ }: BuildKeyItemsParams): SpotlightItem[] { return selectableKeyAccounts(accounts, isCliAgent).map((account) => { const modelIds = enabledAccountModelIds(account); - const groupCount = groupModels(modelIds).length; + const groupCount = groupCatalogModels( + modelIds, + [account], + account.modelType + ).length; const KeyIcon = () => ; const labelContent = ( @@ -234,7 +242,11 @@ export function buildKeyModelItems({ ? enabledAccountModelIds(account).flatMap((model) => groupModels([model], account.modelType) ) - : groupModels(enabledAccountModelIds(account), account.modelType); + : groupCatalogModels( + enabledAccountModelIds(account), + [account], + account.modelType + ); for (const group of groups) { const sortedVariants = [...group.models].sort(compareModelsByVersion); @@ -242,13 +254,11 @@ export function buildKeyModelItems({ if (!representative) continue; const variantInfos = sortedVariants.map((modelId) => - resolveModelVariantFields(modelId) + resolveAccountModelVariant(account, modelId) ); const baseModel = literalModels ? representative - : (parseModelVariant(representative)?.baseModel ?? - variantInfos[0]?.base_model ?? - representative); + : (variantInfos[0]?.base_model ?? representative); const persisted = (account.defaultVariants ?? []).find( (entry) => entry.base_model === baseModel && sortedVariants.includes(entry.model) @@ -297,6 +307,7 @@ export function buildKeyModelItems({ persistDefaultVariantForAccount(account.id, baseModel, nextModelId) } @@ -309,7 +320,7 @@ export function buildKeyModelItems({ withModelRowAttributes({ id: literalModels ? `key-model:${account.id}:${representative}` - : `key-model:${account.id}:${group.label}:${group.sortVersion}`, + : `key-model:${account.id}:${group.label}:${group.sortVersion}:${baseModel}`, label: [displayLabel, ...sortedVariants].join(" "), icon: ModelItemIcon, type: "action" as const, diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSection.ts b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSection.ts index 1bf5b62e10..6bacea7761 100644 --- a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSection.ts +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSection.ts @@ -1,10 +1,13 @@ import type { AdvancedConfig } from "@src/features/SessionCreator/types"; +import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; +import { + groupCatalogModels, + resolveAccountModelVariant, +} from "@src/hooks/models/accountModelCatalog"; import { type RecentModelEntry, marketSelectionsEquivalent, } from "@src/store/session/recentModelEntriesAtom"; -import { groupModels } from "@src/util/modelGrouping"; -import { getModelVariantBaseModel } from "@src/util/modelVariants"; import type { SpotlightItem } from "../../types"; @@ -49,13 +52,20 @@ export function entryMatchesActiveConfig( | "selectedSourceModelType" | "listingModelType" | "cliAgentType" - > + >, + accounts: readonly KeyVaultAccount[] = [] ): boolean { const activeModel = getActiveModelId(config); if ( !activeModel || - getModelVariantBaseModel(entry.modelId) !== - getModelVariantBaseModel(activeModel) + resolveAccountModelVariant( + accounts.find((account) => account.id === entry.accountId), + entry.modelId + ).base_model !== + resolveAccountModelVariant( + accounts.find((account) => account.id === config.selectedAccountId), + activeModel + ).base_model ) { return false; } @@ -94,9 +104,10 @@ export function entryMatchesActiveConfig( } export function buildGroupByModel( - modelIds: Iterable + modelIds: Iterable, + accounts: readonly KeyVaultAccount[] = [] ): Map { - const groups = groupModels(Array.from(modelIds)); + const groups = groupCatalogModels(Array.from(modelIds), accounts); const groupMap = new Map(); for (const group of groups) { for (const modelId of group.models) { diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSelectionCommit.ts b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSelectionCommit.ts new file mode 100644 index 0000000000..3dcdbbfe63 --- /dev/null +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSelectionCommit.ts @@ -0,0 +1,36 @@ +import type { AdvancedConfig } from "@src/features/SessionCreator/types"; +import type { RecentModelEntry } from "@src/store/session/recentModelEntriesAtom"; + +/** `false` means refused or superseded; void keeps creator callers compatible. */ +export type ModelConfigChange = ( + config: AdvancedConfig +) => void | boolean | Promise; + +/** Keep the existing immediate dismissal, but remember only successful picks. */ +export function commitModelSelection(options: { + config: AdvancedConfig; + entry: RecentModelEntry; + apply: ModelConfigChange; + record: (entry: RecentModelEntry) => void; + close?: () => void; + isCurrent: () => boolean; + onError: (error: unknown) => void; +}): Promise | void { + try { + const result = options.apply(options.config); + if (result === false) return; + if (result instanceof Promise) { + options.close?.(); + return result + .then((accepted) => { + if (accepted !== false && options.isCurrent()) + options.record(options.entry); + }) + .catch(options.onError); + } + if (options.isCurrent()) options.record(options.entry); + options.close?.(); + } catch (error) { + options.onError(error); + } +} diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSelectionItems.tsx b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSelectionItems.tsx index 2b629dcddd..b36d915479 100644 --- a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSelectionItems.tsx +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/modelSelectionItems.tsx @@ -3,20 +3,19 @@ import React from "react"; import ModelIcon from "@src/components/ModelIcon"; import type { MarketProfileSource } from "@src/features/MarketConnect/marketProfiles"; import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; +import { + accountHasModel, + groupCatalogModels, + resolveAccountModelVariant, +} from "@src/hooks/models/accountModelCatalog"; import { getModelAliasDisplayName } from "@src/hooks/models/modelAliasRegistry"; import type { ModelAccountInfo } from "@src/hooks/models/types"; -import { accountHasModel } from "@src/hooks/models/useModelAccountLookup"; import type { RecentModelEntry } from "@src/store/session/recentModelEntriesAtom"; import { resolveDefaultVariant } from "@src/util/defaultModelVariant"; import { compareModelsByVersion, formatModelNameFull, } from "@src/util/formatModelName"; -import { groupModels } from "@src/util/modelGrouping"; -import { - parseModelVariant, - resolveModelVariantFields, -} from "@src/util/modelVariants"; import type { SpotlightItem } from "../../types"; import { VariantPill } from "./VariantPill"; @@ -71,7 +70,11 @@ export function buildModelSelectionSpotlightItem({ const concreteModelDisplay = getModelAliasDisplayName(entry.modelId) ?? formatModelNameFull(entry.modelId, entryAgentType); - const groupedModel = groupModels([...family], entryAgentType)[0]; + const groupedModel = groupCatalogModels( + [...family], + recentAccount ? [recentAccount] : accounts, + entryAgentType + )[0]; const modelDisplay = groupedModel && groupedModel.label !== "Other" ? groupedModel.label @@ -110,9 +113,11 @@ export function buildModelSelectionSpotlightItem({ ? family.filter((modelId) => accountHasModel(recentAccount, modelId)) : [entry.modelId]; - const variant = parseModelVariant(entry.modelId); - const previewBaseModel = - variant?.baseModel ?? resolveModelVariantFields(entry.modelId).base_model; + const variant = resolveAccountModelVariant(recentAccount, entry.modelId); + const previewBaseModel = variant.base_model; + const variantInfos = accountFamilyIds.map((modelId) => + resolveAccountModelVariant(recentAccount, modelId) + ); const persistedVariant = recentAccount && previewBaseModel ? (recentAccount.defaultVariants ?? []).find( @@ -125,7 +130,7 @@ export function buildModelSelectionSpotlightItem({ previewBaseModel && accountFamilyIds.length > 0 ? (resolveDefaultVariant( previewBaseModel, - accountFamilyIds.map((modelId) => resolveModelVariantFields(modelId)), + variantInfos, persistedVariant ) ?? entry.modelId) : entry.modelId; @@ -146,16 +151,16 @@ export function buildModelSelectionSpotlightItem({ : undefined; const accountHasMultipleVariants = accountFamilyIds.length > 1; - const trailing: React.ReactNode = - variant && accountHasMultipleVariants ? ( - - ) : ( - - ); + const trailing: React.ReactNode = accountHasMultipleVariants ? ( + + ) : ( + + ); return withModelRowAttributes({ id: `${idPrefix}:${entry.modelId}:${entry.accountId ?? entry.sourceType}`, @@ -201,7 +206,7 @@ export function buildAllModelItems({ const items: SpotlightItem[] = []; const modelIds = Array.from(accountLookup.keys()); - const groups = groupModels(modelIds); + const groups = groupCatalogModels(modelIds, accounts); const getAccountCount = (modelIdsForRow: string[]) => accounts.filter( @@ -304,7 +309,7 @@ export function buildAllModelItems({ items.push( withModelRowAttributes({ - id: `group:${group.label}:${group.sortVersion}`, + id: `group:${group.label}:${group.sortVersion}:${representativeModel}`, label: searchableLabel, icon: GroupItemIcon, type: "action" as const, diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/sourceItems.tsx b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/sourceItems.tsx index 1e67ffc1da..4ac27af19a 100644 --- a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/sourceItems.tsx +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/sourceItems.tsx @@ -4,12 +4,13 @@ import { KEY_SOURCE } from "@src/api/tauri/session"; import ModelIcon from "@src/components/ModelIcon"; import type { MarketProfileSource } from "@src/features/MarketConnect/marketProfiles"; import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; -import { accountHasModel } from "@src/hooks/models/useModelAccountLookup"; -import { resolveDefaultVariant } from "@src/util/defaultModelVariant"; import { - parseModelVariant, - resolveModelVariantFields, -} from "@src/util/modelVariants"; + accountHasModel, + accountModelIds, + isSelectableModelAccount, + resolveAccountModelVariant, +} from "@src/hooks/models/accountModelCatalog"; +import { resolveDefaultVariant } from "@src/util/defaultModelVariant"; import type { SpotlightItem } from "../../types"; import { VariantPill } from "./VariantPill"; @@ -37,16 +38,15 @@ export function buildSourceOptions( const options: SourceOption[] = []; const variantSet = new Set(modelIds.filter(Boolean)); - const readyAccounts = accounts.filter( - (account) => account.status === "ready" && account.hasKey - ); + const readyAccounts = accounts.filter(isSelectableModelAccount); for (const account of readyAccounts) { const hasConcreteModelFilter = variantSet.size > 0; - const hasAnyVariant = - account.availableModels && account.availableModels.length > 0 - ? [...variantSet].some((modelId) => accountHasModel(account, modelId)) - : false; - const modelMatches = hasConcreteModelFilter ? hasAnyVariant : isCliAgent; + const hasAnyVariant = [...variantSet].some((modelId) => + accountHasModel(account, modelId) + ); + const modelMatches = hasConcreteModelFilter + ? hasAnyVariant + : isCliAgent && accountModelIds(account).length === 0; if (modelMatches) { options.push(toSourceOption(account)); } @@ -89,12 +89,6 @@ export function buildSourceItems({ handleSourceSelect, persistDefaultVariantForAccount, }: BuildSourceItemsParams): SpotlightItem[] { - const previewVariantInfo = selectedModelId - ? parseModelVariant(selectedModelId) - : null; - const previewBaseModel = - previewVariantInfo?.baseModel ?? selectedModelId ?? undefined; - const accountById = new Map(accounts.map((account) => [account.id, account])); return sourceOptions.map((source) => { @@ -105,6 +99,9 @@ export function buildSourceItems({ const sourceAccount = source.accountId ? accountById.get(source.accountId) : undefined; + const previewBaseModel = selectedModelId + ? resolveAccountModelVariant(sourceAccount, selectedModelId).base_model + : undefined; const accountVariantIds = source.marketSource ? selectedGroupModelIds.filter((modelId) => @@ -116,6 +113,9 @@ export function buildSourceItems({ ) : []; + const variantInfos = accountVariantIds.map((modelId) => + resolveAccountModelVariant(sourceAccount, modelId) + ); let accountEffectiveModelId: string | undefined; if ( (sourceAccount || source.marketSource) && @@ -127,9 +127,6 @@ export function buildSourceItems({ entry.base_model === previewBaseModel && accountVariantIds.includes(entry.model) )?.model; - const variantInfos = accountVariantIds.map((modelId) => - resolveModelVariantFields(modelId) - ); accountEffectiveModelId = resolveDefaultVariant(previewBaseModel, variantInfos, persisted) ?? accountVariantIds[0]; @@ -158,6 +155,7 @@ export function buildSourceItems({ ); diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/types.ts b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/types.ts index e47879a2ed..34e1f81700 100644 --- a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/types.ts +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/types.ts @@ -8,6 +8,7 @@ import type { MarketProfileSource } from "@src/features/MarketConnect/marketProf import type { AdvancedConfig } from "@src/features/SessionCreator/types"; import type { BasePaletteProps } from "../../shared"; +import type { ModelConfigChange } from "./modelSelectionCommit"; export interface SourceOption { id: string; @@ -24,7 +25,7 @@ export interface SourceOption { export interface UnifiedModelPaletteProps extends BasePaletteProps { advancedConfig: AdvancedConfig; - onConfigChange: (config: AdvancedConfig) => void; + onConfigChange: ModelConfigChange; /** * Display-name override for an already-running conversation's runtime. * Creator surfaces omit this and keep using the SessionCreator selection. diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteData.ts b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteData.ts index d3861ed02d..2ead92f1f6 100644 --- a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteData.ts +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteData.ts @@ -21,11 +21,8 @@ import { type UseKeyVaultReturn, useKeyVault, } from "@src/hooks/keyVault"; -import { withNativeHarnessModels } from "@src/hooks/models/nativeHarnessAccountModels"; -import { - getCliCompatibleAccounts, - useAgentCompatibility, -} from "@src/hooks/models/useAgentCompatibility"; +import { getModelPickerAccounts } from "@src/hooks/models/accountModelCatalog"; +import { useAgentCompatibility } from "@src/hooks/models/useAgentCompatibility"; import { buildAccountLookup } from "@src/hooks/models/useModelAccountLookup"; import { useOrgiiPoolCategories } from "@src/hooks/models/useOrgiiPoolCategories"; import { @@ -122,11 +119,8 @@ export interface UnifiedModelPaletteData { recordRecent: (entry: RecentModelEntry) => void; /** * Persist key edits (e.g. per-account default variants from the - * variant pill). Exposed from the same `useKeyVault` instance that - * supplies `accounts` so optimistic state updates after `saveKey` - * actually flow back into this palette's account list — using a - * second `useKeyVault()` would give us a parallel local-state copy - * that never refreshes until the palette is reopened. + * variant pill). All `useKeyVault` consumers subscribe to the shared + * local account store, so a successful save publishes to every picker. */ saveKey: UseKeyVaultReturn["saveKey"]; /** True while the persisted Key Vault account list is loading. */ @@ -170,13 +164,16 @@ export function useUnifiedModelPaletteData({ const [refreshingAllModels, setRefreshingAllModels] = useState(false); - const accounts = useMemo(() => { - if (dispatchCategory === "cli_agent" && cliAgentType) { - return getCliCompatibleAccounts(registry, cliAgentType, allAccounts); - } - - return withNativeHarnessModels(allAccounts, dispatchCategory); - }, [dispatchCategory, cliAgentType, allAccounts, registry]); + const accounts = useMemo( + () => + getModelPickerAccounts( + registry, + allAccounts, + dispatchCategory, + cliAgentType + ), + [dispatchCategory, cliAgentType, allAccounts, registry] + ); const { sources: marketSources, diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteItems.ts b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteItems.ts index ca24de9a44..91f300539c 100644 --- a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteItems.ts +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteItems.ts @@ -7,6 +7,7 @@ import type { MarketProfileSource } from "@src/features/MarketConnect/marketProf import { findMarketSourceForRecent } from "@src/features/MarketConnect/marketProfiles"; import type { AdvancedConfig } from "@src/features/SessionCreator/types"; import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; +import { resolveCatalogModelVariant } from "@src/hooks/models/accountModelCatalog"; import { isPairCompatible } from "@src/hooks/models/modelPairCompatibility"; import { accountHasModel } from "@src/hooks/models/useModelAccountLookup"; import type { RecentModelEntry } from "@src/store/session/recentModelEntriesAtom"; @@ -17,7 +18,6 @@ import { spotlightModelPinsAtom, } from "@src/store/ui/spotlightPinsAtom"; import { resolveDefaultVariant } from "@src/util/defaultModelVariant"; -import { resolveModelVariantFields } from "@src/util/modelVariants"; import { isModelPinned, toggleModelPin } from "../../pinning/modelPins"; import type { SpotlightItem } from "../../types"; @@ -181,16 +181,11 @@ export function useUnifiedModelPaletteItems({ const account = accounts.find((entry) => entry.id === accountId); if (!account) return; - const nextDefaults = (account.defaultVariants ?? []).filter( - (variant) => variant.base_model !== baseModel - ); - nextDefaults.push({ base_model: baseModel, model: modelId }); - - void saveKey({ + saveKey({ id: account.id, agent_type: account.modelType, - default_variants: nextDefaults, - }); + default_variant_overrides: [{ base_model: baseModel, model: modelId }], + }).catch(() => undefined); }, [accounts, saveKey] ); @@ -198,8 +193,8 @@ export function useUnifiedModelPaletteItems({ // Quick-pick rows group variants over every reachable model, not just the // ones the active source scope lists. const groupByModel = useMemo( - () => buildGroupByModel(fullModelLookup.keys()), - [fullModelLookup] + () => buildGroupByModel(fullModelLookup.keys(), accounts), + [fullModelLookup, accounts] ); const activeModelId = getActiveModelId(advancedConfig); @@ -210,7 +205,7 @@ export function useUnifiedModelPaletteItems({ if (!activeModelId) return null; const fromRecents = compatibleRecentEntries.find((entry) => - entryMatchesActiveConfig(entry, advancedConfig) + entryMatchesActiveConfig(entry, advancedConfig, accounts) ); if (fromRecents) return fromRecents; @@ -293,7 +288,8 @@ export function useUnifiedModelPaletteItems({ (entry: RecentModelEntry, section: ModelSection, index: number) => { const isCurrentSelection = entryMatchesActiveConfig( entry, - advancedConfig + advancedConfig, + accounts ); // The active config may hold another variant of a stored entry. const rowEntry = @@ -353,7 +349,7 @@ export function useUnifiedModelPaletteItems({ if (sortedVariants.length === 0) return ""; const variantInfos = sortedVariants.map((modelId) => - resolveModelVariantFields(modelId) + resolveCatalogModelVariant(accounts, modelId) ); const baseModel = variantInfos[0]?.base_model ?? sortedVariants[0]; const variantModelSet = new Set(sortedVariants); diff --git a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteSelection.ts b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteSelection.ts index 575278bd0d..6d69bffa77 100644 --- a/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteSelection.ts +++ b/src/scaffold/GlobalSpotlight/palettes/UnifiedModelPalette/useUnifiedModelPaletteSelection.ts @@ -14,6 +14,7 @@ import { import type { AdvancedConfig } from "@src/features/SessionCreator/types"; import type { KeyVaultAccount } from "@src/hooks/keyVault/types"; import { createLogger } from "@src/hooks/logger"; +import { resolveAccountModelVariant } from "@src/hooks/models/accountModelCatalog"; import { accountHasModel, accountModelIds as listAccountModelIds, @@ -23,11 +24,11 @@ import { separateEffortPillAtom } from "@src/store/session/separateEffortPillAto import type { ModelSourceScope } from "@src/store/ui/spotlightModelSourceScopeAtom"; import { carryModelEffort } from "@src/util/carryModelEffort"; import { resolveDefaultVariant } from "@src/util/defaultModelVariant"; -import { - parseModelVariant, - resolveModelVariantFields, -} from "@src/util/modelVariants"; +import { + type ModelConfigChange, + commitModelSelection, +} from "./modelSelectionCommit"; import { buildSourceOptions, toSourceOption } from "./sourceItems"; import type { SourceOption } from "./types"; import { resolveVariantReselection } from "./variantReselect"; @@ -57,7 +58,7 @@ interface UseUnifiedModelPaletteSelectionParams { /** Flipping the source scope rebuilds both columns, so the cursor resets. */ sourceScope?: ModelSourceScope; advancedConfig: AdvancedConfig; - onConfigChange: (config: AdvancedConfig) => void; + onConfigChange: ModelConfigChange; onClose: () => void; closeOnSourceSelect?: boolean; recordRecent: (entry: RecentModelEntry) => void; @@ -81,6 +82,33 @@ export function useUnifiedModelPaletteSelection({ }: UseUnifiedModelPaletteSelectionParams) { const { t } = useTranslation("integrations"); const marketSelectionPendingRef = useRef(false); + const selectionGeneration = useRef(0); + const mounted = useRef(true); + useEffect(() => { + mounted.current = true; + return () => { + mounted.current = false; + }; + }, []); + const commitSelection = useCallback( + ( + config: AdvancedConfig, + entry: RecentModelEntry, + close: boolean, + generation: number + ) => + commitModelSelection({ + config, + entry, + apply: onConfigChange, + record: recordRecent, + close: close ? onClose : undefined, + isCurrent: () => generation === selectionGeneration.current, + onError: (error) => + Message.error(error instanceof Error ? error.message : String(error)), + }), + [onConfigChange, recordRecent, onClose] + ); const separateEffortPill = useAtomValue(separateEffortPillAtom); const currentModelId = advancedConfig.model; // With the separate effort pill, a model pick keeps the current effort @@ -150,6 +178,7 @@ export function useUnifiedModelPaletteSelection({ const applySourceSelection = useCallback( (modelId: string, _modelLabel: string, source: SourceOption) => { + const generation = ++selectionGeneration.current; const sourceAccount = source.accountId ? accounts.find((account) => account.id === source.accountId) : undefined; @@ -163,35 +192,40 @@ export function useUnifiedModelPaletteSelection({ : [] ) : advancedConfig.model || ""; - onConfigChange({ - ...advancedConfig, - keySource: KEY_SOURCE.OWN, - selectedAccountId: source.accountId, - credentialSource: undefined, - marketProfileId: undefined, - agent: source.modelType, - provider: source.modelType, - model: resolvedModelId, - nativeHarnessType: source.nativeHarnessType, - selectedSourceLabel: source.label, - selectedSourceModelType: source.modelType, - }); - recordRecent({ - modelId: resolvedModelId, - sourceType: source.type, - accountId: source.accountId, - accountName: source.label, - modelType: source.modelType, - }); - if (closeOnSourceSelect) onClose(); + Promise.resolve( + commitSelection( + { + ...advancedConfig, + keySource: KEY_SOURCE.OWN, + selectedAccountId: source.accountId, + credentialSource: undefined, + marketProfileId: undefined, + agent: source.modelType, + provider: source.modelType, + model: resolvedModelId, + nativeHarnessType: source.nativeHarnessType, + selectedSourceLabel: source.label, + selectedSourceModelType: source.modelType, + }, + { + modelId: resolvedModelId, + sourceType: source.type, + accountId: source.accountId, + accountName: source.label, + modelType: source.modelType, + }, + closeOnSourceSelect, + generation + ) + ).catch((error) => + Message.error(error instanceof Error ? error.message : String(error)) + ); }, [ accounts, advancedConfig, closeOnSourceSelect, - onConfigChange, - onClose, - recordRecent, + commitSelection, withCurrentEffort, ] ); @@ -242,15 +276,17 @@ export function useUnifiedModelPaletteSelection({ : []; if (accountModelIds.length === 0) return selectedModelId; - const selectedVariant = parseModelVariant(selectedModelId); - const baseModel = selectedVariant?.baseModel ?? selectedModelId; + const baseModel = resolveAccountModelVariant( + sourceAccount, + selectedModelId + ).base_model; const persisted = (sourceAccount?.defaultVariants ?? []).find( (entry) => entry.base_model === baseModel && accountModelIds.includes(entry.model) )?.model; const variantInfos = accountModelIds.map((modelId) => - resolveModelVariantFields(modelId) + resolveAccountModelVariant(sourceAccount, modelId) ); return ( resolveDefaultVariant(baseModel, variantInfos, persisted) ?? @@ -267,6 +303,7 @@ export function useUnifiedModelPaletteSelection({ options?: { close?: boolean; keepVariant?: boolean } ) => { if (marketSelectionPendingRef.current) return; + const generation = ++selectionGeneration.current; const modelId = options?.keepVariant ? pickedModelId : withCurrentEffort(pickedModelId, marketSource.modelIds); @@ -283,31 +320,36 @@ export function useUnifiedModelPaletteSelection({ void prepareMarketProfileSource(marketSource, modelId) .then(({ credentialSource }) => { owner.assertCurrent(); + if (!mounted.current || generation !== selectionGeneration.current) + return; const modelType = marketSourceModelType(marketSource, modelId); - onConfigChange({ - ...advancedConfig, - keySource: KEY_SOURCE.OWN, - selectedAccountId: undefined, - credentialSource, - marketProfileId: marketSource.profile.id, - agent: modelType, - provider: modelType, - model: modelId, - nativeHarnessType: undefined, - cliAgentType: marketSource.cliAgentType, - selectedSourceLabel: marketSource.label, - selectedSourceModelType: modelType, - }); - recordRecent({ - modelId, - sourceType: KEY_SOURCE.OWN, - accountName: marketSource.label, - credentialSource, - marketProfileId: marketSource.profile.id, - modelType, - cliAgentType: marketSource.cliAgentType, - }); - if (options?.close ?? closeOnSourceSelect) onClose(); + return commitSelection( + { + ...advancedConfig, + keySource: KEY_SOURCE.OWN, + selectedAccountId: undefined, + credentialSource, + marketProfileId: marketSource.profile.id, + agent: modelType, + provider: modelType, + model: modelId, + nativeHarnessType: undefined, + cliAgentType: marketSource.cliAgentType, + selectedSourceLabel: marketSource.label, + selectedSourceModelType: modelType, + }, + { + modelId, + sourceType: KEY_SOURCE.OWN, + accountName: marketSource.label, + credentialSource, + marketProfileId: marketSource.profile.id, + modelType, + cliAgentType: marketSource.cliAgentType, + }, + options?.close ?? closeOnSourceSelect, + generation + ); }) .catch((error: unknown) => { const code = marketActivationErrorCode(error); @@ -325,15 +367,7 @@ export function useUnifiedModelPaletteSelection({ marketSelectionPendingRef.current = false; }); }, - [ - advancedConfig, - closeOnSourceSelect, - onClose, - onConfigChange, - recordRecent, - t, - withCurrentEffort, - ] + [advancedConfig, closeOnSourceSelect, commitSelection, t, withCurrentEffort] ); const handleSourceSelect = useCallback( @@ -435,6 +469,7 @@ export function useUnifiedModelPaletteSelection({ return; } + const generation = ++selectionGeneration.current; const reboundEntry: RecentModelEntry = { ...entry, modelId: options?.keepVariant @@ -450,31 +485,35 @@ export function useUnifiedModelPaletteSelection({ modelType: reboundAccount.modelType, }; - onConfigChange({ - ...advancedConfig, - keySource: KEY_SOURCE.OWN, - selectedAccountId: reboundAccount.id, - credentialSource: undefined, - marketProfileId: undefined, - agent: reboundAccount.modelType, - provider: reboundAccount.modelType, - model: reboundEntry.modelId, - nativeHarnessType: reboundAccount.nativeHarnessType, - selectedSourceLabel: reboundAccount.name, - selectedSourceModelType: reboundAccount.modelType, - }); - - recordRecent(reboundEntry); - if (options?.close !== false) onClose(); + Promise.resolve( + commitSelection( + { + ...advancedConfig, + keySource: KEY_SOURCE.OWN, + selectedAccountId: reboundAccount.id, + credentialSource: undefined, + marketProfileId: undefined, + agent: reboundAccount.modelType, + provider: reboundAccount.modelType, + model: reboundEntry.modelId, + nativeHarnessType: reboundAccount.nativeHarnessType, + selectedSourceLabel: reboundAccount.name, + selectedSourceModelType: reboundAccount.modelType, + }, + reboundEntry, + options?.close !== false, + generation + ) + ).catch((error) => + Message.error(error instanceof Error ? error.message : String(error)) + ); }, [ accounts, advancedConfig, applyMarketSourceSelection, marketSources, - onConfigChange, - onClose, - recordRecent, + commitSelection, t, withCurrentEffort, ] diff --git a/src/store/session/creatorDefaultModelAtom.ts b/src/store/session/creatorDefaultModelAtom.ts index f641139d61..7ea4d0a503 100644 --- a/src/store/session/creatorDefaultModelAtom.ts +++ b/src/store/session/creatorDefaultModelAtom.ts @@ -43,7 +43,11 @@ import type { CliAgentType, ModelType, } from "@src/api/tauri/rpc/schemas/validation"; -import { KEY_SOURCE, isHostedKey } from "@src/api/tauri/session"; +import { + type DispatchCategory, + KEY_SOURCE, + isHostedKey, +} from "@src/api/tauri/session"; import { formatAgentType } from "@src/assets/providers"; import type { AdvancedConfig, @@ -282,9 +286,15 @@ export const creatorDefaultModelSelectionAtom = atom( const pair = get(creatorDefaultModelPairAtom); return pair ? deriveLastModelSelection(pair) : null; }, - (get, set, entry: RecentModelEntry | null) => { + ( + get, + set, + entry: RecentModelEntry | null, + categoryOverride?: DispatchCategory + ) => { const map = get(creatorDefaultModelMapAtom); - const category = get(dispatchCategoryAtom); + // Session selections may finish after the creator changed categories. + const category = categoryOverride ?? get(dispatchCategoryAtom); const newMap: LastModelPairMap = { ...map, diff --git a/src/util/__tests__/modelVariants.test.ts b/src/util/__tests__/modelVariants.test.ts index ba8d1e45d3..17507ef191 100644 --- a/src/util/__tests__/modelVariants.test.ts +++ b/src/util/__tests__/modelVariants.test.ts @@ -384,7 +384,7 @@ describe("parseModelVariant", () => { }); }); - it("prefers frontend parse over stale backend model variant metadata", () => { + it("uses backend catalog fields before model ID grammar", () => { expect( resolveModelVariantFields("gpt-5.1-codex-max-medium", { model: "gpt-5.1-codex-max-medium", @@ -394,9 +394,24 @@ describe("parseModelVariant", () => { }) ).toEqual({ model: "gpt-5.1-codex-max-medium", - base_model: "gpt-5.1-codex-max", - reasoning: MODEL_REASONING_LEVEL.MEDIUM, + base_model: "gpt-5.1-codex-max-medium", + reasoning: MODEL_REASONING_LEVEL.MAX, fast: false, }); }); + + it.each(["o4-mini", "o4-nano"])( + "keeps the %s size suffix in its inferred effort family", + (model) => { + expect(resolveModelVariantFields(model).base_model).toBe(model); + expect(resolveModelVariantFields(`${model}-high`).base_model).toBe(model); + expect( + resolveModelVariantFields(model, { + model, + base_model: "o4", + fast: false, + }).base_model + ).toBe("o4"); + } + ); }); diff --git a/src/util/__tests__/variantEditOptions.test.ts b/src/util/__tests__/variantEditOptions.test.ts index f1d4aae11b..a29b6f0705 100644 --- a/src/util/__tests__/variantEditOptions.test.ts +++ b/src/util/__tests__/variantEditOptions.test.ts @@ -81,3 +81,32 @@ describe("buildVariantEditOptions", () => { } }); }); + +it("uses wire effort and fast fields while preserving the independent thinking dimension", () => { + const modelIds = ["claude-opus-4-8-thinking-high", "deployment-b"]; + const metadata = [ + { + model: modelIds[0], + base_model: "catalog-family", + reasoning: "low", + fast: true, + }, + { + model: modelIds[1], + base_model: "catalog-family", + reasoning: "high", + fast: false, + }, + ]; + const options = buildVariantEditOptions(modelIds, metadata); + expect(options.availableLevels).toEqual(["low", "high"]); + expect(options.thinkingToggleable).toBe(true); + expect(options.parseSelection(modelIds[0])).toEqual({ + thinking: true, + level: "low", + fast: true, + }); + expect( + options.resolveVariantId({ thinking: false, level: "high", fast: false }) + ).toBe("deployment-b"); +}); diff --git a/src/util/defaultModelVariant.ts b/src/util/defaultModelVariant.ts index b61e0f585b..9146b7817c 100644 --- a/src/util/defaultModelVariant.ts +++ b/src/util/defaultModelVariant.ts @@ -2,7 +2,7 @@ import type { ModelTableVariantInfo } from "@src/types/modelTable"; import { MODEL_REASONING_LEVEL, type ModelReasoningLevel, - parseModelVariant, + resolveModelVariantFields, toModelReasoningLevel, } from "@src/util/modelVariants"; @@ -128,10 +128,15 @@ export function resolveDefaultVariant( ) { return persistedModel; } - if (persistedModel && !parseModelVariant(persistedModel)?.reasoning) { - const previous = parseModelVariant(persistedModel); + const previous = persistedModel + ? resolveModelVariantFields( + persistedModel, + variants.find((variant) => variant.model === persistedModel) + ) + : undefined; + if (persistedModel && !previous?.reasoning) { const sameOptions = selectable.filter((variant) => { - const parsed = parseModelVariant(variant.model); + const parsed = resolveModelVariantFields(variant.model, variant); return ( (parsed?.thinking ?? false) === (previous?.thinking ?? false) && (parsed?.fast ?? false) === (previous?.fast ?? false) diff --git a/src/util/modelVariants.ts b/src/util/modelVariants.ts index eae44a45fa..b794443068 100644 --- a/src/util/modelVariants.ts +++ b/src/util/modelVariants.ts @@ -353,26 +353,37 @@ export interface ResolvedModelVariantFields { reasoning?: string | null; fast: boolean; context_window?: number | null; + /** Wire catalogs currently omit thinking; retain the ID grammar for this dimension. */ + thinking?: boolean; } -/** Frontend parse wins over backend model_variants wire metadata. */ +/** Catalog fields are authoritative; parse IDs only when metadata is absent. */ export function resolveModelVariantFields( model: string, - fallback?: ResolvedModelVariantFields + metadata?: ResolvedModelVariantFields ): ResolvedModelVariantFields { const parsed = parseModelVariant(model); + if (metadata) { + return { + ...metadata, + model, + ...(metadata.thinking === undefined && parsed?.thinking + ? { thinking: true } + : {}), + }; + } if (parsed) { return { model: parsed.model, - base_model: parsed.baseModel, + base_model: + parsed.reasoning || parsed.thinking || parsed.fast + ? parsed.baseModel + : model, reasoning: parsed.reasoning ?? null, fast: parsed.fast, - context_window: fallback?.context_window, + ...(parsed.thinking ? { thinking: true } : {}), }; } - if (fallback) { - return fallback; - } return { model, base_model: model, fast: false }; } diff --git a/src/util/selectableModelVariants.ts b/src/util/selectableModelVariants.ts index d29c3803b6..ea338656fb 100644 --- a/src/util/selectableModelVariants.ts +++ b/src/util/selectableModelVariants.ts @@ -1,4 +1,8 @@ -import { parseModelVariant } from "./modelVariants"; +import { + type ResolvedModelVariantFields, + parseModelVariant, + resolveModelVariantFields, +} from "./modelVariants"; /** Size variants such as o4-mini own an effort ladder; mini is not effort. */ export function getModelEffortBaseModel(model: string): string { @@ -13,18 +17,28 @@ export function getModelEffortBaseModel(model: string): string { * same thinking/speed combination. Keep bare ids for models without that * ladder, including speed-only and thinking-only families. */ -export function selectableModelVariants( - variants: readonly T[] -): T[] { +export function selectableModelVariants< + T extends { + model: string; + base_model?: string; + reasoning?: string | null; + fast?: boolean; + thinking?: boolean; + }, +>(variants: readonly T[]): T[] { const entries = variants.map((variant) => { - const parsed = parseModelVariant(variant.model); + const metadata = + variant.base_model !== undefined && variant.fast !== undefined + ? (variant as ResolvedModelVariantFields) + : undefined; + const resolved = resolveModelVariantFields(variant.model, metadata); return { variant, - reasoning: parsed?.reasoning, + reasoning: resolved.reasoning, key: JSON.stringify([ - getModelEffortBaseModel(variant.model), - parsed?.thinking ?? false, - parsed?.fast ?? false, + metadata?.base_model ?? getModelEffortBaseModel(variant.model), + resolved.thinking ?? false, + resolved.fast, ]), }; }); diff --git a/src/util/session/__tests__/selectionFromSession.test.ts b/src/util/session/__tests__/selectionFromSession.test.ts new file mode 100644 index 0000000000..5ea8c15e07 --- /dev/null +++ b/src/util/session/__tests__/selectionFromSession.test.ts @@ -0,0 +1,58 @@ +import { describe, expect, it } from "vitest"; + +import type { LastModelSelection } from "@src/store/session/creatorDefaultModelAtom"; +import type { Session } from "@src/store/session/sessionAtom/types"; + +import { resolveModelForMessage } from "../resolveModelForMessage"; +import { selectionFromSession } from "../selectionFromSession"; + +const fallback: LastModelSelection = { + keySource: "own_key", + model: "creator-model", + selectedAccountId: "creator-account", + cliAgentType: "codex", +}; +const session: Session = { + session_id: "session", + created_at: "2026-09-23", + updated_at: "2026-09-23", + status: "idle", + model: "session-model", +}; + +describe("session model identity", () => { + it("uses a complete default only before a session has a model", () => { + expect(selectionFromSession(undefined, fallback)).toBe(fallback); + expect( + selectionFromSession({ ...session, model: undefined }, fallback) + ).toBe(fallback); + }); + it("never sends a creator account alongside an existing session model", () => { + const selected = selectionFromSession(session, fallback); + expect(resolveModelForMessage(selected)).toEqual({ + model: "session-model", + accountId: undefined, + }); + expect(selected?.cliAgentType).toBeUndefined(); + expect(selected?.keySource).toBeUndefined(); + }); + it("preserves a Market credential without borrowing a personal account", () => { + const selected = selectionFromSession( + { + ...session, + keySource: "own_key", + credentialSource: "market:selection", + }, + fallback + ); + expect(selected?.credentialSource).toBe("market:selection"); + expect(resolveModelForMessage(selected).accountId).toBeUndefined(); + }); + it("resolves hosted models from the session while excluding own-key routing", () => { + expect( + resolveModelForMessage( + selectionFromSession({ ...session, keySource: "hosted_key" }, fallback) + ) + ).toEqual({ model: "session-model", accountId: undefined }); + }); +}); diff --git a/src/util/session/selectionFromSession.ts b/src/util/session/selectionFromSession.ts index c2d7687e77..9f4dfc9fc6 100644 --- a/src/util/session/selectionFromSession.ts +++ b/src/util/session/selectionFromSession.ts @@ -16,9 +16,11 @@ export function selectionFromSession( session: Session | undefined, fallback: LastModelSelection | null ): LastModelSelection | null { - if (!session) return fallback; + // A model/source pair is one identity. Never fill an existing session's + // missing account or routing fields from an unrelated creator preference. + if (!session?.model) return fallback; - const keySource = session.keySource ?? fallback?.keySource; + const keySource = session.keySource; // Rust persists market sessions with `listingModel` written into // `code_sessions.model`, so we can read either as the market `model` // identifier without a separate column. @@ -26,20 +28,11 @@ export function selectionFromSession( return { keySource, - model: isHosted ? undefined : (session.model ?? fallback?.model), - listingModel: isHosted - ? (session.model ?? fallback?.listingModel) - : undefined, - selectedAccountId: session.accountId ?? fallback?.selectedAccountId, - cliAgentType: session.cliAgentType ?? fallback?.cliAgentType, - tier: session.tier ?? fallback?.tier, - // Display-only fields: carry forward from fallback so the UI side - // preserves whatever it last rendered. - listingModelDisplay: fallback?.listingModelDisplay, - listingModelType: fallback?.listingModelType, - listingName: fallback?.listingName, - selectedSourceLabel: fallback?.selectedSourceLabel, - selectedSourceModelType: fallback?.selectedSourceModelType, - provider: fallback?.provider, + model: isHosted ? undefined : session.model, + listingModel: isHosted ? session.model : undefined, + selectedAccountId: session.accountId, + cliAgentType: session.cliAgentType, + tier: session.tier, + credentialSource: session.credentialSource, }; } diff --git a/src/util/variantEditOptions.ts b/src/util/variantEditOptions.ts index 0892844800..0da20efcb4 100644 --- a/src/util/variantEditOptions.ts +++ b/src/util/variantEditOptions.ts @@ -18,7 +18,9 @@ import { computeSeedDefaultVariant } from "./defaultModelVariant"; import { MODEL_REASONING_LEVEL, type ModelReasoningLevel, - parseModelVariant, + type ResolvedModelVariantFields, + resolveModelVariantFields, + toModelReasoningLevel, } from "./modelVariants"; import { getModelEffortBaseModel, @@ -95,35 +97,17 @@ interface IndexedVariant { fast: boolean; } -function indexVariants(modelIds: readonly string[]): IndexedVariant[] { - const out: IndexedVariant[] = []; - for (const modelId of modelIds) { - const parsed = parseModelVariant(modelId); - if (parsed) { - // `thinking` is the parsed thinking flag (Anthropic extended - // thinking). Effort level is independent — a variant can have - // a reasoning level without being a thinking variant. Variants - // with no parsed level (e.g. `claude-opus-4-6-thinking`, - // `composer-2.5-fast`) are surfaced as the `Default` effort row - // only when concrete effort variants do not cover that combination. - out.push({ - modelId, - thinking: parsed.thinking, - level: parsed.reasoning ?? MODEL_REASONING_LEVEL.BASELINE, - fast: parsed.fast, - }); - continue; - } - // Unparsed ids (e.g. `claude-sonnet-4-6`) are treated as the - // unsuffixed Default variant. - out.push({ - modelId, - thinking: false, - level: MODEL_REASONING_LEVEL.BASELINE, - fast: false, - }); - } - return out; +function indexVariants( + variants: readonly ResolvedModelVariantFields[] +): IndexedVariant[] { + return variants.map((variant) => ({ + modelId: variant.model, + thinking: variant.thinking ?? false, + level: + toModelReasoningLevel(variant.reasoning) ?? + MODEL_REASONING_LEVEL.BASELINE, + fast: variant.fast, + })); } function selectionKey(selection: { @@ -140,11 +124,17 @@ function selectionKey(selection: { } export function buildVariantEditOptions( - modelIds: readonly string[] + modelIds: readonly string[], + metadata: readonly ResolvedModelVariantFields[] = [] ): VariantEditOptions { - const indexed = selectableModelVariants( - modelIds.map((model) => ({ model })) - ).flatMap(({ model }) => indexVariants([model])); + const metadataById = new Map( + metadata.map((variant) => [variant.model, variant]) + ); + const resolve = (modelId: string) => + resolveModelVariantFields(modelId, metadataById.get(modelId)); + const effortBase = (modelId: string) => + metadataById.get(modelId)?.base_model ?? getModelEffortBaseModel(modelId); + const indexed = indexVariants(selectableModelVariants(modelIds.map(resolve))); // Effort levels are collected across BOTH thinking and non-thinking // variants. After the parser split, a non-thinking Claude variant @@ -197,14 +187,14 @@ export function buildVariantEditOptions( // Previously saved bare family ids resolve through the same seed rule // as family selection, rather than adding an invented effort rung. if ( - !parseModelVariant(modelId)?.reasoning && + !resolve(modelId).reasoning && !indexed.some((variant) => variant.modelId === modelId) ) { - const original = parseModelVariant(modelId); - const base = getModelEffortBaseModel(modelId); + const original = resolve(modelId); + const base = effortBase(modelId); const candidates = indexed.filter((variant) => { return ( - getModelEffortBaseModel(variant.modelId) === base && + effortBase(variant.modelId) === base && variant.thinking === (original?.thinking ?? false) && variant.fast === (original?.fast ?? false) ); @@ -220,18 +210,13 @@ export function buildVariantEditOptions( ); if (resolved) return parseSelection(resolved); } - const parsed = parseModelVariant(modelId); - if (!parsed) { - return { - thinking: false, - level: MODEL_REASONING_LEVEL.BASELINE, - fast: false, - }; - } + const variant = resolve(modelId); return { - thinking: parsed.thinking, - level: parsed.reasoning ?? MODEL_REASONING_LEVEL.BASELINE, - fast: parsed.fast, + thinking: variant.thinking ?? false, + level: + toModelReasoningLevel(variant.reasoning) ?? + MODEL_REASONING_LEVEL.BASELINE, + fast: variant.fast, }; };