mirror of
https://github.com/openclaw/openclaw.git
synced 2026-08-12 21:53:00 -06:00
fix(lmstudio): resolve JIT embedding variants
This commit is contained in:
@@ -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({
|
||||
|
||||
@@ -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({
|
||||
|
||||
Reference in New Issue
Block a user