From 7c3f8b0a286f4f6b29684dfa1ced3a63fea6a28b Mon Sep 17 00:00:00 2001 From: Vincent Koc Date: Tue, 4 Aug 2026 01:33:03 +0800 Subject: [PATCH] fix(lmstudio): resolve JIT embedding variants --- .../lmstudio/src/embedding-provider.test.ts | 90 +++++++++++++++++-- extensions/lmstudio/src/embedding-provider.ts | 49 +++++++++- 2 files changed, 130 insertions(+), 9 deletions(-) diff --git a/extensions/lmstudio/src/embedding-provider.test.ts b/extensions/lmstudio/src/embedding-provider.test.ts index a92752b5b1b2..7248d8693402 100644 --- a/extensions/lmstudio/src/embedding-provider.test.ts +++ b/extensions/lmstudio/src/embedding-provider.test.ts @@ -9,19 +9,42 @@ const ensureLmstudioModelLoadedMock = vi.hoisted(() => async (_params?: { requestedContextLength?: number }) => "text-embedding-nomic-embed-text-v1.5", ), ); +const fetchLmstudioModelsMock = vi.hoisted(() => + vi.fn(async (_params?: unknown) => ({ + reachable: true, + status: 200, + models: [] as Array<{ + type?: "llm" | "embedding"; + key?: string; + variants?: unknown; + selected_variant?: unknown; + loaded_instances?: unknown[]; + }>, + })), +); const resolveLmstudioProviderHeadersMock = vi.hoisted(() => vi.fn(async (_params?: unknown) => undefined), ); const resolveLmstudioRuntimeApiKeyMock = vi.hoisted(() => vi.fn(async (_params?: unknown) => undefined), ); +const embeddedModels = vi.hoisted(() => [] as string[]); const createRemoteEmbeddingProviderMock = vi.hoisted(() => - vi.fn(() => ({ - id: "lmstudio", - model: "text-embedding-nomic-embed-text-v1.5", - embedQuery: vi.fn(async () => [1, 0]), - embedBatch: vi.fn(async (texts: string[]) => texts.map(() => [1, 0])), - })), + vi.fn((params: { client: { model: string } }) => { + const providerModel = params.client.model; + return { + id: "lmstudio", + model: providerModel, + embedQuery: vi.fn(async () => { + embeddedModels.push(params.client.model); + return [1, 0]; + }), + embedBatch: vi.fn(async (texts: string[]) => { + embeddedModels.push(params.client.model); + return texts.map(() => [1, 0]); + }), + }; + }), ); vi.mock("openclaw/plugin-sdk/memory-core-host-engine-embeddings", async (importOriginal) => { @@ -39,6 +62,7 @@ vi.mock("./models.fetch.js", async (importOriginal) => { ...actual, ensureLmstudioModelLoaded: (params: { requestedContextLength?: number }) => ensureLmstudioModelLoadedMock(params), + fetchLmstudioModels: (params: unknown) => fetchLmstudioModelsMock(params), }; }); @@ -84,7 +108,10 @@ async function readRequestedContextLength(config: OpenClawConfig): Promise { beforeEach(() => { ensureLmstudioModelLoadedMock.mockClear(); + fetchLmstudioModelsMock.mockClear(); + fetchLmstudioModelsMock.mockResolvedValue({ reachable: true, status: 200, models: [] }); createRemoteEmbeddingProviderMock.mockClear(); + embeddedModels.length = 0; resolveLmstudioProviderHeadersMock.mockClear(); resolveLmstudioRuntimeApiKeyMock.mockClear(); }); @@ -159,6 +186,7 @@ describe("createLmstudioEmbeddingProvider preload context length", () => { }); expect(ensureLmstudioModelLoadedMock).not.toHaveBeenCalled(); + expect(fetchLmstudioModelsMock).not.toHaveBeenCalled(); expect(acquireLocalService).not.toHaveBeenCalled(); await expect(provider.embedQuery("hello")).resolves.toEqual([1, 0]); @@ -172,6 +200,56 @@ describe("createLmstudioEmbeddingProvider preload context length", () => { }, ); + it("resolves a JIT variant before freezing provider and cache identity", async () => { + const requestedVariant = `${EMBEDDING_MODEL}@q4_k_m`; + const release = vi.fn(); + const acquireLocalService = vi.fn(async () => ({ release })); + fetchLmstudioModelsMock.mockResolvedValueOnce({ + reachable: true, + status: 200, + models: [ + { + type: "embedding", + key: EMBEDDING_MODEL, + variants: [requestedVariant], + selected_variant: requestedVariant, + loaded_instances: [], + }, + ], + }); + + const options = { + config: buildConfig({ + model: { id: requestedVariant }, + provider: { + params: { preload: false }, + localService: { command: "/usr/bin/lms" }, + }, + }), + provider: "lmstudio", + model: `lmstudio/${requestedVariant}`, + fallback: "none", + acquireLocalService, + }; + const result = await lmstudioMemoryEmbeddingProviderAdapter.create(options); + if (!result.provider) { + throw new Error("expected LM Studio embedding provider"); + } + + expect(ensureLmstudioModelLoadedMock).not.toHaveBeenCalled(); + expect(fetchLmstudioModelsMock).toHaveBeenCalledOnce(); + expect(acquireLocalService).toHaveBeenCalledOnce(); + expect(release).toHaveBeenCalledOnce(); + expect(result.provider.model).toBe(EMBEDDING_MODEL); + expect(result.runtime?.cacheKeyData).toMatchObject({ model: EMBEDDING_MODEL }); + + await expect(result.provider.embedQuery("hello")).resolves.toEqual([1, 0]); + + expect(embeddedModels).toEqual([EMBEDDING_MODEL]); + expect(acquireLocalService).toHaveBeenCalledTimes(2); + expect(release).toHaveBeenCalledTimes(2); + }); + it("uses the canonical preloaded model for embedding requests and memory identity", async () => { const requestedVariant = `${EMBEDDING_MODEL}@q4_k_m`; const result = await lmstudioMemoryEmbeddingProviderAdapter.create({ diff --git a/extensions/lmstudio/src/embedding-provider.ts b/extensions/lmstudio/src/embedding-provider.ts index 7dcafda0990d..d233d4276278 100644 --- a/extensions/lmstudio/src/embedding-provider.ts +++ b/extensions/lmstudio/src/embedding-provider.ts @@ -12,9 +12,10 @@ import { normalizeProviderId } from "openclaw/plugin-sdk/provider-model-shared"; import { formatErrorMessage, type SsrFPolicy } from "openclaw/plugin-sdk/ssrf-runtime"; import { asPositiveSafeInteger } from "openclaw/plugin-sdk/string-coerce-runtime"; import { LMSTUDIO_DEFAULT_EMBEDDING_MODEL, LMSTUDIO_PROVIDER_ID } from "./defaults.js"; -import { ensureLmstudioModelLoaded } from "./models.fetch.js"; +import { ensureLmstudioModelLoaded, fetchLmstudioModels } from "./models.fetch.js"; import { normalizeLmstudioConfiguredCatalogEntries, + resolveLmstudioCanonicalModelKey, resolveLmstudioInferenceBase, resolveLmstudioServerBase, } from "./models.js"; @@ -156,9 +157,31 @@ function resolveLmstudioLocalServiceBaseUrl( return /\/api\/v1$/iu.test(configuredPath) ? `${serverBaseUrl}/api/v1` : `${serverBaseUrl}/v1`; } +async function resolveLmstudioEmbeddingModelKey(params: { + baseUrl: string; + apiKey?: string; + headers: Record; + ssrfPolicy?: SsrFPolicy; + model: string; +}): Promise { + const discovered = await fetchLmstudioModels({ + baseUrl: params.baseUrl, + apiKey: params.apiKey, + headers: params.headers, + ssrfPolicy: params.ssrfPolicy, + }); + if (!discovered.reachable || (discovered.status !== undefined && discovered.status >= 400)) { + return params.model; + } + return resolveLmstudioCanonicalModelKey({ + modelKey: params.model, + models: discovered.models, + }); +} + /** Creates the LM Studio embedding provider client and preloads the target model before return. */ export async function createLmstudioEmbeddingProvider( - options: MemoryEmbeddingProviderCreateOptions, + options: LocalServiceAwareEmbeddingOptions, ): Promise<{ provider: MemoryEmbeddingProvider; client: LmstudioEmbeddingClient }> { const resolvedProvider = resolveConfiguredLmstudioProvider(options); const providerConfig = resolvedProvider?.config; @@ -225,7 +248,7 @@ export async function createLmstudioEmbeddingProvider( headers, } : undefined; - const acquireLocalService = (options as LocalServiceAwareEmbeddingOptions).acquireLocalService; + const acquireLocalService = options.acquireLocalService; const withLocalServiceLease = async ( signal: AbortSignal | undefined, action: () => Promise, @@ -269,6 +292,26 @@ export async function createLmstudioEmbeddingProvider( }); } }); + } else if (model.includes("@")) { + // Variant aliases are not accepted by LM Studio's inference routes. Resolve + // only the stable wire/cache identity here; JIT still owns the actual load. + try { + await withLocalServiceLease(undefined, async () => { + client.model = await resolveLmstudioEmbeddingModelKey({ + baseUrl, + apiKey, + headers: headerOverrides, + ssrfPolicy, + model, + }); + }); + } catch (error) { + log.debug("lmstudio embedding variant discovery failed; using requested model", { + baseUrl, + model, + error: formatErrorMessage(error), + }); + } } const remoteProvider = createRemoteEmbeddingProvider({