fix(lmstudio): honor embedding JIT and canonical model identity (#117320)

* codex/qa400-lmstudio-jit-canonical-identity:
  fix(lmstudio): resolve JIT embedding variants
  fix(lmstudio): preserve embedding preload and model identity
This commit is contained in:
Vincent Koc
2026-08-04 02:22:51 +08:00
2 changed files with 245 additions and 20 deletions
@@ -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<unkno
describe("createLmstudioEmbeddingProvider preload context length", () => {
beforeEach(() => {
ensureLmstudioModelLoadedMock.mockClear();
fetchLmstudioModelsMock.mockClear();
fetchLmstudioModelsMock.mockResolvedValue({ reachable: true, status: 200, models: [] });
createRemoteEmbeddingProviderMock.mockClear();
embeddedModels.length = 0;
resolveLmstudioProviderHeadersMock.mockClear();
resolveLmstudioRuntimeApiKeyMock.mockClear();
});
@@ -134,6 +161,151 @@ describe("createLmstudioEmbeddingProvider preload context length", () => {
);
});
it.each(["lmstudio", "lmstudio-spark"])(
"honors the preload opt-out for %s while retaining request-time service leases",
async (providerId) => {
const release = vi.fn();
const acquireLocalService = vi.fn(async () => ({ release }));
const { provider } = await createLmstudioEmbeddingProvider({
config: {
models: {
providers: {
[providerId]: {
baseUrl: "http://spark.local:1234/v1",
params: { preload: false },
localService: { command: "/usr/bin/lms-spark" },
models: [{ id: EMBEDDING_MODEL }],
},
},
},
} as unknown as OpenClawConfig,
provider: providerId,
model: `${providerId}/${EMBEDDING_MODEL}`,
fallback: "none",
acquireLocalService,
});
expect(ensureLmstudioModelLoadedMock).not.toHaveBeenCalled();
expect(fetchLmstudioModelsMock).not.toHaveBeenCalled();
expect(acquireLocalService).not.toHaveBeenCalled();
await expect(provider.embedQuery("hello")).resolves.toEqual([1, 0]);
expect(acquireLocalService).toHaveBeenCalledOnce();
expect(acquireLocalService).toHaveBeenCalledWith(
expect.objectContaining({ providerId, baseUrl: "http://spark.local:1234/v1" }),
undefined,
);
expect(release).toHaveBeenCalledOnce();
},
);
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({
config: buildConfig({ model: { id: requestedVariant } }),
provider: "lmstudio",
model: `lmstudio/${requestedVariant}`,
fallback: "none",
});
expect(ensureLmstudioModelLoadedMock).toHaveBeenCalledWith(
expect.objectContaining({ modelKey: requestedVariant }),
);
expect(createRemoteEmbeddingProviderMock).toHaveBeenCalledWith(
expect.objectContaining({ client: expect.objectContaining({ model: EMBEDDING_MODEL }) }),
);
expect(result.provider?.model).toBe(EMBEDDING_MODEL);
expect(result.runtime?.cacheKeyData).toMatchObject({ model: EMBEDDING_MODEL });
});
it("retains the discovered canonical model when its preload subsequently fails", async () => {
const requestedVariant = `${EMBEDDING_MODEL}@q4_k_m`;
ensureLmstudioModelLoadedMock.mockRejectedValueOnce(
Object.assign(new Error("fixture preload rejected"), { resolvedModelKey: EMBEDDING_MODEL }),
);
const result = await lmstudioMemoryEmbeddingProviderAdapter.create({
config: buildConfig({ model: { id: requestedVariant } }),
provider: "lmstudio",
model: `lmstudio/${requestedVariant}`,
fallback: "none",
});
expect(createRemoteEmbeddingProviderMock).toHaveBeenCalledWith(
expect.objectContaining({ client: expect.objectContaining({ model: EMBEDDING_MODEL }) }),
);
expect(result.provider?.model).toBe(EMBEDDING_MODEL);
expect(result.runtime?.cacheKeyData).toMatchObject({ model: EMBEDDING_MODEL });
});
it("keeps the requested model when preload fails before discovering its identity", async () => {
const requestedVariant = `${EMBEDDING_MODEL}@q4_k_m`;
ensureLmstudioModelLoadedMock.mockRejectedValueOnce(new Error("fixture discovery unavailable"));
const { client } = await createLmstudioEmbeddingProvider({
config: buildConfig({ model: { id: requestedVariant } }),
provider: "lmstudio",
model: `lmstudio/${requestedVariant}`,
fallback: "none",
});
expect(client.model).toBe(requestedVariant);
expect(createRemoteEmbeddingProviderMock).toHaveBeenCalledWith(
expect.objectContaining({ client: expect.objectContaining({ model: requestedVariant }) }),
);
});
it("leases the exact configured alias for preload and embedding requests", async () => {
const release = vi.fn();
const acquireLocalService = vi.fn(async (_target: unknown) => ({ release }));
+67 -14
View File
@@ -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<string, string>;
ssrfPolicy?: SsrFPolicy;
model: string;
}): Promise<string> {
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 <T>(
signal: AbortSignal | undefined,
action: () => Promise<T>,
@@ -241,25 +264,55 @@ export async function createLmstudioEmbeddingProvider(
}
};
await withLocalServiceLease(undefined, async () => {
// The provider-owned JIT opt-out applies to embeddings as well as chat.
if (providerConfig?.params?.preload !== false) {
await withLocalServiceLease(undefined, async () => {
try {
client.model = await ensureLmstudioModelLoaded({
baseUrl,
apiKey,
headers: headerOverrides,
ssrfPolicy,
modelKey: model,
requestedContextLength,
timeoutMs: 120_000,
});
} catch (error) {
// Discovery still identifies the wire model when the subsequent load fails.
if (error instanceof Error && "resolvedModelKey" in error) {
const resolvedModelKey = error.resolvedModelKey;
if (typeof resolvedModelKey === "string" && resolvedModelKey.trim()) {
client.model = resolvedModelKey.trim();
}
}
log.warn("lmstudio embeddings warmup failed; continuing without preload", {
baseUrl,
model,
error: formatErrorMessage(error),
});
}
});
} 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 ensureLmstudioModelLoaded({
baseUrl,
apiKey,
headers: headerOverrides,
ssrfPolicy,
modelKey: model,
requestedContextLength,
timeoutMs: 120_000,
await withLocalServiceLease(undefined, async () => {
client.model = await resolveLmstudioEmbeddingModelKey({
baseUrl,
apiKey,
headers: headerOverrides,
ssrfPolicy,
model,
});
});
} catch (error) {
log.warn("lmstudio embeddings warmup failed; continuing without preload", {
log.debug("lmstudio embedding variant discovery failed; using requested model", {
baseUrl,
model,
error: formatErrorMessage(error),
});
}
});
}
const remoteProvider = createRemoteEmbeddingProvider({
id: LMSTUDIO_PROVIDER_ID,