From 21a548d57505a33670ac022a1fa788bcacb64d36 Mon Sep 17 00:00:00 2001 From: Bhagirath Mehta Date: Fri, 18 Sep 2026 13:16:16 -0500 Subject: [PATCH 1/6] Restore model download cancellation parity Carry the legacy AbortSignal and threading.Event contracts into the v2 bindings so callers can stop native downloads instead of receiving progress-only APIs. Preserve completion races and callback exceptions while mapping signal-driven JavaScript cancellation to AbortError. Files changed: - sdk_v2/js/native/src/model.cc: bridge AbortSignal state into the native progress callback - sdk_v2/js/src and test: restore overloads and verify cancellation semantics - sdk_v2/python/src and test: restore cancel_event and preserve callback errors - sdk_v2/python/README.md: document cooperative download cancellation Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 2d2c76d2-0c86-4672-b91b-3fcb54ed7906 --- sdk_v2/js/native/src/model.cc | 99 ++++++++++++++++--- sdk_v2/js/src/detail/native.ts | 2 +- sdk_v2/js/src/imodel.ts | 5 +- sdk_v2/js/src/model.ts | 49 ++++++++- sdk_v2/js/test/model-download.types.ts | 11 +++ sdk_v2/js/test/model-lifecycle.test.ts | 17 ++++ sdk_v2/js/test/model.test.ts | 23 +++++ sdk_v2/js/tsconfig.types.json | 2 +- sdk_v2/python/README.md | 13 +++ sdk_v2/python/src/foundry_local_sdk/imodel.py | 49 ++++++--- .../unit/test_model_download_cancellation.py | 99 +++++++++++++++++++ 11 files changed, 335 insertions(+), 34 deletions(-) create mode 100644 sdk_v2/js/test/model-download.types.ts create mode 100644 sdk_v2/python/test/unit/test_model_download_cancellation.py diff --git a/sdk_v2/js/native/src/model.cc b/sdk_v2/js/native/src/model.cc index 6d1844c9e..d948d3897 100644 --- a/sdk_v2/js/native/src/model.cc +++ b/sdk_v2/js/native/src/model.cc @@ -9,6 +9,7 @@ #include #include +#include #include #include #include @@ -283,27 +284,41 @@ namespace { class DownloadWorker : public Napi::AsyncWorker { public: DownloadWorker(Napi::Env env, foundry_local::IModel* impl, Napi::ObjectReference owner, - Napi::ThreadSafeFunction tsfn) + Napi::ThreadSafeFunction tsfn, std::shared_ptr> abort_requested, + Napi::ObjectReference abort_signal, Napi::FunctionReference abort_listener) : Napi::AsyncWorker(env), deferred_(Napi::Promise::Deferred::New(env)), impl_(impl), owner_(std::move(owner)), - tsfn_(std::move(tsfn)) {} + tsfn_(std::move(tsfn)), + abort_requested_(std::move(abort_requested)), + abort_signal_(std::move(abort_signal)), + abort_listener_(std::move(abort_listener)) {} Napi::Promise Promise() { return deferred_.Promise(); } void Execute() override { try { - auto progress_cb = tsfn_ ? std::function([this](float percent) { - // BlockingCall keeps backpressure on the worker thread: if JS is - // slow to drain the queue we'll wait rather than dropping reports. - // Callback return value is unused on the JS side; we always continue. - tsfn_.BlockingCall([percent](Napi::Env env, Napi::Function js_cb) { - js_cb.Call({Napi::Number::New(env, static_cast(percent))}); - }); - return 0; // 0 = continue per flProgressCallback contract. + const bool has_cancellation = abort_requested_ != nullptr; + auto progress_cb = (tsfn_ || has_cancellation) ? std::function([this](float percent) { + if (IsAbortRequested()) { + cancelled_by_signal_ = true; + return 1; + } + if (tsfn_) { + // BlockingCall keeps backpressure on the worker thread: if JS is slow to drain the queue we'll wait rather + // than dropping reports. + tsfn_.BlockingCall([percent](Napi::Env env, Napi::Function js_cb) { + js_cb.Call({Napi::Number::New(env, static_cast(percent))}); + }); + } + if (IsAbortRequested()) { + cancelled_by_signal_ = true; + return 1; + } + return 0; }) - : std::function(nullptr); + : std::function(nullptr); impl_->Download(std::move(progress_cb)); } catch (const foundry_local::Error& e) { err_code_ = static_cast(e.Code()); @@ -321,18 +336,20 @@ class DownloadWorker : public Napi::AsyncWorker { void OnOK() override { Napi::HandleScope scope(Env()); - ReleaseTsfn(); + CleanupJsReferences(); deferred_.Resolve(Env().Undefined()); } void OnError(const Napi::Error& /*unused*/) override { Napi::Env env = Env(); Napi::HandleScope scope(env); - ReleaseTsfn(); + CleanupJsReferences(); if (tagged_) { Napi::Error err = Napi::Error::New(env, err_msg_); Napi::Object value = err.Value(); - value.Set("name", Napi::String::New(env, "FoundryLocalError")); + const bool is_signal_cancellation = + cancelled_by_signal_ && err_code_ == FOUNDRY_LOCAL_ERROR_OPERATION_CANCELLED; + value.Set("name", Napi::String::New(env, is_signal_cancellation ? "AbortError" : "FoundryLocalError")); value.Set("code", Napi::Number::New(env, err_code_)); deferred_.Reject(value); } else { @@ -341,7 +358,21 @@ class DownloadWorker : public Napi::AsyncWorker { } private: - void ReleaseTsfn() { + bool IsAbortRequested() const { + return abort_requested_ != nullptr && abort_requested_->load(std::memory_order_acquire); + } + + void CleanupJsReferences() { + if (!abort_signal_.IsEmpty() && !abort_listener_.IsEmpty()) { + Napi::Object signal = abort_signal_.Value(); + Napi::Value remove_value = signal.Get("removeEventListener"); + if (remove_value.IsFunction()) { + remove_value.As().Call( + signal, {Napi::String::New(Env(), "abort"), abort_listener_.Value()}); + } + abort_listener_.Reset(); + abort_signal_.Reset(); + } if (tsfn_) { tsfn_.Release(); tsfn_ = Napi::ThreadSafeFunction(); @@ -352,9 +383,13 @@ class DownloadWorker : public Napi::AsyncWorker { foundry_local::IModel* impl_; Napi::ObjectReference owner_; Napi::ThreadSafeFunction tsfn_; + std::shared_ptr> abort_requested_; + Napi::ObjectReference abort_signal_; + Napi::FunctionReference abort_listener_; std::string err_msg_; int err_code_ = 0; bool tagged_ = false; + bool cancelled_by_signal_ = false; }; } // namespace @@ -378,8 +413,40 @@ Napi::Value Model::Download(const Napi::CallbackInfo& info) { return env.Undefined(); } + std::shared_ptr> abort_requested; + Napi::ObjectReference abort_signal; + Napi::FunctionReference abort_listener; + if (info.Length() >= 2 && !info[1].IsUndefined() && !info[1].IsNull()) { + if (!info[1].IsObject()) { + Napi::TypeError::New(env, "Model.download: signal must be an AbortSignal").ThrowAsJavaScriptException(); + return env.Undefined(); + } + Napi::Object signal = info[1].As(); + Napi::Value aborted = signal.Get("aborted"); + Napi::Value add_value = signal.Get("addEventListener"); + Napi::Value remove_value = signal.Get("removeEventListener"); + if (!aborted.IsBoolean() || !add_value.IsFunction() || !remove_value.IsFunction()) { + Napi::TypeError::New(env, "Model.download: signal must be an AbortSignal").ThrowAsJavaScriptException(); + return env.Undefined(); + } + + abort_requested = std::make_shared>(aborted.As().Value()); + Napi::Function listener = Napi::Function::New( + env, [abort_requested](const Napi::CallbackInfo&) { + abort_requested->store(true, std::memory_order_release); + }); + add_value.As().Call( + signal, {Napi::String::New(env, "abort"), listener}); + if (signal.Get("aborted").As().Value()) { + abort_requested->store(true, std::memory_order_release); + } + abort_signal = Napi::Persistent(signal); + abort_listener = Napi::Persistent(listener); + } + Napi::ObjectReference owner = Napi::Reference::New(manager_.Value(), 1); - auto* w = new DownloadWorker(env, impl_, std::move(owner), std::move(tsfn)); + auto* w = new DownloadWorker(env, impl_, std::move(owner), std::move(tsfn), std::move(abort_requested), + std::move(abort_signal), std::move(abort_listener)); Napi::Promise p = w->Promise(); w->Queue(); return p; diff --git a/sdk_v2/js/src/detail/native.ts b/sdk_v2/js/src/detail/native.ts index 4167acdd3..bd74857ac 100644 --- a/sdk_v2/js/src/detail/native.ts +++ b/sdk_v2/js/src/detail/native.ts @@ -92,7 +92,7 @@ export interface NativeModel { selectVariant(variant: NativeModel): void; load(): Promise; unload(): Promise; - download(progress?: (percent: number) => void): Promise; + download(progress?: (percent: number) => void, signal?: AbortSignal): Promise; removeFromCache(): void; } diff --git a/sdk_v2/js/src/imodel.ts b/sdk_v2/js/src/imodel.ts index fe5449df0..e653df933 100644 --- a/sdk_v2/js/src/imodel.ts +++ b/sdk_v2/js/src/imodel.ts @@ -19,7 +19,10 @@ export interface IModel { get capabilities(): string | null; get supportsToolCalling(): boolean | null; - download(progressCallback?: (progress: number) => void): Promise; + download(): Promise; + download(signal: AbortSignal): Promise; + download(progressCallback: (progress: number) => void, signal?: AbortSignal): Promise; + download(progressCallback: undefined, signal: AbortSignal): Promise; get path(): string; load(): Promise; removeFromCache(): void; diff --git a/sdk_v2/js/src/model.ts b/sdk_v2/js/src/model.ts index 6bb5d52ff..8053395e7 100644 --- a/sdk_v2/js/src/model.ts +++ b/sdk_v2/js/src/model.ts @@ -16,6 +16,22 @@ const internalCtorKey = Symbol("Model.internal"); const nativeByModel = new WeakMap(); +function isAbortSignal(value: unknown): value is AbortSignal { + return ( + typeof value === "object" && + value !== null && + typeof (value as AbortSignal).aborted === "boolean" && + typeof (value as AbortSignal).addEventListener === "function" && + typeof (value as AbortSignal).removeEventListener === "function" + ); +} + +function makeAbortError(message: string): Error { + const error = new Error(message); + error.name = "AbortError"; + return error; +} + function toDeviceType(value: NativeModelInfo["deviceType"]): DeviceType { switch (value) { case "CPU": @@ -24,7 +40,6 @@ function toDeviceType(value: NativeModelInfo["deviceType"]): DeviceType { return DeviceType.GPU; case "NPU": return DeviceType.NPU; - case "Invalid": default: return DeviceType.Invalid; } @@ -157,8 +172,36 @@ export class Model implements IModel { await this.#native.unload(); } - async download(progressCallback?: (progress: number) => void): Promise { - await this.#native.download(progressCallback); + async download(): Promise; + async download(signal: AbortSignal): Promise; + async download(progressCallback: (progress: number) => void, signal?: AbortSignal): Promise; + async download(progressCallback: undefined, signal: AbortSignal): Promise; + async download( + progressCallbackOrSignal?: ((progress: number) => void) | AbortSignal, + signal?: AbortSignal, + ): Promise { + const progressCallback = + typeof progressCallbackOrSignal === "function" ? progressCallbackOrSignal : undefined; + const abortSignal = isAbortSignal(progressCallbackOrSignal) ? progressCallbackOrSignal : signal; + + if ( + progressCallbackOrSignal !== undefined && + typeof progressCallbackOrSignal !== "function" && + !isAbortSignal(progressCallbackOrSignal) + ) { + throw new TypeError("Model.download: first argument must be a progress callback or AbortSignal"); + } + if (signal !== undefined && !isAbortSignal(signal)) { + throw new TypeError("Model.download: second argument must be an AbortSignal"); + } + if (isAbortSignal(progressCallbackOrSignal) && signal !== undefined) { + throw new TypeError("Model.download: signal must not be provided twice"); + } + if (abortSignal?.aborted === true) { + throw makeAbortError("Model download aborted before start"); + } + + await this.#native.download(progressCallback, abortSignal); } removeFromCache(): void { diff --git a/sdk_v2/js/test/model-download.types.ts b/sdk_v2/js/test/model-download.types.ts new file mode 100644 index 000000000..20f01068e --- /dev/null +++ b/sdk_v2/js/test/model-download.types.ts @@ -0,0 +1,11 @@ +import type { IModel } from "../src/imodel.js"; + +declare const model: IModel; +declare const progress: (percent: number) => void; +declare const signal: AbortSignal; + +void model.download(); +void model.download(signal); +void model.download(progress); +void model.download(progress, signal); +void model.download(undefined, signal); diff --git a/sdk_v2/js/test/model-lifecycle.test.ts b/sdk_v2/js/test/model-lifecycle.test.ts index 0596cef3e..69bbfb1d1 100644 --- a/sdk_v2/js/test/model-lifecycle.test.ts +++ b/sdk_v2/js/test/model-lifecycle.test.ts @@ -55,6 +55,23 @@ describe.skipIf(!haveTestModelCache)("Model lifecycle (real model)", () => { 2 * 60_000, ); + it("download() accepts an AbortSignal and preserves completion when a cache hit wins the race", async () => { + const m = fixture?.model; + if (m === undefined) throw new Error("fixture missing"); + const controller = new AbortController(); + + await expect(m.download(() => controller.abort(), controller.signal)).resolves.toBeUndefined(); + }); + + it("download() rejects a pre-aborted AbortSignal before native submission", async () => { + const m = fixture?.model; + if (m === undefined) throw new Error("fixture missing"); + const controller = new AbortController(); + controller.abort(); + + await expect(m.download(controller.signal)).rejects.toMatchObject({ name: "AbortError" }); + }); + it("calling load() on an already-loaded model is idempotent (or surfaces a clear error)", async () => { const m = fixture?.model; if (m === undefined) throw new Error("fixture missing"); diff --git a/sdk_v2/js/test/model.test.ts b/sdk_v2/js/test/model.test.ts index 07bc32b2d..a930b2019 100644 --- a/sdk_v2/js/test/model.test.ts +++ b/sdk_v2/js/test/model.test.ts @@ -4,6 +4,7 @@ import { afterAll, beforeAll, describe, expect, it } from "vitest"; import type { Catalog } from "../src/catalog.js"; +import { FlErrorCode } from "../src/detail/errors.js"; import { Model } from "../src/model.js"; import { @@ -88,6 +89,28 @@ describeIfBuilt("Model (cache-only)", () => { expect(typeof model.path).toBe("string"); }); + it("download() maps AbortSignal cancellation from a native progress checkpoint to AbortError", async () => { + const controller = new AbortController(); + + await expect( + model.download((progress) => { + if (progress === 0) controller.abort(); + }, controller.signal), + ).rejects.toMatchObject({ + name: "AbortError", + code: FlErrorCode.OperationCancelled, + }); + expect(model.isCached).toBe(false); + }); + + it("download() rejects a pre-aborted AbortSignal before native submission", async () => { + const controller = new AbortController(); + controller.abort(); + + await expect(model.download(controller.signal)).rejects.toMatchObject({ name: "AbortError" }); + expect(model.isCached).toBe(false); + }); + it("id and alias match info", () => { expect(model.id).toBe(model.info.id); expect(model.alias).toBe(model.info.alias); diff --git a/sdk_v2/js/tsconfig.types.json b/sdk_v2/js/tsconfig.types.json index d78e2d054..e5c1de9f5 100644 --- a/sdk_v2/js/tsconfig.types.json +++ b/sdk_v2/js/tsconfig.types.json @@ -5,5 +5,5 @@ "moduleResolution": "NodeNext", "noEmit": true }, - "include": ["src/**/*", "test/tool-definition.types.ts"] + "include": ["src/**/*", "test/model-download.types.ts", "test/tool-definition.types.ts"] } diff --git a/sdk_v2/python/README.md b/sdk_v2/python/README.md index 69990b531..e4bcdc6cf 100644 --- a/sdk_v2/python/README.md +++ b/sdk_v2/python/README.md @@ -110,6 +110,19 @@ with ChatSession(model) as session: model.unload() ``` +Pass a `threading.Event` as `cancel_event` to cancel an active download at the next native progress checkpoint: + +```python +from threading import Event + +cancel_event = Event() + +def on_progress(percent: float) -> None: + print(f"\rDownloading: {percent:.1f}%", end="", flush=True) + +model.download(progress_callback=on_progress, cancel_event=cancel_event) +``` + Runnable end-to-end examples live under [`samples/python/`](https://github.com/microsoft/Foundry-Local/tree/main/samples/python). ## Usage diff --git a/sdk_v2/python/src/foundry_local_sdk/imodel.py b/sdk_v2/python/src/foundry_local_sdk/imodel.py index 5b9043d2a..2052df475 100644 --- a/sdk_v2/python/src/foundry_local_sdk/imodel.py +++ b/sdk_v2/python/src/foundry_local_sdk/imodel.py @@ -5,6 +5,7 @@ from __future__ import annotations from abc import ABC, abstractmethod +from threading import Event from typing import TYPE_CHECKING, Callable from typing_extensions import deprecated @@ -78,12 +79,17 @@ def supports_tool_calling(self) -> bool | None: """Whether the model supports tool/function calling, or ``None`` if unknown.""" @abstractmethod - def download(self, progress_callback: Callable[[float], None] | None = None) -> None: + def download( + self, + progress_callback: Callable[[float], None] | None = None, + cancel_event: Event | None = None, + ) -> None: """Download the model to the local cache if not already present. Args: progress_callback: Optional callback receiving download progress as a percentage (0.0–100.0). + cancel_event: Optional event that cancels the download when set. """ @abstractmethod @@ -326,29 +332,48 @@ def supports_tool_calling(self) -> bool | None: # Model lifecycle # ------------------------------------------------------------------ - def download(self, progress_callback: Callable[[float], None] | None = None) -> None: + def download( + self, + progress_callback: Callable[[float], None] | None = None, + cancel_event: Event | None = None, + ) -> None: from foundry_local_sdk._native.api import api, ffi cb = ffi.NULL user_data = ffi.NULL + callback_error: BaseException | None = None - if progress_callback is not None: - self._progress_cb_handle = ffi.new_handle(progress_callback) + if progress_callback is not None or cancel_event is not None: + callback_state = (progress_callback, cancel_event) + progress_cb_handle = ffi.new_handle(callback_state) - @ffi.callback("flProgressCallback") - def _cb(value: float, ud: object) -> int: + def _progress_callback(value: float, ud: object) -> int: + nonlocal callback_error try: - fn = ffi.from_handle(ud) - fn(float(value)) + fn, event = ffi.from_handle(ud) + if event is not None and event.is_set(): + return 1 + if fn is not None: + fn(float(value)) + if event is not None and event.is_set(): + return 1 return 0 - except Exception: + except BaseException as exc: + callback_error = exc return 1 - self._progress_cb = _cb # keep alive + _cb = ffi.callback("flProgressCallback")(_progress_callback) cb = _cb - user_data = self._progress_cb_handle + user_data = progress_cb_handle - api.check_status(api.model.Download(self._ptr, cb, user_data)) + try: + api.check_status(api.model.Download(self._ptr, cb, user_data)) + except FoundryLocalException: + if callback_error is not None: + raise callback_error + raise + if callback_error is not None: + raise callback_error def get_path(self) -> str: from foundry_local_sdk._native.api import api, ffi diff --git a/sdk_v2/python/test/unit/test_model_download_cancellation.py b/sdk_v2/python/test/unit/test_model_download_cancellation.py new file mode 100644 index 000000000..5d3053788 --- /dev/null +++ b/sdk_v2/python/test/unit/test_model_download_cancellation.py @@ -0,0 +1,99 @@ +from __future__ import annotations + +import sys +from threading import Event +from types import SimpleNamespace + +import pytest + +from foundry_local_sdk.exception import FoundryLocalException +from foundry_local_sdk.imodel import _ModelImpl + + +class FakeFfi: + NULL = None + + @staticmethod + def new_handle(value): + return value + + @staticmethod + def from_handle(value): + return value + + @staticmethod + def callback(_signature): + return lambda fn: fn + + +def make_model(monkeypatch, invoke_callback): + def download(_ptr, callback, user_data): + return invoke_callback(callback, user_data) + + def check_status(status): + if status is not None: + raise FoundryLocalException("download cancelled", error_code=5) + + fake_api = SimpleNamespace(model=SimpleNamespace(Download=download), check_status=check_status) + monkeypatch.setitem( + sys.modules, + "foundry_local_sdk._native.api", + SimpleNamespace(api=fake_api, ffi=FakeFfi()), + ) + model = _ModelImpl.__new__(_ModelImpl) + model._ptr = object() + return model + + +def test_download_cancel_event_returns_nonzero(monkeypatch): + cancel_event = Event() + cancel_event.set() + + def invoke(callback, user_data): + assert callback(25.0, user_data) == 1 + return object() + + model = make_model(monkeypatch, invoke) + + with pytest.raises(FoundryLocalException, match="download cancelled") as exc: + model.download(cancel_event=cancel_event) + + assert exc.value.error_code == 5 + + +def test_download_cancel_event_is_checked_after_progress(monkeypatch): + cancel_event = Event() + progress: list[float] = [] + + def on_progress(value: float) -> None: + progress.append(value) + cancel_event.set() + + def invoke(callback, user_data): + assert callback(50.0, user_data) == 1 + return object() + + model = make_model(monkeypatch, invoke) + + with pytest.raises(FoundryLocalException, match="download cancelled"): + model.download(on_progress, cancel_event) + + assert progress == [50.0] + + +def test_download_preserves_progress_callback_exception(monkeypatch): + expected = RuntimeError("progress failed") + + def on_progress(_value: float) -> None: + raise expected + + def invoke(callback, user_data): + assert callback(10.0, user_data) == 1 + return object() + + model = make_model(monkeypatch, invoke) + + with pytest.raises(RuntimeError, match="progress failed") as exc: + model.download(on_progress) + + assert exc.value is expected From 36ebfeba58e9c7ffe67f54c8aee557ce14c01101 Mon Sep 17 00:00:00 2001 From: Bhagirath Mehta Date: Fri, 18 Sep 2026 18:22:12 -0500 Subject: [PATCH 2/6] Preserve download overload source compatibility Keep optional AbortSignal and callback-or-undefined variables accepted while rejecting duplicate signals, and make the Python cancellation example demonstrate cancellation.\n\nFiles changed:\n- sdk_v2/js/src/imodel.ts\n- sdk_v2/js/src/model.ts\n- sdk_v2/js/test/model-download.types.ts\n- sdk_v2/python/README.md\n\nCo-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>\nCopilot-Session: 2d2c76d2-0c86-4672-b91b-3fcb54ed7906 --- sdk_v2/js/src/imodel.ts | 6 ++---- sdk_v2/js/src/model.ts | 9 +++------ sdk_v2/js/test/model-download.types.ts | 9 +++++++++ sdk_v2/python/README.md | 1 + 4 files changed, 15 insertions(+), 10 deletions(-) diff --git a/sdk_v2/js/src/imodel.ts b/sdk_v2/js/src/imodel.ts index e653df933..4af8fc076 100644 --- a/sdk_v2/js/src/imodel.ts +++ b/sdk_v2/js/src/imodel.ts @@ -19,10 +19,8 @@ export interface IModel { get capabilities(): string | null; get supportsToolCalling(): boolean | null; - download(): Promise; - download(signal: AbortSignal): Promise; - download(progressCallback: (progress: number) => void, signal?: AbortSignal): Promise; - download(progressCallback: undefined, signal: AbortSignal): Promise; + download(signal?: AbortSignal): Promise; + download(progressCallback: ((progress: number) => void) | undefined, signal?: AbortSignal): Promise; get path(): string; load(): Promise; removeFromCache(): void; diff --git a/sdk_v2/js/src/model.ts b/sdk_v2/js/src/model.ts index 8053395e7..f7e0b06e7 100644 --- a/sdk_v2/js/src/model.ts +++ b/sdk_v2/js/src/model.ts @@ -172,16 +172,13 @@ export class Model implements IModel { await this.#native.unload(); } - async download(): Promise; - async download(signal: AbortSignal): Promise; - async download(progressCallback: (progress: number) => void, signal?: AbortSignal): Promise; - async download(progressCallback: undefined, signal: AbortSignal): Promise; + async download(signal?: AbortSignal): Promise; + async download(progressCallback: ((progress: number) => void) | undefined, signal?: AbortSignal): Promise; async download( progressCallbackOrSignal?: ((progress: number) => void) | AbortSignal, signal?: AbortSignal, ): Promise { - const progressCallback = - typeof progressCallbackOrSignal === "function" ? progressCallbackOrSignal : undefined; + const progressCallback = typeof progressCallbackOrSignal === "function" ? progressCallbackOrSignal : undefined; const abortSignal = isAbortSignal(progressCallbackOrSignal) ? progressCallbackOrSignal : signal; if ( diff --git a/sdk_v2/js/test/model-download.types.ts b/sdk_v2/js/test/model-download.types.ts index 20f01068e..b51fb595d 100644 --- a/sdk_v2/js/test/model-download.types.ts +++ b/sdk_v2/js/test/model-download.types.ts @@ -2,10 +2,19 @@ import type { IModel } from "../src/imodel.js"; declare const model: IModel; declare const progress: (percent: number) => void; +declare const maybeProgress: ((percent: number) => void) | undefined; declare const signal: AbortSignal; +declare const maybeSignal: AbortSignal | undefined; +declare const secondSignal: AbortSignal; void model.download(); +void model.download(undefined); void model.download(signal); +void model.download(maybeSignal); void model.download(progress); +void model.download(maybeProgress); void model.download(progress, signal); void model.download(undefined, signal); + +// @ts-expect-error A second signal is only valid when the first argument is a progress callback. +void model.download(signal, secondSignal); diff --git a/sdk_v2/python/README.md b/sdk_v2/python/README.md index e4bcdc6cf..3f184ec80 100644 --- a/sdk_v2/python/README.md +++ b/sdk_v2/python/README.md @@ -119,6 +119,7 @@ cancel_event = Event() def on_progress(percent: float) -> None: print(f"\rDownloading: {percent:.1f}%", end="", flush=True) + cancel_event.set() model.download(progress_callback=on_progress, cancel_event=cancel_event) ``` From f9bd4a66f440846535281596f6d9331a43a2514f Mon Sep 17 00:00:00 2001 From: Bhagirath Mehta Date: Fri, 18 Sep 2026 20:10:21 -0500 Subject: [PATCH 3/6] Preserve the callback-first download API Keep AbortSignal as an optional second argument so existing IModel implementations remain source-compatible while model downloads still support cancellation. Files changed: - sdk_v2/js/src/imodel.ts - sdk_v2/js/src/model.ts - sdk_v2/js/test/model-download.types.ts - sdk_v2/js/test/model-lifecycle.test.ts - sdk_v2/js/test/model.test.ts Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 2d2c76d2-0c86-4672-b91b-3fcb54ed7906 --- sdk_v2/js/src/imodel.ts | 3 +-- sdk_v2/js/src/model.ts | 21 +++------------------ sdk_v2/js/test/model-download.types.ts | 13 ++++++++----- sdk_v2/js/test/model-lifecycle.test.ts | 2 +- sdk_v2/js/test/model.test.ts | 2 +- 5 files changed, 14 insertions(+), 27 deletions(-) diff --git a/sdk_v2/js/src/imodel.ts b/sdk_v2/js/src/imodel.ts index 4af8fc076..b1481838b 100644 --- a/sdk_v2/js/src/imodel.ts +++ b/sdk_v2/js/src/imodel.ts @@ -19,8 +19,7 @@ export interface IModel { get capabilities(): string | null; get supportsToolCalling(): boolean | null; - download(signal?: AbortSignal): Promise; - download(progressCallback: ((progress: number) => void) | undefined, signal?: AbortSignal): Promise; + download(progressCallback?: (progress: number) => void, signal?: AbortSignal): Promise; get path(): string; load(): Promise; removeFromCache(): void; diff --git a/sdk_v2/js/src/model.ts b/sdk_v2/js/src/model.ts index f7e0b06e7..50fce8dde 100644 --- a/sdk_v2/js/src/model.ts +++ b/sdk_v2/js/src/model.ts @@ -172,33 +172,18 @@ export class Model implements IModel { await this.#native.unload(); } - async download(signal?: AbortSignal): Promise; - async download(progressCallback: ((progress: number) => void) | undefined, signal?: AbortSignal): Promise; async download( - progressCallbackOrSignal?: ((progress: number) => void) | AbortSignal, + progressCallback?: (progress: number) => void, signal?: AbortSignal, ): Promise { - const progressCallback = typeof progressCallbackOrSignal === "function" ? progressCallbackOrSignal : undefined; - const abortSignal = isAbortSignal(progressCallbackOrSignal) ? progressCallbackOrSignal : signal; - - if ( - progressCallbackOrSignal !== undefined && - typeof progressCallbackOrSignal !== "function" && - !isAbortSignal(progressCallbackOrSignal) - ) { - throw new TypeError("Model.download: first argument must be a progress callback or AbortSignal"); - } if (signal !== undefined && !isAbortSignal(signal)) { throw new TypeError("Model.download: second argument must be an AbortSignal"); } - if (isAbortSignal(progressCallbackOrSignal) && signal !== undefined) { - throw new TypeError("Model.download: signal must not be provided twice"); - } - if (abortSignal?.aborted === true) { + if (signal?.aborted === true) { throw makeAbortError("Model download aborted before start"); } - await this.#native.download(progressCallback, abortSignal); + await this.#native.download(progressCallback, signal); } removeFromCache(): void { diff --git a/sdk_v2/js/test/model-download.types.ts b/sdk_v2/js/test/model-download.types.ts index b51fb595d..1e464349b 100644 --- a/sdk_v2/js/test/model-download.types.ts +++ b/sdk_v2/js/test/model-download.types.ts @@ -5,16 +5,19 @@ declare const progress: (percent: number) => void; declare const maybeProgress: ((percent: number) => void) | undefined; declare const signal: AbortSignal; declare const maybeSignal: AbortSignal | undefined; -declare const secondSignal: AbortSignal; void model.download(); void model.download(undefined); -void model.download(signal); -void model.download(maybeSignal); void model.download(progress); void model.download(maybeProgress); void model.download(progress, signal); void model.download(undefined, signal); +void model.download(progress, maybeSignal); + +// Existing structural implementations remain compatible after adding the optional signal parameter. +declare const legacyDownload: (progressCallback?: (percent: number) => void) => Promise; +const compatibleDownload: IModel["download"] = legacyDownload; +void compatibleDownload; -// @ts-expect-error A second signal is only valid when the first argument is a progress callback. -void model.download(signal, secondSignal); +// @ts-expect-error AbortSignal remains the optional second argument so the original callback-first API is preserved. +void model.download(signal); diff --git a/sdk_v2/js/test/model-lifecycle.test.ts b/sdk_v2/js/test/model-lifecycle.test.ts index 69bbfb1d1..e1d76c8fd 100644 --- a/sdk_v2/js/test/model-lifecycle.test.ts +++ b/sdk_v2/js/test/model-lifecycle.test.ts @@ -69,7 +69,7 @@ describe.skipIf(!haveTestModelCache)("Model lifecycle (real model)", () => { const controller = new AbortController(); controller.abort(); - await expect(m.download(controller.signal)).rejects.toMatchObject({ name: "AbortError" }); + await expect(m.download(undefined, controller.signal)).rejects.toMatchObject({ name: "AbortError" }); }); it("calling load() on an already-loaded model is idempotent (or surfaces a clear error)", async () => { diff --git a/sdk_v2/js/test/model.test.ts b/sdk_v2/js/test/model.test.ts index a930b2019..7fb984db6 100644 --- a/sdk_v2/js/test/model.test.ts +++ b/sdk_v2/js/test/model.test.ts @@ -107,7 +107,7 @@ describeIfBuilt("Model (cache-only)", () => { const controller = new AbortController(); controller.abort(); - await expect(model.download(controller.signal)).rejects.toMatchObject({ name: "AbortError" }); + await expect(model.download(undefined, controller.signal)).rejects.toMatchObject({ name: "AbortError" }); expect(model.isCached).toBe(false); }); From fcff01433c89c4aa63381c1f12c4a9b416a5a941 Mon Sep 17 00:00:00 2001 From: Bhagirath Mehta Date: Fri, 18 Sep 2026 20:36:00 -0500 Subject: [PATCH 4/6] Remove obsolete download callback fields Keep CFFI callback state local to the synchronous download call instead of retaining unused per-model references. Files changed: - sdk_v2/python/src/foundry_local_sdk/imodel.py Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 2d2c76d2-0c86-4672-b91b-3fcb54ed7906 --- sdk_v2/python/src/foundry_local_sdk/imodel.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/sdk_v2/python/src/foundry_local_sdk/imodel.py b/sdk_v2/python/src/foundry_local_sdk/imodel.py index 2052df475..0d88e0851 100644 --- a/sdk_v2/python/src/foundry_local_sdk/imodel.py +++ b/sdk_v2/python/src/foundry_local_sdk/imodel.py @@ -257,9 +257,6 @@ def __init__(self, native_ptr: object, *, parent: object | None = None) -> None: # is owned by the catalog; without this reference, GC could release the # catalog (and the manager behind it) first and dangle our pointer. self._parent = parent - # Callback references — stored to prevent premature GC. - self._progress_cb = None - self._progress_cb_handle = None @property def _native_ptr(self) -> object: From ab109e03e5c4e7b1fae37b03481633fdc4787c3c Mon Sep 17 00:00:00 2001 From: Bhagirath Mehta Date: Fri, 18 Sep 2026 20:57:10 -0500 Subject: [PATCH 5/6] Serialize Python download progress callbacks Prevent overlapping native callbacks from repeating user side effects or replacing the first callback exception during cooperative cancellation. Files changed: - sdk_v2/python/src/foundry_local_sdk/imodel.py - sdk_v2/python/test/unit/test_model_download_cancellation.py Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 2d2c76d2-0c86-4672-b91b-3fcb54ed7906 --- sdk_v2/python/src/foundry_local_sdk/imodel.py | 26 +++++----- .../unit/test_model_download_cancellation.py | 48 ++++++++++++++++++- 2 files changed, 62 insertions(+), 12 deletions(-) diff --git a/sdk_v2/python/src/foundry_local_sdk/imodel.py b/sdk_v2/python/src/foundry_local_sdk/imodel.py index 0d88e0851..5caf2dbfc 100644 --- a/sdk_v2/python/src/foundry_local_sdk/imodel.py +++ b/sdk_v2/python/src/foundry_local_sdk/imodel.py @@ -5,7 +5,7 @@ from __future__ import annotations from abc import ABC, abstractmethod -from threading import Event +from threading import Event, Lock from typing import TYPE_CHECKING, Callable from typing_extensions import deprecated @@ -343,21 +343,25 @@ def download( if progress_callback is not None or cancel_event is not None: callback_state = (progress_callback, cancel_event) progress_cb_handle = ffi.new_handle(callback_state) + callback_lock = Lock() def _progress_callback(value: float, ud: object) -> int: nonlocal callback_error - try: - fn, event = ffi.from_handle(ud) - if event is not None and event.is_set(): + with callback_lock: + if callback_error is not None: return 1 - if fn is not None: - fn(float(value)) - if event is not None and event.is_set(): + try: + fn, event = ffi.from_handle(ud) + if event is not None and event.is_set(): + return 1 + if fn is not None: + fn(float(value)) + if event is not None and event.is_set(): + return 1 + return 0 + except BaseException as exc: + callback_error = exc return 1 - return 0 - except BaseException as exc: - callback_error = exc - return 1 _cb = ffi.callback("flProgressCallback")(_progress_callback) cb = _cb diff --git a/sdk_v2/python/test/unit/test_model_download_cancellation.py b/sdk_v2/python/test/unit/test_model_download_cancellation.py index 5d3053788..02b3a6690 100644 --- a/sdk_v2/python/test/unit/test_model_download_cancellation.py +++ b/sdk_v2/python/test/unit/test_model_download_cancellation.py @@ -1,7 +1,7 @@ from __future__ import annotations import sys -from threading import Event +from threading import Event, Thread from types import SimpleNamespace import pytest @@ -97,3 +97,49 @@ def invoke(callback, user_data): model.download(on_progress) assert exc.value is expected + + +def test_download_preserves_first_exception_across_concurrent_callbacks(monkeypatch): + expected = RuntimeError("first progress failed") + first_entered = Event() + second_native_started = Event() + second_user_entered = Event() + progress: list[float] = [] + + def on_progress(value: float) -> None: + progress.append(value) + if value == 10.0: + first_entered.set() + second_native_started.wait(timeout=1) + second_user_entered.wait(timeout=0.1) + raise expected + second_user_entered.set() + + def invoke(callback, user_data): + results: dict[str, int] = {} + + first = Thread(target=lambda: results.setdefault("first", callback(10.0, user_data))) + first.start() + assert first_entered.wait(timeout=1) + + def call_second() -> None: + second_native_started.set() + results["second"] = callback(20.0, user_data) + + second = Thread(target=call_second) + second.start() + first.join(timeout=2) + second.join(timeout=2) + + assert not first.is_alive() + assert not second.is_alive() + assert results == {"first": 1, "second": 1} + return object() + + model = make_model(monkeypatch, invoke) + + with pytest.raises(RuntimeError, match="first progress failed") as exc: + model.download(on_progress) + + assert exc.value is expected + assert progress == [10.0] From c6ee368c939670da657f84520123c361e707e6ec Mon Sep 17 00:00:00 2001 From: Bhagirath Mehta Date: Fri, 18 Sep 2026 21:08:53 -0500 Subject: [PATCH 6/6] Wait for JavaScript download progress callbacks Acknowledge each TSFN dispatch before rechecking AbortSignal state so cancellation triggered by a progress callback stops at the same native checkpoint. Files changed: - sdk_v2/js/native/src/model.cc Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 2d2c76d2-0c86-4672-b91b-3fcb54ed7906 --- sdk_v2/js/native/src/model.cc | 55 +++++++++++++++++++++++++++++++---- 1 file changed, 50 insertions(+), 5 deletions(-) diff --git a/sdk_v2/js/native/src/model.cc b/sdk_v2/js/native/src/model.cc index d948d3897..38ecc9b1f 100644 --- a/sdk_v2/js/native/src/model.cc +++ b/sdk_v2/js/native/src/model.cc @@ -10,7 +10,9 @@ #include #include +#include #include +#include #include #include #include @@ -277,6 +279,41 @@ Napi::Value Model::Unload(const Napi::CallbackInfo& info) { namespace { +class ProgressDispatch { + public: + explicit ProgressDispatch(float percent) : percent_(percent) {} + + float Percent() const { return percent_; } + + void Wait() { + std::unique_lock lock(mutex_); + completed_cv_.wait(lock, [this] { return completed_; }); + } + + void Complete() { + { + std::lock_guard lock(mutex_); + completed_ = true; + } + completed_cv_.notify_one(); + } + + private: + const float percent_; + std::mutex mutex_; + std::condition_variable completed_cv_; + bool completed_ = false; +}; + +class ProgressDispatchCompletion { + public: + explicit ProgressDispatchCompletion(ProgressDispatch& dispatch) : dispatch_(dispatch) {} + ~ProgressDispatchCompletion() { dispatch_.Complete(); } + + private: + ProgressDispatch& dispatch_; +}; + // AsyncWorker variant that drives IModel::Download with an optional JS // progress callback. The callback runs on the libuv worker thread; we bounce // each (float percent) to JS via a ThreadSafeFunction acquired before the @@ -306,11 +343,19 @@ class DownloadWorker : public Napi::AsyncWorker { return 1; } if (tsfn_) { - // BlockingCall keeps backpressure on the worker thread: if JS is slow to drain the queue we'll wait rather - // than dropping reports. - tsfn_.BlockingCall([percent](Napi::Env env, Napi::Function js_cb) { - js_cb.Call({Napi::Number::New(env, static_cast(percent))}); - }); + ProgressDispatch dispatch(percent); + const napi_status status = tsfn_.BlockingCall( + &dispatch, [](Napi::Env env, Napi::Function js_cb, ProgressDispatch* pending) { + ProgressDispatchCompletion completion(*pending); + if (env == nullptr || js_cb.IsEmpty()) { + return; + } + js_cb.Call({Napi::Number::New(env, static_cast(pending->Percent()))}); + }); + if (status != napi_ok) { + return 1; + } + dispatch.Wait(); } if (IsAbortRequested()) { cancelled_by_signal_ = true;