diff --git a/sdk_v2/js/native/src/model.cc b/sdk_v2/js/native/src/model.cc index 6d1844c9e..38ecc9b1f 100644 --- a/sdk_v2/js/native/src/model.cc +++ b/sdk_v2/js/native/src/model.cc @@ -9,7 +9,10 @@ #include #include +#include +#include #include +#include #include #include #include @@ -276,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 @@ -283,27 +321,49 @@ 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_) { + 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; + 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 +381,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 +403,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 +428,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 +458,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..b1481838b 100644 --- a/sdk_v2/js/src/imodel.ts +++ b/sdk_v2/js/src/imodel.ts @@ -19,7 +19,7 @@ export interface IModel { get capabilities(): string | null; get supportsToolCalling(): boolean | null; - download(progressCallback?: (progress: number) => void): 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 6bb5d52ff..50fce8dde 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,18 @@ export class Model implements IModel { await this.#native.unload(); } - async download(progressCallback?: (progress: number) => void): Promise { - await this.#native.download(progressCallback); + async download( + progressCallback?: (progress: number) => void, + signal?: AbortSignal, + ): Promise { + if (signal !== undefined && !isAbortSignal(signal)) { + throw new TypeError("Model.download: second argument must be an AbortSignal"); + } + if (signal?.aborted === true) { + throw makeAbortError("Model download aborted before start"); + } + + 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 new file mode 100644 index 000000000..1e464349b --- /dev/null +++ b/sdk_v2/js/test/model-download.types.ts @@ -0,0 +1,23 @@ +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; + +void model.download(); +void model.download(undefined); +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 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 0596cef3e..e1d76c8fd 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(undefined, 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..7fb984db6 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(undefined, 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..3f184ec80 100644 --- a/sdk_v2/python/README.md +++ b/sdk_v2/python/README.md @@ -110,6 +110,20 @@ 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) + cancel_event.set() + +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..5caf2dbfc 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, Lock 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 @@ -251,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: @@ -326,29 +329,52 @@ 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 - - if progress_callback is not None: - self._progress_cb_handle = ffi.new_handle(progress_callback) - - @ffi.callback("flProgressCallback") - def _cb(value: float, ud: object) -> int: - try: - fn = ffi.from_handle(ud) - fn(float(value)) - return 0 - except Exception: - return 1 - - self._progress_cb = _cb # keep alive + callback_error: BaseException | None = None + + 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 + with callback_lock: + if callback_error is not None: + return 1 + 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 + + _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..02b3a6690 --- /dev/null +++ b/sdk_v2/python/test/unit/test_model_download_cancellation.py @@ -0,0 +1,145 @@ +from __future__ import annotations + +import sys +from threading import Event, Thread +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 + + +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]