From 258e82a0e94b22e2f65b18003856fe6e3a8fff2d Mon Sep 17 00:00:00 2001 From: Daniel Riccio Date: Tue, 24 Jun 2025 14:02:55 -0500 Subject: [PATCH] fix: integrate LmStudioHandler with centralized model cache (#5075) - Updated LmStudioHandler.fetchModel() to use getModels() from model cache - Removed custom getModelWithFetch() method for consistency with other providers - Updated createMessage() and completePrompt() to use await this.fetchModel() - Fixed tests to mock getModels instead of getLMStudioModels - Ensures LM Studio models show correct context length instead of default value Fixes #5075 --- src/api/providers/__tests__/lmstudio.spec.ts | 35 +++++--------- src/api/providers/lm-studio.ts | 51 ++++++++++---------- 2 files changed, 37 insertions(+), 49 deletions(-) diff --git a/src/api/providers/__tests__/lmstudio.spec.ts b/src/api/providers/__tests__/lmstudio.spec.ts index 954d166930..da5e731a28 100644 --- a/src/api/providers/__tests__/lmstudio.spec.ts +++ b/src/api/providers/__tests__/lmstudio.spec.ts @@ -58,9 +58,9 @@ vi.mock("openai", () => { } }) -// Mock LM Studio fetcher -vi.mock("../fetchers/lmstudio", () => ({ - getLMStudioModels: vi.fn(), +// Mock model cache +vi.mock("../fetchers/modelCache", () => ({ + getModels: vi.fn(), })) import type { Anthropic } from "@anthropic-ai/sdk" @@ -68,10 +68,10 @@ import type { ModelInfo } from "@roo-code/types" import { LmStudioHandler } from "../lm-studio" import type { ApiHandlerOptions } from "../../../shared/api" -import { getLMStudioModels } from "../fetchers/lmstudio" +import { getModels } from "../fetchers/modelCache" // Get the mocked function -const mockGetLMStudioModels = vi.mocked(getLMStudioModels) +const mockGetModels = vi.mocked(getModels) describe("LmStudioHandler", () => { let handler: LmStudioHandler @@ -98,7 +98,7 @@ describe("LmStudioHandler", () => { } handler = new LmStudioHandler(mockOptions) mockCreate.mockClear() - mockGetLMStudioModels.mockClear() + mockGetModels.mockClear() }) describe("constructor", () => { @@ -190,7 +190,7 @@ describe("LmStudioHandler", () => { it("should return fetched model info when available", async () => { // Mock the fetched models - mockGetLMStudioModels.mockResolvedValueOnce({ + mockGetModels.mockResolvedValueOnce({ "local-model": mockModelInfo, }) @@ -204,7 +204,7 @@ describe("LmStudioHandler", () => { it("should fallback to default when model not found in fetched models", async () => { // Mock fetched models without our target model - mockGetLMStudioModels.mockResolvedValueOnce({ + mockGetModels.mockResolvedValueOnce({ "other-model": mockModelInfo, }) @@ -219,32 +219,21 @@ describe("LmStudioHandler", () => { describe("fetchModel", () => { it("should fetch models successfully", async () => { - mockGetLMStudioModels.mockResolvedValueOnce({ + mockGetModels.mockResolvedValueOnce({ "local-model": mockModelInfo, }) const result = await handler.fetchModel() - expect(mockGetLMStudioModels).toHaveBeenCalledWith(mockOptions.lmStudioBaseUrl) + expect(mockGetModels).toHaveBeenCalledWith({ provider: "lmstudio" }) expect(result.id).toBe(mockOptions.lmStudioModelId) expect(result.info).toEqual(mockModelInfo) }) it("should handle fetch errors gracefully", async () => { - const consoleSpy = vi.spyOn(console, "warn").mockImplementation(() => {}) - mockGetLMStudioModels.mockRejectedValueOnce(new Error("Connection failed")) + mockGetModels.mockRejectedValueOnce(new Error("Connection failed")) - const result = await handler.fetchModel() - - expect(consoleSpy).toHaveBeenCalledWith( - "Failed to fetch LM Studio models, using defaults:", - expect.any(Error), - ) - expect(result.id).toBe(mockOptions.lmStudioModelId) - expect(result.info.maxTokens).toBe(-1) - expect(result.info.contextWindow).toBe(128_000) - - consoleSpy.mockRestore() + await expect(handler.fetchModel()).rejects.toThrow("Connection failed") }) }) }) diff --git a/src/api/providers/lm-studio.ts b/src/api/providers/lm-studio.ts index 27d1090246..833afd4f42 100644 --- a/src/api/providers/lm-studio.ts +++ b/src/api/providers/lm-studio.ts @@ -1,10 +1,9 @@ import { Anthropic } from "@anthropic-ai/sdk" import OpenAI from "openai" -import axios from "axios" import { type ModelInfo, openAiModelInfoSaneDefaults, LMSTUDIO_DEFAULT_TEMPERATURE } from "@roo-code/types" -import type { ApiHandlerOptions } from "../../shared/api" +import type { ApiHandlerOptions, ModelRecord } from "../../shared/api" import { XmlMatcher } from "../../utils/xml-matcher" @@ -13,7 +12,7 @@ import { ApiStream } from "../transform/stream" import { BaseProvider } from "./base-provider" import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from "../index" -import { getLMStudioModels } from "./fetchers/lmstudio" +import { getModels } from "./fetchers/modelCache" export class LmStudioHandler extends BaseProvider implements SingleCompletionHandler { protected options: ApiHandlerOptions @@ -75,8 +74,9 @@ export class LmStudioHandler extends BaseProvider implements SingleCompletionHan let assistantText = "" try { + const modelInfo = await this.fetchModel() const params: OpenAI.Chat.ChatCompletionCreateParamsStreaming & { draft_model?: string } = { - model: this.getModel().id, + model: modelInfo.id, messages: openAiMessages, temperature: this.options.modelTemperature ?? LMSTUDIO_DEFAULT_TEMPERATURE, stream: true, @@ -133,12 +133,7 @@ export class LmStudioHandler extends BaseProvider implements SingleCompletionHan } public async fetchModel() { - try { - this.models = await getLMStudioModels(this.options.lmStudioBaseUrl) - } catch (error) { - console.warn("Failed to fetch LM Studio models, using defaults:", error) - this.models = {} - } + this.models = await getModels({ provider: "lmstudio" }) return this.getModel() } @@ -146,7 +141,24 @@ export class LmStudioHandler extends BaseProvider implements SingleCompletionHan const id = this.options.lmStudioModelId || "" // Try to get the actual model info from fetched models - const info = this.models[id] || openAiModelInfoSaneDefaults + // The fetcher uses model.path or modelKey as keys, so try both + let info: ModelInfo | undefined = undefined + + if (this.models && Object.keys(this.models).length > 0) { + info = this.models[id] + + // If not found by exact ID, try to find by partial match (for model paths) + if (!info) { + const modelKeys = Object.keys(this.models) + const matchingKey = modelKeys.find((key) => key === id || key.includes(id) || id.includes(key)) + if (matchingKey) { + info = this.models[matchingKey] + } + } + } + + // Fall back to defaults if still not found + info = info || openAiModelInfoSaneDefaults return { id, @@ -157,8 +169,9 @@ export class LmStudioHandler extends BaseProvider implements SingleCompletionHan async completePrompt(prompt: string): Promise { try { // Create params object with optional draft model + const modelInfo = await this.fetchModel() const params: any = { - model: this.getModel().id, + model: modelInfo.id, messages: [{ role: "user", content: prompt }], temperature: this.options.modelTemperature ?? LMSTUDIO_DEFAULT_TEMPERATURE, stream: false, @@ -178,17 +191,3 @@ export class LmStudioHandler extends BaseProvider implements SingleCompletionHan } } } - -export async function getLmStudioModels(baseUrl = "http://localhost:1234") { - try { - if (!URL.canParse(baseUrl)) { - return [] - } - - const response = await axios.get(`${baseUrl}/v1/models`) - const modelsArray = response.data?.data?.map((model: any) => model.id) || [] - return [...new Set(modelsArray)] - } catch (error) { - return [] - } -}