refactor: use Immer for immutable state updates (#527)

* refactor: use Immer for immutable state updates

Replace Set/Map with Record types and use Immer's produce() for
immutable state updates. This fixes mutation bugs where .add()/.set()
were mutating state before copying (e.g., `new Set(prev.add(id))`).

Changes:
- Add immer dependency
- Convert Set<string> to Record<string, true> (sparse hash set pattern)
- Convert Map<K,V> to Record<K,V>
- Use produce() for all state mutations
- Update prop types in child components

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>

* move deps back to where they started

---------

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
Co-authored-by: CJ Pais <cj@cjpais.com>
This commit is contained in:
Josh Ribakoff 2026-01-18 16:47:43 -08:00 committed by GitHub
parent c84e863423
commit f9d2aa68c3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 97 additions and 88 deletions

View file

@ -18,6 +18,7 @@
"@tauri-apps/plugin-store": "~2.4.1", "@tauri-apps/plugin-store": "~2.4.1",
"@tauri-apps/plugin-updater": "~2.9.0", "@tauri-apps/plugin-updater": "~2.9.0",
"i18next": "^25.7.2", "i18next": "^25.7.2",
"immer": "^11.1.3",
"lucide-react": "^0.542.0", "lucide-react": "^0.542.0",
"react": "^18.3.1", "react": "^18.3.1",
"react-dom": "^18.3.1", "react-dom": "^18.3.1",
@ -493,6 +494,8 @@
"ignore": ["ignore@7.0.5", "", {}, "sha512-Hs59xBNfUIunMFgWAbGX5cq6893IbWg4KnrjbYwX3tx0ztorVgTDA6B2sxf8ejHJ4wz8BqGUMYlnzNBer5NvGg=="], "ignore": ["ignore@7.0.5", "", {}, "sha512-Hs59xBNfUIunMFgWAbGX5cq6893IbWg4KnrjbYwX3tx0ztorVgTDA6B2sxf8ejHJ4wz8BqGUMYlnzNBer5NvGg=="],
"immer": ["immer@11.1.3", "", {}, "sha512-6jQTc5z0KJFtr1UgFpIL3N9XSC3saRaI9PwWtzM2pSqkNGtiNkYY2OSwkOGDK2XcTRcLb1pi/aNkKZz0nxVH4Q=="],
"import-fresh": ["import-fresh@3.3.1", "", { "dependencies": { "parent-module": "^1.0.0", "resolve-from": "^4.0.0" } }, "sha512-TR3KfrTZTYLPB6jUjfx6MF9WcWrHL9su5TObK4ZkYgBdWKPOFoSoQIdEuTuR82pmtxH2spWG9h6etwfr1pLBqQ=="], "import-fresh": ["import-fresh@3.3.1", "", { "dependencies": { "parent-module": "^1.0.0", "resolve-from": "^4.0.0" } }, "sha512-TR3KfrTZTYLPB6jUjfx6MF9WcWrHL9su5TObK4ZkYgBdWKPOFoSoQIdEuTuR82pmtxH2spWG9h6etwfr1pLBqQ=="],
"imurmurhash": ["imurmurhash@0.1.4", "", {}, "sha512-JmXMZ6wuvDmLiHEml9ykzqO6lwFbof0GG4IkcGaENdCRDDmMVnny7s5HsIgHCbaq0w2MyPhDqkhTUgS2LU2PHA=="], "imurmurhash": ["imurmurhash@0.1.4", "", {}, "sha512-JmXMZ6wuvDmLiHEml9ykzqO6lwFbof0GG4IkcGaENdCRDDmMVnny7s5HsIgHCbaq0w2MyPhDqkhTUgS2LU2PHA=="],

View file

@ -31,6 +31,7 @@
"react-select": "^5.8.0", "react-select": "^5.8.0",
"tauri-plugin-macos-permissions-api": "2.3.0", "tauri-plugin-macos-permissions-api": "2.3.0",
"i18next": "^25.7.2", "i18next": "^25.7.2",
"immer": "^11.1.3",
"lucide-react": "^0.542.0", "lucide-react": "^0.542.0",
"react": "^18.3.1", "react": "^18.3.1",
"react-dom": "^18.3.1", "react-dom": "^18.3.1",

View file

@ -16,8 +16,8 @@ interface DownloadStats {
} }
interface DownloadProgressDisplayProps { interface DownloadProgressDisplayProps {
downloadProgress: Map<string, DownloadProgress>; downloadProgress: Record<string, DownloadProgress>;
downloadStats: Map<string, DownloadStats>; downloadStats: Record<string, DownloadStats>;
className?: string; className?: string;
} }
@ -26,14 +26,13 @@ const DownloadProgressDisplay: React.FC<DownloadProgressDisplayProps> = ({
downloadStats, downloadStats,
className = "", className = "",
}) => { }) => {
if (downloadProgress.size === 0) { const progressValues = Object.values(downloadProgress);
if (progressValues.length === 0) {
return null; return null;
} }
const progressData: ProgressData[] = Array.from( const progressData: ProgressData[] = progressValues.map((progress) => {
downloadProgress.values(), const stats = downloadStats[progress.model_id];
).map((progress) => {
const stats = downloadStats.get(progress.model_id);
return { return {
id: progress.model_id, id: progress.model_id,
percentage: progress.percentage, percentage: progress.percentage,
@ -45,7 +44,7 @@ const DownloadProgressDisplay: React.FC<DownloadProgressDisplayProps> = ({
<ProgressBar <ProgressBar
progress={progressData} progress={progressData}
className={className} className={className}
showSpeed={downloadProgress.size === 1} showSpeed={progressValues.length === 1}
size="medium" size="medium"
/> />
); );

View file

@ -18,7 +18,7 @@ interface DownloadProgress {
interface ModelDropdownProps { interface ModelDropdownProps {
models: ModelInfo[]; models: ModelInfo[];
currentModelId: string; currentModelId: string;
downloadProgress: Map<string, DownloadProgress>; downloadProgress: Record<string, DownloadProgress>;
onModelSelect: (modelId: string) => void; onModelSelect: (modelId: string) => void;
onModelDownload: (modelId: string) => void; onModelDownload: (modelId: string) => void;
onModelDelete: (modelId: string) => Promise<void>; onModelDelete: (modelId: string) => Promise<void>;
@ -52,14 +52,14 @@ const ModelDropdown: React.FC<ModelDropdownProps> = ({
}; };
const handleModelClick = (modelId: string) => { const handleModelClick = (modelId: string) => {
if (downloadProgress.has(modelId)) { if (modelId in downloadProgress) {
return; // Don't allow interaction while downloading return; // Don't allow interaction while downloading
} }
onModelSelect(modelId); onModelSelect(modelId);
}; };
const handleDownloadClick = (modelId: string) => { const handleDownloadClick = (modelId: string) => {
if (downloadProgress.has(modelId)) { if (modelId in downloadProgress) {
return; // Don't allow interaction while downloading return; // Don't allow interaction while downloading
} }
onModelDownload(modelId); onModelDownload(modelId);
@ -158,8 +158,8 @@ const ModelDropdown: React.FC<ModelDropdownProps> = ({
: t("modelSelector.downloadModels")} : t("modelSelector.downloadModels")}
</div> </div>
{downloadableModels.map((model) => { {downloadableModels.map((model) => {
const isDownloading = downloadProgress.has(model.id); const isDownloading = model.id in downloadProgress;
const progress = downloadProgress.get(model.id); const progress = downloadProgress[model.id];
return ( return (
<div <div

View file

@ -1,6 +1,7 @@
import React, { useState, useRef, useEffect } from "react"; import React, { useState, useRef, useEffect } from "react";
import { useTranslation } from "react-i18next"; import { useTranslation } from "react-i18next";
import { listen } from "@tauri-apps/api/event"; import { listen } from "@tauri-apps/api/event";
import { produce } from "immer";
import { commands, type ModelInfo } from "@/bindings"; import { commands, type ModelInfo } from "@/bindings";
import { getTranslatedModelName } from "../../lib/utils/modelTranslation"; import { getTranslatedModelName } from "../../lib/utils/modelTranslation";
import ModelStatusButton from "./ModelStatusButton"; import ModelStatusButton from "./ModelStatusButton";
@ -48,15 +49,15 @@ const ModelSelector: React.FC<ModelSelectorProps> = ({ onError }) => {
const [modelStatus, setModelStatus] = useState<ModelStatus>("unloaded"); const [modelStatus, setModelStatus] = useState<ModelStatus>("unloaded");
const [modelError, setModelError] = useState<string | null>(null); const [modelError, setModelError] = useState<string | null>(null);
const [modelDownloadProgress, setModelDownloadProgress] = useState< const [modelDownloadProgress, setModelDownloadProgress] = useState<
Map<string, DownloadProgress> Record<string, DownloadProgress>
>(new Map()); >({});
const [showModelDropdown, setShowModelDropdown] = useState(false); const [showModelDropdown, setShowModelDropdown] = useState(false);
const [downloadStats, setDownloadStats] = useState< const [downloadStats, setDownloadStats] = useState<
Map<string, DownloadStats> Record<string, DownloadStats>
>(new Map()); >({});
const [extractingModels, setExtractingModels] = useState<Set<string>>( const [extractingModels, setExtractingModels] = useState<
new Set(), Record<string, true>
); >({});
const dropdownRef = useRef<HTMLDivElement>(null); const dropdownRef = useRef<HTMLDivElement>(null);
@ -97,53 +98,52 @@ const ModelSelector: React.FC<ModelSelectorProps> = ({ onError }) => {
"model-download-progress", "model-download-progress",
(event) => { (event) => {
const progress = event.payload; const progress = event.payload;
setModelDownloadProgress((prev) => { setModelDownloadProgress(
const newMap = new Map(prev); produce((downloadProgress) => {
newMap.set(progress.model_id, progress); downloadProgress[progress.model_id] = progress;
return newMap; }),
}); );
setModelStatus("downloading"); setModelStatus("downloading");
// Update download stats for speed calculation // Update download stats for speed calculation
const now = Date.now(); const now = Date.now();
setDownloadStats((prev) => { setDownloadStats(
const current = prev.get(progress.model_id); produce((stats) => {
const newStats = new Map(prev); const current = stats[progress.model_id];
if (!current) { if (!current) {
// First progress update - initialize // First progress update - initialize
newStats.set(progress.model_id, { stats[progress.model_id] = {
startTime: now, startTime: now,
lastUpdate: now,
totalDownloaded: progress.downloaded,
speed: 0,
});
} else {
// Calculate speed over last few seconds
const timeDiff = (now - current.lastUpdate) / 1000; // seconds
const bytesDiff = progress.downloaded - current.totalDownloaded;
if (timeDiff > 0.5) {
// Update speed every 500ms
const currentSpeed = bytesDiff / (1024 * 1024) / timeDiff; // MB/s
// Smooth the speed with exponential moving average, but ensure positive values
const validCurrentSpeed = Math.max(0, currentSpeed);
const smoothedSpeed =
current.speed > 0
? current.speed * 0.8 + validCurrentSpeed * 0.2
: validCurrentSpeed;
newStats.set(progress.model_id, {
startTime: current.startTime,
lastUpdate: now, lastUpdate: now,
totalDownloaded: progress.downloaded, totalDownloaded: progress.downloaded,
speed: Math.max(0, smoothedSpeed), speed: 0,
}); };
} } else {
} // Calculate speed over last few seconds
const timeDiff = (now - current.lastUpdate) / 1000; // seconds
const bytesDiff = progress.downloaded - current.totalDownloaded;
return newStats; if (timeDiff > 0.5) {
}); // Update speed every 500ms
const currentSpeed = bytesDiff / (1024 * 1024) / timeDiff; // MB/s
// Smooth the speed with exponential moving average, but ensure positive values
const validCurrentSpeed = Math.max(0, currentSpeed);
const smoothedSpeed =
current.speed > 0
? current.speed * 0.8 + validCurrentSpeed * 0.2
: validCurrentSpeed;
stats[progress.model_id] = {
startTime: current.startTime,
lastUpdate: now,
totalDownloaded: progress.downloaded,
speed: Math.max(0, smoothedSpeed),
};
}
}
}),
);
}, },
); );
@ -152,16 +152,16 @@ const ModelSelector: React.FC<ModelSelectorProps> = ({ onError }) => {
"model-download-complete", "model-download-complete",
(event) => { (event) => {
const modelId = event.payload; const modelId = event.payload;
setModelDownloadProgress((prev) => { setModelDownloadProgress(
const newMap = new Map(prev); produce((progress) => {
newMap.delete(modelId); delete progress[modelId];
return newMap; }),
}); );
setDownloadStats((prev) => { setDownloadStats(
const newStats = new Map(prev); produce((stats) => {
newStats.delete(modelId); delete stats[modelId];
return newStats; }),
}); );
loadModels(); // Refresh models list loadModels(); // Refresh models list
// Auto-select the newly downloaded model (skip if recording in progress) // Auto-select the newly downloaded model (skip if recording in progress)
@ -181,7 +181,11 @@ const ModelSelector: React.FC<ModelSelectorProps> = ({ onError }) => {
"model-extraction-started", "model-extraction-started",
(event) => { (event) => {
const modelId = event.payload; const modelId = event.payload;
setExtractingModels((prev) => new Set(prev.add(modelId))); setExtractingModels(
produce((extracting) => {
extracting[modelId] = true;
}),
);
setModelStatus("extracting"); setModelStatus("extracting");
}, },
); );
@ -190,11 +194,11 @@ const ModelSelector: React.FC<ModelSelectorProps> = ({ onError }) => {
"model-extraction-completed", "model-extraction-completed",
(event) => { (event) => {
const modelId = event.payload; const modelId = event.payload;
setExtractingModels((prev) => { setExtractingModels(
const next = new Set(prev); produce((extracting) => {
next.delete(modelId); delete extracting[modelId];
return next; }),
}); );
loadModels(); // Refresh models list loadModels(); // Refresh models list
// Auto-select the newly extracted model (skip if recording in progress) // Auto-select the newly extracted model (skip if recording in progress)
@ -214,11 +218,11 @@ const ModelSelector: React.FC<ModelSelectorProps> = ({ onError }) => {
error: string; error: string;
}>("model-extraction-failed", (event) => { }>("model-extraction-failed", (event) => {
const modelId = event.payload.model_id; const modelId = event.payload.model_id;
setExtractingModels((prev) => { setExtractingModels(
const next = new Set(prev); produce((extracting) => {
next.delete(modelId); delete extracting[modelId];
return next; }),
}); );
setModelError(`Failed to extract model: ${event.payload.error}`); setModelError(`Failed to extract model: ${event.payload.error}`);
setModelStatus("error"); setModelStatus("error");
}); });
@ -329,9 +333,10 @@ const ModelSelector: React.FC<ModelSelectorProps> = ({ onError }) => {
}; };
const getModelDisplayText = (): string => { const getModelDisplayText = (): string => {
if (extractingModels.size > 0) { const extractingKeys = Object.keys(extractingModels);
if (extractingModels.size === 1) { if (extractingKeys.length > 0) {
const [modelId] = Array.from(extractingModels); if (extractingKeys.length === 1) {
const modelId = extractingKeys[0];
const model = models.find((m) => m.id === modelId); const model = models.find((m) => m.id === modelId);
const modelName = model const modelName = model
? getTranslatedModelName(model, t) ? getTranslatedModelName(model, t)
@ -339,14 +344,15 @@ const ModelSelector: React.FC<ModelSelectorProps> = ({ onError }) => {
return t("modelSelector.extracting", { modelName }); return t("modelSelector.extracting", { modelName });
} else { } else {
return t("modelSelector.extractingMultiple", { return t("modelSelector.extractingMultiple", {
count: extractingModels.size, count: extractingKeys.length,
}); });
} }
} }
if (modelDownloadProgress.size > 0) { const progressValues = Object.values(modelDownloadProgress);
if (modelDownloadProgress.size === 1) { if (progressValues.length > 0) {
const [progress] = Array.from(modelDownloadProgress.values()); if (progressValues.length === 1) {
const progress = progressValues[0];
const percentage = Math.max( const percentage = Math.max(
0, 0,
Math.min(100, Math.round(progress.percentage)), Math.min(100, Math.round(progress.percentage)),
@ -354,7 +360,7 @@ const ModelSelector: React.FC<ModelSelectorProps> = ({ onError }) => {
return t("modelSelector.downloading", { percentage }); return t("modelSelector.downloading", { percentage });
} else { } else {
return t("modelSelector.downloadingMultiple", { return t("modelSelector.downloadingMultiple", {
count: modelDownloadProgress.size, count: progressValues.length,
}); });
} }
} }