diff --git a/extensions/llama-cpp/src/inference-messages.ts b/extensions/llama-cpp/src/inference-messages.ts new file mode 100644 index 000000000000..7ab1d76b5a5d --- /dev/null +++ b/extensions/llama-cpp/src/inference-messages.ts @@ -0,0 +1,51 @@ +import type { StreamFn } from "openclaw/plugin-sdk/agent-core"; +import type { AssistantMessage, StopReason, Usage } from "openclaw/plugin-sdk/llm"; + +export function zeroCostUsage(input = 0, output = 0): Usage { + return { + input, + output, + cacheRead: 0, + cacheWrite: 0, + totalTokens: input + output, + cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, + }; +} + +export function buildMessage(params: { + model: Parameters[0]; + content: AssistantMessage["content"]; + stopReason: StopReason; + usage?: Usage; + errorMessage?: string; +}): AssistantMessage { + return { + role: "assistant", + content: params.content, + api: params.model.api, + provider: params.model.provider, + model: params.model.id, + stopReason: params.stopReason, + usage: params.usage ?? zeroCostUsage(), + timestamp: Date.now(), + ...(params.errorMessage ? { errorMessage: params.errorMessage } : {}), + }; +} + +export function runtimeUnavailableErrorMessage(restartRequired: boolean): string { + return restartRequired + ? "llama.cpp runtime stopped after cleanup failed. Run `openclaw gateway restart` to recover." + : "llama.cpp runtime is stopping"; +} + +export function runtimeUnavailableMessage( + model: Parameters[0], + restartRequired: boolean, +): AssistantMessage { + return buildMessage({ + model, + content: [], + stopReason: "error", + errorMessage: runtimeUnavailableErrorMessage(restartRequired), + }); +} diff --git a/extensions/llama-cpp/src/inference-provider.test.ts b/extensions/llama-cpp/src/inference-provider.test.ts index e35030c26f6b..73157965b220 100644 --- a/extensions/llama-cpp/src/inference-provider.test.ts +++ b/extensions/llama-cpp/src/inference-provider.test.ts @@ -58,6 +58,11 @@ import { createLlamaCppInferenceRuntime } from "./inference-provider.js"; type LlamaCppInferenceRuntime = ReturnType; let inferenceRuntime: LlamaCppInferenceRuntime; +const testApi = (globalThis as Record)[ + Symbol.for("openclaw.llamaCppInferenceTestApi") +] as { + resetInferenceRuntimeCoordinator: () => void; +}; const model: Model = { id: "test.gguf", @@ -145,6 +150,7 @@ function expectDisposeCalls(contextCount: number, modelCount: number, llamaCount } beforeEach(() => { + testApi.resetInferenceRuntimeCoordinator(); inferenceRuntime = createLlamaCppInferenceRuntime(); vi.clearAllMocks(); mocks.generateResponse.mockResolvedValue({ @@ -817,6 +823,86 @@ describe("llama.cpp inference provider", () => { ); }); + it("blocks a published replacement until predecessor service stop finishes", async () => { + const retiringRuntime = inferenceRuntime; + await collectTestEvents(); + let finishCleanup!: () => void; + mocks.contextDispose.mockImplementationOnce( + async () => + await new Promise((resolve) => { + finishCleanup = resolve; + }), + ); + // Gateway publishes the replacement registry before stopping old services. + const disposing = retiringRuntime.dispose(); + await vi.waitFor(() => expect(mocks.contextDispose).toHaveBeenCalledOnce()); + + const replacementRuntime = createLlamaCppInferenceRuntime(); + const replacementEvents = collectEvents( + await replacementRuntime.createStreamFn({})(model, { + messages: [{ role: "user", content: "after reload", timestamp: 2 }], + }), + ); + await Promise.resolve(); + expect(mocks.llama.loadModel).toHaveBeenCalledOnce(); + + finishCleanup(); + await disposing; + await expect(replacementEvents).resolves.toEqual( + expect.arrayContaining([expect.objectContaining({ type: "done", reason: "stop" })]), + ); + expect(mocks.llama.loadModel).toHaveBeenCalledTimes(2); + inferenceRuntime = replacementRuntime; + }); + + it("keeps a published replacement blocked when predecessor service stop fails", async () => { + const retiringRuntime = inferenceRuntime; + await collectTestEvents(); + mocks.contextDispose.mockRejectedValueOnce(new Error("retiring cleanup failed")); + + const replacementRuntime = createLlamaCppInferenceRuntime(); + // Match reload ordering: replacement is reachable before old service stop settles. + const disposing = retiringRuntime.dispose(); + const replacementEvents = collectEvents( + await replacementRuntime.createStreamFn({})(model, { + messages: [{ role: "user", content: "after failed reload", timestamp: 2 }], + }), + ); + + await expect(disposing).rejects.toThrow("retiring cleanup failed"); + await expect(replacementEvents).resolves.toEqual( + expect.arrayContaining([ + expect.objectContaining({ + type: "error", + error: expect.objectContaining({ + errorMessage: expect.stringContaining("openclaw gateway restart"), + }), + }), + ]), + ); + expect(mocks.llama.loadModel).toHaveBeenCalledOnce(); + await expect(replacementRuntime.dispose()).resolves.toBeUndefined(); + inferenceRuntime = replacementRuntime; + }); + + it("lets a superseded waiting replacement stop before its predecessor retires", async () => { + await collectTestEvents(); + const waitingRuntime = createLlamaCppInferenceRuntime(); + const waitingStream = await waitingRuntime.createStreamFn({})(model, { + messages: [{ role: "user", content: "superseded reload", timestamp: 2 }], + }); + await Promise.resolve(); + + const disposingWaitingRuntime = waitingRuntime.dispose(); + + await expect(waitingStream.result()).resolves.toMatchObject({ + stopReason: "error", + errorMessage: "llama.cpp runtime is stopping", + }); + await expect(disposingWaitingRuntime).resolves.toBeUndefined(); + expect(mocks.llama.loadModel).toHaveBeenCalledOnce(); + }); + it("waits for admitted inference before disposing the runtime", async () => { const finishGeneration = deferGeneration(); const stream = await createTestStream(); diff --git a/extensions/llama-cpp/src/inference-provider.ts b/extensions/llama-cpp/src/inference-provider.ts index f077aca88deb..22096b9e5559 100644 --- a/extensions/llama-cpp/src/inference-provider.ts +++ b/extensions/llama-cpp/src/inference-provider.ts @@ -10,13 +10,7 @@ import type { LlamaModel, } from "node-llama-cpp"; import type { StreamFn } from "openclaw/plugin-sdk/agent-core"; -import type { - AssistantMessage, - Context, - StopReason, - ToolCall, - Usage, -} from "openclaw/plugin-sdk/llm"; +import type { AssistantMessage, Context, StopReason, ToolCall } from "openclaw/plugin-sdk/llm"; import { createAssistantMessageEventStream, parseStreamingJson } from "openclaw/plugin-sdk/llm"; import type { ModelProviderConfig } from "openclaw/plugin-sdk/provider-model-shared"; import { createPlainTextToolCallCompatWrapper } from "openclaw/plugin-sdk/provider-stream-shared"; @@ -25,6 +19,16 @@ import { resolveLlamaCppModelCacheDir, resolveLlamaCppModelSource, } from "./defaults.js"; +import { + buildMessage, + runtimeUnavailableErrorMessage, + runtimeUnavailableMessage, + zeroCostUsage, +} from "./inference-messages.js"; +import { + createLlamaCppInferenceRuntimeToken, + type LlamaCppInferenceRuntimeToken, +} from "./inference-runtime-coordinator.js"; import { formatLlamaCppSetupError, importNodeLlamaCpp, @@ -42,11 +46,13 @@ type LoadedModel = { type LlamaJsonSchemaInput = Parameters[0]; type LlamaCppInferenceRuntimeState = { + admission: LlamaCppInferenceRuntimeToken; loadedModel?: LoadedModel; llamaInstance?: Llama; operationQueue: Promise; lifecycle: "open" | "closing" | "closed"; cleanupFailure?: { error: Error }; + retiringRuntimeFailure?: boolean; disposePromise?: Promise; }; @@ -55,53 +61,8 @@ type LlamaCppInferenceRuntime = { dispose: () => Promise; }; -function zeroCostUsage(input = 0, output = 0): Usage { - return { - input, - output, - cacheRead: 0, - cacheWrite: 0, - totalTokens: input + output, - cost: { input: 0, output: 0, cacheRead: 0, cacheWrite: 0, total: 0 }, - }; -} - -function buildMessage(params: { - model: Parameters[0]; - content: AssistantMessage["content"]; - stopReason: StopReason; - usage?: Usage; - errorMessage?: string; -}): AssistantMessage { - return { - role: "assistant", - content: params.content, - api: params.model.api, - provider: params.model.provider, - model: params.model.id, - stopReason: params.stopReason, - usage: params.usage ?? zeroCostUsage(), - timestamp: Date.now(), - ...(params.errorMessage ? { errorMessage: params.errorMessage } : {}), - }; -} - -function runtimeUnavailableErrorMessage(state: LlamaCppInferenceRuntimeState): string { - return state.cleanupFailure - ? "llama.cpp runtime stopped after cleanup failed. Run `openclaw gateway restart` to recover." - : "llama.cpp runtime is stopping"; -} - -function runtimeUnavailableMessage( - state: LlamaCppInferenceRuntimeState, - model: Parameters[0], -): AssistantMessage { - return buildMessage({ - model, - content: [], - stopReason: "error", - errorMessage: runtimeUnavailableErrorMessage(state), - }); +function runtimeRequiresRestart(state: LlamaCppInferenceRuntimeState): boolean { + return Boolean(state.cleanupFailure || state.retiringRuntimeFailure); } function extractText(content: unknown): string { @@ -279,6 +240,12 @@ async function disposeLoadedModel(state: LlamaCppInferenceRuntimeState): Promise function recordCleanupFailure(state: LlamaCppInferenceRuntimeState, error: unknown): void { state.cleanupFailure ??= { error: error instanceof Error ? error : new Error(String(error)) }; state.lifecycle = "closed"; + state.admission.fail(state.cleanupFailure.error); +} + +function recordRetiringRuntimeFailure(state: LlamaCppInferenceRuntimeState): void { + state.retiringRuntimeFailure = true; + state.lifecycle = "closed"; } async function getLoadedModel(params: { @@ -345,6 +312,7 @@ function disposeLlamaCppInferenceRuntime(state: LlamaCppInferenceRuntimeState): return state.disposePromise; } state.lifecycle = "closing"; + state.admission.close(); // node-llama-cpp disposers are one-shot and child cleanup releases the // parent's disposal guard. Do not force parent cleanup after a child rejects: // the retained guard can make that parent disposer wait forever. @@ -360,6 +328,7 @@ function disposeLlamaCppInferenceRuntime(state: LlamaCppInferenceRuntimeState): state.llamaInstance = undefined; } } + state.admission.release(); }) .catch((error: unknown) => { recordCleanupFailure(state, error); @@ -381,7 +350,7 @@ function createLlamaCppStreamFnForRuntime( stream.push({ type: "error", reason: "error", - error: runtimeUnavailableMessage(state, model), + error: runtimeUnavailableMessage(model, runtimeRequiresRestart(state)), }); stream.end(); return stream; @@ -424,10 +393,18 @@ function createLlamaCppStreamFnForRuntime( stream.push({ type: "error", reason: "error", - error: runtimeUnavailableMessage(state, model), + error: runtimeUnavailableMessage(model, runtimeRequiresRestart(state)), }); return; } + await state.admission.acquire({ + signal, + isLive: () => state.lifecycle === "open", + onRestartRequired: () => recordRetiringRuntimeFailure(state), + }); + if (state.lifecycle !== "open") { + throw new Error("llama.cpp runtime is stopping"); + } const runtime = await importNodeLlamaCpp(); const loaded = await getLoadedModel({ state, @@ -688,8 +665,8 @@ function createLlamaCppStreamFnForRuntime( const reason = aborted ? "aborted" : "error"; const errorMessage = aborted ? "Request was aborted" - : state.cleanupFailure - ? runtimeUnavailableErrorMessage(state) + : state.lifecycle !== "open" + ? runtimeUnavailableErrorMessage(runtimeRequiresRestart(state)) : formatLlamaCppSetupError(error); stream.push({ type: "error", @@ -715,6 +692,7 @@ function createLlamaCppStreamFnForRuntime( export function createLlamaCppInferenceRuntime(): LlamaCppInferenceRuntime { const state: LlamaCppInferenceRuntimeState = { + admission: createLlamaCppInferenceRuntimeToken(), operationQueue: Promise.resolve(), lifecycle: "open", }; @@ -725,8 +703,10 @@ export function createLlamaCppInferenceRuntime(): LlamaCppInferenceRuntime { } if (process.env.VITEST || process.env.NODE_ENV === "test") { - (globalThis as Record)[Symbol.for("openclaw.llamaCppInferenceTestApi")] = { + const globalStore = globalThis as Record; + const testApiKey = Symbol.for("openclaw.llamaCppInferenceTestApi"); + Object.assign((globalStore[testApiKey] ??= {}), { mapContextToLlamaChatHistory, mapToolsToLlamaFunctions, - }; + }); } diff --git a/extensions/llama-cpp/src/inference-runtime-coordinator.test.ts b/extensions/llama-cpp/src/inference-runtime-coordinator.test.ts new file mode 100644 index 000000000000..d215a4e421f0 --- /dev/null +++ b/extensions/llama-cpp/src/inference-runtime-coordinator.test.ts @@ -0,0 +1,113 @@ +import { beforeEach, describe, expect, it } from "vitest"; +import { + createLlamaCppInferenceRuntimeToken, + LlamaCppInferenceRestartRequiredError, +} from "./inference-runtime-coordinator.js"; + +const testApi = (globalThis as Record)[ + Symbol.for("openclaw.llamaCppInferenceTestApi") +] as { + resetInferenceRuntimeCoordinator: () => void; +}; + +function createLiveToken() { + let live = true; + const token = createLlamaCppInferenceRuntimeToken(); + return { + acquire: (signal?: AbortSignal) => + token.acquire({ signal, isLive: () => live, onRestartRequired: () => undefined }), + dispose: () => { + live = false; + token.close(); + }, + fail: token.fail, + release: token.release, + }; +} + +beforeEach(() => { + testApi.resetInferenceRuntimeCoordinator(); +}); + +describe("llama.cpp inference runtime coordinator", () => { + it("does not reserve native ownership until a runtime acquires it", async () => { + createLiveToken(); + const replacement = createLiveToken(); + + await expect(replacement.acquire()).resolves.toBeUndefined(); + }); + + it("hands native ownership to rapid contenders one generation at a time", async () => { + const first = createLiveToken(); + const second = createLiveToken(); + const third = createLiveToken(); + await first.acquire(); + let secondAcquired = false; + let thirdAcquired = false; + const secondAcquisition = second.acquire().then(() => { + secondAcquired = true; + }); + const thirdAcquisition = third.acquire().then(() => { + thirdAcquired = true; + }); + + first.release(); + await secondAcquisition; + expect(secondAcquired).toBe(true); + expect(thirdAcquired).toBe(false); + + second.release(); + await thirdAcquisition; + expect(thirdAcquired).toBe(true); + }); + + it("skips a disposed waiter before handing ownership to the next runtime", async () => { + const first = createLiveToken(); + const disposed = createLiveToken(); + const replacement = createLiveToken(); + await first.acquire(); + const disposedAcquisition = disposed.acquire(); + let replacementAcquired = false; + const replacementAcquisition = replacement.acquire().then(() => { + replacementAcquired = true; + }); + + disposed.dispose(); + + await expect(disposedAcquisition).rejects.toThrow("runtime is stopping"); + expect(replacementAcquired).toBe(false); + + first.release(); + await expect(replacementAcquisition).resolves.toBeUndefined(); + }); + + it("removes an aborted waiter without letting it claim ownership", async () => { + const first = createLiveToken(); + const aborted = createLiveToken(); + const replacement = createLiveToken(); + const abortController = new AbortController(); + await first.acquire(); + const abortedAcquisition = aborted.acquire(abortController.signal); + const replacementAcquisition = replacement.acquire(); + + abortController.abort(); + first.release(); + + await expect(abortedAcquisition).rejects.toThrow(); + await expect(replacementAcquisition).resolves.toBeUndefined(); + }); + + it("latches cleanup failure for current and future waiters", async () => { + const first = createLiveToken(); + const waiting = createLiveToken(); + await first.acquire(); + const waitingAcquisition = waiting.acquire(); + + first.fail(new Error("native cleanup failed")); + + await expect(waitingAcquisition).rejects.toBeInstanceOf(LlamaCppInferenceRestartRequiredError); + await expect(createLiveToken().acquire()).rejects.toBeInstanceOf( + LlamaCppInferenceRestartRequiredError, + ); + }); +}); diff --git a/extensions/llama-cpp/src/inference-runtime-coordinator.ts b/extensions/llama-cpp/src/inference-runtime-coordinator.ts new file mode 100644 index 000000000000..288f1c0a106c --- /dev/null +++ b/extensions/llama-cpp/src/inference-runtime-coordinator.ts @@ -0,0 +1,174 @@ +import { resolveGlobalSingleton } from "openclaw/plugin-sdk/global-singleton"; + +const COORDINATOR_KEY = Symbol.for("openclaw.llamaCppInferenceRuntimeCoordinator"); +const TEST_API_KEY = Symbol.for("openclaw.llamaCppInferenceTestApi"); +const RESTART_REQUIRED_CODE = "LLAMA_CPP_INFERENCE_RESTART_REQUIRED"; + +type Completion = { + promise: Promise; + complete: () => void; +}; + +type CoordinatorState = { + owner?: Completion; + restartRequired?: Error; +}; + +type RuntimeToken = { + closing: Completion; + owner?: Completion; + released: boolean; +}; + +export type LlamaCppInferenceRuntimeToken = { + acquire: (params: { + signal?: AbortSignal; + isLive: () => boolean; + onRestartRequired: () => void; + }) => Promise; + close: () => void; + fail: (error: unknown) => void; + release: () => void; +}; + +export class LlamaCppInferenceRestartRequiredError extends Error { + readonly code = RESTART_REQUIRED_CODE; + + constructor(cause: Error) { + super("A previous llama.cpp runtime failed to release native resources", { cause }); + this.name = "LlamaCppInferenceRestartRequiredError"; + } +} + +function getCoordinatorState(): CoordinatorState { + return resolveGlobalSingleton(COORDINATOR_KEY, () => ({})); +} + +function createCompletion(): Completion { + let complete!: () => void; + const promise = new Promise((resolve) => { + complete = resolve; + }); + return { promise, complete }; +} + +function abortedError(signal: AbortSignal): Error { + return signal.reason instanceof Error ? signal.reason : new Error("Request was aborted"); +} + +function waitForTurn(params: { + closing: Promise; + retired: Promise; + signal?: AbortSignal; +}): Promise<"closed" | "retired"> { + const signal = params.signal; + if (signal?.aborted) { + return Promise.reject(abortedError(signal)); + } + return new Promise((resolve, reject) => { + const cleanup = () => signal?.removeEventListener("abort", abort); + const abort = () => { + cleanup(); + reject(signal ? abortedError(signal) : new Error("Request was aborted")); + }; + const finish = (result: "closed" | "retired") => { + cleanup(); + resolve(result); + }; + signal?.addEventListener("abort", abort, { once: true }); + void params.closing.then(() => finish("closed")); + void params.retired.then(() => finish("retired")); + }); +} + +function isRestartRequiredError(error: unknown): error is LlamaCppInferenceRestartRequiredError { + return error instanceof Error && "code" in error && error.code === RESTART_REQUIRED_CODE; +} + +async function acquireRuntime( + token: RuntimeToken, + params: { signal?: AbortSignal; isLive: () => boolean }, +): Promise { + if (token.owner) { + return; + } + while (true) { + if (token.released || !params.isLive()) { + throw new Error("llama.cpp runtime is stopping"); + } + if (params.signal?.aborted) { + throw abortedError(params.signal); + } + const state = getCoordinatorState(); + if (state.restartRequired) { + throw new LlamaCppInferenceRestartRequiredError(state.restartRequired); + } + if (!state.owner) { + const owner = createCompletion(); + state.owner = owner; + token.owner = owner; + return; + } + if ( + (await waitForTurn({ + closing: token.closing.promise, + retired: state.owner.promise, + signal: params.signal, + })) === "closed" + ) { + throw new Error("llama.cpp runtime is stopping"); + } + } +} + +function releaseRuntime(token: RuntimeToken): void { + token.released = true; + if (!token.owner) { + return; + } + const state = getCoordinatorState(); + if (state.owner === token.owner) { + state.owner = undefined; + } + token.owner.complete(); + token.owner = undefined; +} + +function failRuntime(token: RuntimeToken, error: unknown): void { + const state = getCoordinatorState(); + state.restartRequired ??= error instanceof Error ? error : new Error(String(error)); + token.owner?.complete(); +} + +export function createLlamaCppInferenceRuntimeToken(): LlamaCppInferenceRuntimeToken { + // Registration creates only a local token; native ownership remains lazy. + const token: RuntimeToken = { closing: createCompletion(), released: false }; + return { + acquire: async ({ onRestartRequired, ...params }) => { + try { + await acquireRuntime(token, params); + } catch (error) { + // Reloaded plugin chunks have distinct class identities. + if (isRestartRequiredError(error)) { + onRestartRequired(); + } + throw error; + } + }, + close: () => token.closing.complete(), + fail: (error) => failRuntime(token, error), + release: () => releaseRuntime(token), + }; +} + +if (process.env.VITEST || process.env.NODE_ENV === "test") { + const globalStore = globalThis as Record; + const testApi = (globalStore[TEST_API_KEY] ?? {}) as Record; + testApi.resetInferenceRuntimeCoordinator = () => { + const state = getCoordinatorState(); + state.owner?.complete(); + state.owner = undefined; + state.restartRequired = undefined; + }; + globalStore[TEST_API_KEY] = testApi; +}