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
This commit is contained in:
Daniel Riccio 2025-06-24 14:02:55 -05:00
parent ac85aa6816
commit 258e82a0e9
No known key found for this signature in database
GPG key ID: FFD5FD825F8E8209
2 changed files with 37 additions and 49 deletions

View file

@ -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")
})
})
})

View file

@ -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<string> {
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<string>(modelsArray)]
} catch (error) {
return []
}
}