fix(lmstudio): resolve JIT embedding variants

This commit is contained in:
Vincent Koc
2026-08-04 01:33:03 +08:00
parent f1e7e20bad
commit 7c3f8b0a28
2 changed files with 130 additions and 9 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();
});
@@ -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({
+46 -3
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>,
@@ -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({