diff --git a/src/api/providers/__tests__/lmstudio.spec.ts b/src/api/providers/__tests__/lmstudio.spec.ts index 2679d225df4..954d1669308 100644 --- a/src/api/providers/__tests__/lmstudio.spec.ts +++ b/src/api/providers/__tests__/lmstudio.spec.ts @@ -58,23 +58,47 @@ vi.mock("openai", () => { } }) +// Mock LM Studio fetcher +vi.mock("../fetchers/lmstudio", () => ({ + getLMStudioModels: vi.fn(), +})) + import type { Anthropic } from "@anthropic-ai/sdk" +import type { ModelInfo } from "@roo-code/types" import { LmStudioHandler } from "../lm-studio" import type { ApiHandlerOptions } from "../../../shared/api" +import { getLMStudioModels } from "../fetchers/lmstudio" + +// Get the mocked function +const mockGetLMStudioModels = vi.mocked(getLMStudioModels) describe("LmStudioHandler", () => { let handler: LmStudioHandler let mockOptions: ApiHandlerOptions + const mockModelInfo: ModelInfo = { + maxTokens: 8192, + contextWindow: 32768, + supportsImages: false, + supportsComputerUse: false, + supportsPromptCache: true, + inputPrice: 0, + outputPrice: 0, + cacheWritesPrice: 0, + cacheReadsPrice: 0, + description: "Test Model - local-model", + } + beforeEach(() => { mockOptions = { apiModelId: "local-model", lmStudioModelId: "local-model", - lmStudioBaseUrl: "http://localhost:1234/v1", + lmStudioBaseUrl: "http://localhost:1234", } handler = new LmStudioHandler(mockOptions) mockCreate.mockClear() + mockGetLMStudioModels.mockClear() }) describe("constructor", () => { @@ -156,12 +180,71 @@ describe("LmStudioHandler", () => { }) describe("getModel", () => { - it("should return model info", () => { + it("should return default model info when no models fetched", () => { const modelInfo = handler.getModel() expect(modelInfo.id).toBe(mockOptions.lmStudioModelId) expect(modelInfo.info).toBeDefined() expect(modelInfo.info.maxTokens).toBe(-1) expect(modelInfo.info.contextWindow).toBe(128_000) }) + + it("should return fetched model info when available", async () => { + // Mock the fetched models + mockGetLMStudioModels.mockResolvedValueOnce({ + "local-model": mockModelInfo, + }) + + await handler.fetchModel() + const modelInfo = handler.getModel() + + expect(modelInfo.id).toBe(mockOptions.lmStudioModelId) + expect(modelInfo.info).toEqual(mockModelInfo) + expect(modelInfo.info.contextWindow).toBe(32768) + }) + + it("should fallback to default when model not found in fetched models", async () => { + // Mock fetched models without our target model + mockGetLMStudioModels.mockResolvedValueOnce({ + "other-model": mockModelInfo, + }) + + await handler.fetchModel() + const modelInfo = handler.getModel() + + expect(modelInfo.id).toBe(mockOptions.lmStudioModelId) + expect(modelInfo.info.maxTokens).toBe(-1) + expect(modelInfo.info.contextWindow).toBe(128_000) + }) + }) + + describe("fetchModel", () => { + it("should fetch models successfully", async () => { + mockGetLMStudioModels.mockResolvedValueOnce({ + "local-model": mockModelInfo, + }) + + const result = await handler.fetchModel() + + expect(mockGetLMStudioModels).toHaveBeenCalledWith(mockOptions.lmStudioBaseUrl) + 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")) + + 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() + }) }) }) diff --git a/src/api/providers/lm-studio.ts b/src/api/providers/lm-studio.ts index f032e2d5605..27d1090246b 100644 --- a/src/api/providers/lm-studio.ts +++ b/src/api/providers/lm-studio.ts @@ -13,10 +13,12 @@ import { ApiStream } from "../transform/stream" import { BaseProvider } from "./base-provider" import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from "../index" +import { getLMStudioModels } from "./fetchers/lmstudio" export class LmStudioHandler extends BaseProvider implements SingleCompletionHandler { protected options: ApiHandlerOptions private client: OpenAI + private models: Record = {} constructor(options: ApiHandlerOptions) { super() @@ -130,10 +132,25 @@ 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 = {} + } + return this.getModel() + } + override getModel(): { id: string; info: ModelInfo } { + const id = this.options.lmStudioModelId || "" + + // Try to get the actual model info from fetched models + const info = this.models[id] || openAiModelInfoSaneDefaults + return { - id: this.options.lmStudioModelId || "", - info: openAiModelInfoSaneDefaults, + id, + info, } }