mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-10-07 02:58:15 +00:00
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:
parent
ac85aa6816
commit
258e82a0e9
2 changed files with 37 additions and 49 deletions
|
|
@ -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")
|
||||
})
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue