diff --git a/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts index 6620541904..095f74dc36 100644 --- a/studio/frontend/src/features/chat/api/chat-adapter.ts +++ b/studio/frontend/src/features/chat/api/chat-adapter.ts @@ -1520,11 +1520,24 @@ function localRowToCandidate( * the same variants API the picker card uses and the smallest complete, * auto-loadable quant is chosen. Returns null when no quant can be resolved. */ +// Unknown sizes (0) sort last so a sizeless row can't shadow a real one. +function sizeOrUnknownBytes(bytes?: number | null): number { + return bytes && bytes > 0 ? bytes : Number.MAX_SAFE_INTEGER; +} + +type ResolvedLocalCandidate = { + candidate: AutoLoadCandidate; + /** Size of what would actually load: the resolved quant's own size for a + * multi-quant folder (whose row size_bytes SUMS every quant), else the + * row size. Orders the smallest-first cascade. */ + sizeBytes: number; +}; + async function resolveLocalRowCandidate( row: LocalModelInfo, rememberedVariant: string | null = null, isSkippedCandidate?: (candidate: AutoLoadCandidate) => boolean, -): Promise { +): Promise { const isGguf = row.model_format === "gguf"; if (row.capabilities?.requires_variant === true) { // Only GGUF folders have an automatic quant resolution path. @@ -1550,12 +1563,15 @@ async function resolveLocalRowCandidate( for (const entry of downloaded) { const candidate = localRowToCandidate(row, entry.quant); if (isSkippedCandidate?.(candidate)) continue; - return candidate; + return { candidate, sizeBytes: sizeOrUnknownBytes(entry.size_bytes) }; } return null; } } - return localRowToCandidate(row, isGguf ? rememberedVariant : null); + return { + candidate: localRowToCandidate(row, isGguf ? rememberedVariant : null), + sizeBytes: sizeOrUnknownBytes(row.size_bytes), + }; } /** Resolve a remembered local model against current backend inventory. */ @@ -1952,17 +1968,27 @@ export async function autoLoadOnDeviceModel(): Promise<{ const modelRepos = allModelRepos.filter(isAutoLoadableCachedRepo); const localRows = allLocalRows.filter(isAutoLoadableLocalRow); // Dedupe candidates that resolve to the SAME load target (e.g. a custom - // scan folder pointing into an HF cache). Keyed on load targets and - // on-disk paths only: a shared model_id does not mean the same files, and - // a distinct local copy must stay available when the cached copy fails. + // scan folder pointing into an HF cache). Keyed on kind + load target / + // on-disk path: a shared model_id does not mean the same files (a distinct + // local copy stays available when the cached copy fails), and a folder + // emitting both GGUF and safetensors rows shares a path while holding two + // different models, so the format is part of the key. const seenLoadTargets = new Set(); - const markSeen = (...values: (string | null | undefined)[]): void => { + const markSeen = ( + kind: LastLocalModelKind, + ...values: (string | null | undefined)[] + ): void => { for (const value of values) { - if (value) seenLoadTargets.add(value.toLowerCase()); + if (value) seenLoadTargets.add(`${kind}:${value.toLowerCase()}`); } }; - const isSeen = (...values: (string | null | undefined)[]): boolean => - values.some((value) => !!value && seenLoadTargets.has(value.toLowerCase())); + const isSeen = ( + kind: LastLocalModelKind, + ...values: (string | null | undefined)[] + ): boolean => + values.some( + (value) => !!value && seenLoadTargets.has(`${kind}:${value.toLowerCase()}`), + ); try { if (lastLoaded) { @@ -1977,10 +2003,9 @@ export async function autoLoadOnDeviceModel(): Promise<{ // managed-cache remembered path). let rememberedCandidate: AutoLoadCandidate | null = null; try { - rememberedCandidate = await resolveLocalRowCandidate( - row, - lastLoaded.ggufVariant, - ); + rememberedCandidate = + (await resolveLocalRowCandidate(row, lastLoaded.ggufVariant)) + ?.candidate ?? null; if (rememberedCandidate) { toast("Loading last used model…", { id: toastId, @@ -2006,7 +2031,7 @@ export async function autoLoadOnDeviceModel(): Promise<{ } else if (lastLoaded.kind === "gguf") { const repo = findCachedRepo(ggufRepos, lastLoaded.id); if (repo && lastLoaded.ggufVariant) { - markSeen(repo.load_id || repo.repo_id, repo.cache_path); + markSeen("gguf", repo.load_id || repo.repo_id, repo.cache_path); try { const variants = await listGgufVariants(repo.repo_id, undefined, { preferLocalCache: true, @@ -2051,7 +2076,7 @@ export async function autoLoadOnDeviceModel(): Promise<{ } else { const repo = findCachedRepo(modelRepos, lastLoaded.id); if (repo) { - markSeen(repo.load_id || repo.repo_id, repo.cache_path); + markSeen("model", repo.load_id || repo.repo_id, repo.cache_path); try { toast("Loading last used model…", { id: toastId, @@ -2093,53 +2118,80 @@ export async function autoLoadOnDeviceModel(): Promise<{ type FallbackCandidate = | { type: "cached-gguf"; repo: CachedGgufRepo; sizeBytes: number } | { type: "cached-model"; repo: CachedModelRepo; sizeBytes: number } - | { type: "local"; row: LocalModelInfo; sizeBytes: number }; - // Unknown sizes (0) sort last so a sizeless row can't shadow a real one. - const sizeOrUnknown = (bytes?: number | null): number => - bytes && bytes > 0 ? bytes : Number.MAX_SAFE_INTEGER; + | { + type: "local"; + row: LocalModelInfo; + candidate: AutoLoadCandidate; + sizeBytes: number; + }; const bySizeAsc = (a: FallbackCandidate, b: FallbackCandidate): number => a.sizeBytes - b.sizeBytes; - // Directory-based GGUF rows resolve a quant automatically below; only - // non-GGUF variant-requiring rows have no background resolution path. + const isSkippedAutoLoadCandidate = (c: AutoLoadCandidate): boolean => + skippedAutoLoadCandidates.has( + autoLoadCandidateKey(c.kind, c.id, c.ggufVariant), + ); + // Directory-based GGUF rows resolve a quant automatically; only non-GGUF + // variant-requiring rows have no background resolution path. const cascadeLocalRows = localRows.filter( (row) => row.model_format === "gguf" || row.capabilities?.requires_variant !== true, ); + // Resolve local candidates BEFORE ordering: a multi-quant folder's row + // size_bytes sums every quant in it, so the cascade must order on the + // resolved quant's own size or a folder with a small quant would lose to + // a larger single-quant model. + const localEntries = ( + await Promise.all( + cascadeLocalRows.map( + async (row): Promise => { + try { + const resolved = await resolveLocalRowCandidate( + row, + null, + isSkippedAutoLoadCandidate, + ); + if (!resolved) return null; + return { + type: "local" as const, + row, + candidate: resolved.candidate, + sizeBytes: resolved.sizeBytes, + }; + } catch { + hadNonTrustFailure = true; + return null; + } + }, + ), + ) + ).filter((entry): entry is FallbackCandidate => entry !== null); const ggufGroup: FallbackCandidate[] = [ ...ggufRepos.map((repo) => ({ type: "cached-gguf" as const, repo, - sizeBytes: sizeOrUnknown(repo.size_bytes), + sizeBytes: sizeOrUnknownBytes(repo.size_bytes), })), - ...cascadeLocalRows - .filter((row) => row.model_format === "gguf") - .map((row) => ({ - type: "local" as const, - row, - sizeBytes: sizeOrUnknown(row.size_bytes), - })), + ...localEntries.filter( + (entry) => entry.type === "local" && entry.candidate.kind === "gguf", + ), ].sort(bySizeAsc); const modelGroup: FallbackCandidate[] = [ ...modelRepos.map((repo) => ({ type: "cached-model" as const, repo, - sizeBytes: sizeOrUnknown(repo.size_bytes), + sizeBytes: sizeOrUnknownBytes(repo.size_bytes), })), - ...cascadeLocalRows - .filter((row) => row.model_format !== "gguf") - .map((row) => ({ - type: "local" as const, - row, - sizeBytes: sizeOrUnknown(row.size_bytes), - })), + ...localEntries.filter( + (entry) => entry.type === "local" && entry.candidate.kind === "model", + ), ].sort(bySizeAsc); for (const candidate of [...ggufGroup, ...modelGroup]) { if (loadAttempts >= MAX_AUTO_LOAD_ATTEMPTS) break; if (candidate.type === "cached-gguf") { const repo = candidate.repo; - markSeen(repo.load_id || repo.repo_id, repo.cache_path); + markSeen("gguf", repo.load_id || repo.repo_id, repo.cache_path); try { const variants = await listGgufVariants(repo.repo_id, undefined, { preferLocalCache: true, @@ -2181,7 +2233,7 @@ export async function autoLoadOnDeviceModel(): Promise<{ } if (candidate.type === "cached-model") { const repo = candidate.repo; - markSeen(repo.load_id || repo.repo_id, repo.cache_path); + markSeen("model", repo.load_id || repo.repo_id, repo.cache_path); if ( skippedAutoLoadCandidates.has( autoLoadCandidateKey("model", repo.repo_id), @@ -2210,30 +2262,15 @@ export async function autoLoadOnDeviceModel(): Promise<{ continue; } const row = candidate.row; - if (isSeen(row.load_id, row.id, row.path)) { + const localCandidate = candidate.candidate; + if (isSeen(localCandidate.kind, row.load_id, row.id, row.path)) { + continue; + } + markSeen(localCandidate.kind, row.load_id, row.id, row.path); + if (isSkippedAutoLoadCandidate(localCandidate)) { continue; } - markSeen(row.load_id, row.id, row.path); try { - const localCandidate = await resolveLocalRowCandidate(row, null, (c) => - skippedAutoLoadCandidates.has( - autoLoadCandidateKey(c.kind, c.id, c.ggufVariant), - ), - ); - if (!localCandidate) { - continue; - } - if ( - skippedAutoLoadCandidates.has( - autoLoadCandidateKey( - localCandidate.kind, - localCandidate.id, - localCandidate.ggufVariant, - ), - ) - ) { - continue; - } if (await loadAutoLoadCandidate(localCandidate)) { return { loaded: true, blockedByTrustRemoteCode: false }; } diff --git a/tests/studio/test_model_picker_contracts.py b/tests/studio/test_model_picker_contracts.py index 1c18ca39ea..f1c7037e1a 100644 --- a/tests/studio/test_model_picker_contracts.py +++ b/tests/studio/test_model_picker_contracts.py @@ -572,8 +572,12 @@ def test_autoload_deduplicates_cached_and_local_candidates(): available when the cached copy fails or has no usable quant.""" auto_load = _autoload_section() assert "const seenLoadTargets = new Set()" in auto_load - assert "markSeen(repo.load_id || repo.repo_id, repo.cache_path)" in auto_load - assert "isSeen(row.load_id, row.id, row.path)" in auto_load + # Keys carry the model kind: a folder emitting both GGUF and safetensors + # rows shares a path while holding two different models. + assert "seenLoadTargets.add(`${kind}:${value.toLowerCase()}`)" in auto_load + assert 'markSeen("gguf", repo.load_id || repo.repo_id, repo.cache_path)' in auto_load + assert 'markSeen("model", repo.load_id || repo.repo_id, repo.cache_path)' in auto_load + assert "isSeen(localCandidate.kind, row.load_id, row.id, row.path)" in auto_load # The repo-id-based dedupe that shadowed distinct local copies is gone. assert "isSeen(row.load_id, row.id, row.path, row.model_id)" not in auto_load assert "markSeen(repo.repo_id," not in auto_load @@ -590,7 +594,7 @@ def test_local_quant_resolution_skips_failed_quants(): assert "if (isSkippedCandidate?.(candidate)) continue;" in resolve_fn # The fallback loop feeds the skip set into resolution. auto_load = _autoload_section() - assert "await resolveLocalRowCandidate(row, null, (c) =>" in auto_load + assert "isSkippedAutoLoadCandidate," in auto_load def test_autoload_trust_guard_still_blocks_background_loads(): @@ -671,7 +675,7 @@ def test_directory_gguf_rows_resolve_variant_like_picker(): # The cascade must keep directory GGUF rows as candidates. auto_load = src.split("async function autoLoadOnDeviceModel", 1)[1] assert 'row.model_format === "gguf" ||' in auto_load - assert "await resolveLocalRowCandidate(row, null, (c) =>" in auto_load + assert "await resolveLocalRowCandidate(" in auto_load def test_remembered_local_failure_does_not_block_folder_fallback(): @@ -686,3 +690,21 @@ def test_remembered_local_failure_does_not_block_folder_fallback(): "markSeen(" not in remembered_block ), "remembered-local retry must not pre-mark the row as deduped" assert "rememberedCandidate?.ggufVariant ?? lastLoaded.ggufVariant" in remembered_block + + +def test_local_fallback_orders_by_resolved_quant_size(): + """A GGUF folder row's size_bytes sums every quant in the folder, so the + smallest-first cascade must order local candidates by the resolved + quant's own size; otherwise a folder with a small quant loses to a + larger single-quant model.""" + src = _read("features/chat/api/chat-adapter.ts") + resolve_fn = src.split("async function resolveLocalRowCandidate", 1)[1] + resolve_fn = resolve_fn.split("\nfunction ", 1)[0] + assert "sizeBytes: sizeOrUnknownBytes(entry.size_bytes)" in resolve_fn + auto_load = _autoload_section() + # Local candidates are resolved BEFORE the groups are sorted. + assert "const localEntries = (" in auto_load + assert auto_load.index("const localEntries = (") < auto_load.index( + "const ggufGroup: FallbackCandidate[]" + ) + assert "sizeBytes: resolved.sizeBytes" in auto_load