mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-09-05 08:10:14 +00:00
feat: add openai-compatible provider support for token usage display
- Add 'openai-compatible' as a valid provider in buildApiHandler - Add 'openai-compatible' to provider types and schemas - Update ProfileValidator to handle openai-compatible provider - Add tests for openai-compatible provider functionality Fixes #8543 - Token usage now displays correctly when using OpenAI Compatible API provider
This commit is contained in:
parent
5a3f911321
commit
c20ef5a27a
4 changed files with 146 additions and 2 deletions
|
|
@ -88,7 +88,7 @@ export const isInternalProvider = (key: string): key is InternalProvider =>
|
|||
* Custom providers are completely configurable within Roo Code settings.
|
||||
*/
|
||||
|
||||
export const customProviders = ["openai"] as const
|
||||
export const customProviders = ["openai", "openai-compatible"] as const
|
||||
|
||||
export type CustomProvider = (typeof customProviders)[number]
|
||||
|
||||
|
|
@ -138,6 +138,7 @@ export const providerNames = [
|
|||
"vertex",
|
||||
"xai",
|
||||
"zai",
|
||||
"openai-compatible",
|
||||
] as const
|
||||
|
||||
export const providerNamesSchema = z.enum(providerNames)
|
||||
|
|
@ -424,6 +425,7 @@ export const providerSettingsSchemaDiscriminated = z.discriminatedUnion("apiProv
|
|||
bedrockSchema.merge(z.object({ apiProvider: z.literal("bedrock") })),
|
||||
vertexSchema.merge(z.object({ apiProvider: z.literal("vertex") })),
|
||||
openAiSchema.merge(z.object({ apiProvider: z.literal("openai") })),
|
||||
openAiSchema.merge(z.object({ apiProvider: z.literal("openai-compatible") })),
|
||||
ollamaSchema.merge(z.object({ apiProvider: z.literal("ollama") })),
|
||||
vsCodeLmSchema.merge(z.object({ apiProvider: z.literal("vscode-lm") })),
|
||||
lmStudioSchema.merge(z.object({ apiProvider: z.literal("lmstudio") })),
|
||||
|
|
@ -610,7 +612,7 @@ export const getApiProtocol = (provider: ProviderName | undefined, modelId?: str
|
|||
*/
|
||||
|
||||
export const MODELS_BY_PROVIDER: Record<
|
||||
Exclude<ProviderName, "fake-ai" | "human-relay" | "gemini-cli" | "openai">,
|
||||
Exclude<ProviderName, "fake-ai" | "human-relay" | "gemini-cli" | "openai" | "openai-compatible">,
|
||||
{ id: ProviderName; label: string; models: string[] }
|
||||
> = {
|
||||
anthropic: {
|
||||
|
|
|
|||
|
|
@ -106,6 +106,7 @@ export function buildApiHandler(configuration: ProviderSettings): ApiHandler {
|
|||
? new AnthropicVertexHandler(options)
|
||||
: new VertexHandler(options)
|
||||
case "openai":
|
||||
case "openai-compatible":
|
||||
return new OpenAiHandler(options)
|
||||
case "ollama":
|
||||
return new NativeOllamaHandler(options)
|
||||
|
|
|
|||
140
src/api/providers/__tests__/openai-compatible.spec.ts
Normal file
140
src/api/providers/__tests__/openai-compatible.spec.ts
Normal file
|
|
@ -0,0 +1,140 @@
|
|||
import { describe, it, expect, vi, beforeEach } from "vitest"
|
||||
import { buildApiHandler } from "../../index"
|
||||
import { OpenAiHandler } from "../openai"
|
||||
|
||||
vi.mock("openai", () => {
|
||||
const mockCreate = vi.fn()
|
||||
return {
|
||||
default: vi.fn().mockImplementation(() => ({
|
||||
chat: {
|
||||
completions: {
|
||||
create: mockCreate,
|
||||
},
|
||||
},
|
||||
})),
|
||||
OpenAI: vi.fn().mockImplementation(() => ({
|
||||
chat: {
|
||||
completions: {
|
||||
create: mockCreate,
|
||||
},
|
||||
},
|
||||
})),
|
||||
AzureOpenAI: vi.fn().mockImplementation(() => ({
|
||||
chat: {
|
||||
completions: {
|
||||
create: mockCreate,
|
||||
},
|
||||
},
|
||||
})),
|
||||
}
|
||||
})
|
||||
|
||||
describe("OpenAI Compatible Provider", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
})
|
||||
|
||||
it("should create OpenAiHandler when apiProvider is 'openai-compatible'", () => {
|
||||
const handler = buildApiHandler({
|
||||
apiProvider: "openai-compatible",
|
||||
openAiApiKey: "test-key",
|
||||
openAiBaseUrl: "https://api.example.com/v1",
|
||||
openAiModelId: "test-model",
|
||||
})
|
||||
|
||||
expect(handler).toBeInstanceOf(OpenAiHandler)
|
||||
})
|
||||
|
||||
it("should handle token usage correctly for openai-compatible provider", async () => {
|
||||
const mockStream = {
|
||||
async *[Symbol.asyncIterator]() {
|
||||
yield {
|
||||
choices: [{ delta: { content: "Hello" } }],
|
||||
}
|
||||
yield {
|
||||
choices: [{ delta: { content: " world" } }],
|
||||
}
|
||||
yield {
|
||||
choices: [{ delta: {} }],
|
||||
usage: {
|
||||
prompt_tokens: 10,
|
||||
completion_tokens: 5,
|
||||
total_tokens: 15,
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
const OpenAI = (await import("openai")).default
|
||||
const mockCreate = vi.fn().mockResolvedValue(mockStream)
|
||||
;(OpenAI as any).mockImplementation(() => ({
|
||||
chat: {
|
||||
completions: {
|
||||
create: mockCreate,
|
||||
},
|
||||
},
|
||||
}))
|
||||
|
||||
const handler = buildApiHandler({
|
||||
apiProvider: "openai-compatible",
|
||||
openAiApiKey: "test-key",
|
||||
openAiBaseUrl: "https://api.example.com/v1",
|
||||
openAiModelId: "test-model",
|
||||
})
|
||||
|
||||
const messages = [{ role: "user" as const, content: "Test message" }]
|
||||
const stream = handler.createMessage("System prompt", messages)
|
||||
|
||||
const chunks = []
|
||||
for await (const chunk of stream) {
|
||||
chunks.push(chunk)
|
||||
}
|
||||
|
||||
// Check that we got text chunks
|
||||
const textChunks = chunks.filter((c) => c.type === "text")
|
||||
expect(textChunks).toHaveLength(2)
|
||||
expect(textChunks[0].text).toBe("Hello")
|
||||
expect(textChunks[1].text).toBe(" world")
|
||||
|
||||
// Check that we got usage data
|
||||
const usageChunk = chunks.find((c) => c.type === "usage")
|
||||
expect(usageChunk).toBeDefined()
|
||||
expect(usageChunk).toEqual({
|
||||
type: "usage",
|
||||
inputTokens: 10,
|
||||
outputTokens: 5,
|
||||
})
|
||||
})
|
||||
|
||||
it("should use the same configuration as openai provider", () => {
|
||||
const config = {
|
||||
openAiApiKey: "test-key",
|
||||
openAiBaseUrl: "https://api.example.com/v1",
|
||||
openAiModelId: "test-model",
|
||||
openAiCustomModelInfo: {
|
||||
maxTokens: 4096,
|
||||
contextWindow: 8192,
|
||||
supportsPromptCache: false,
|
||||
inputPrice: 0.001,
|
||||
outputPrice: 0.002,
|
||||
},
|
||||
}
|
||||
|
||||
const openaiHandler = buildApiHandler({
|
||||
apiProvider: "openai",
|
||||
...config,
|
||||
})
|
||||
|
||||
const openaiCompatibleHandler = buildApiHandler({
|
||||
apiProvider: "openai-compatible",
|
||||
...config,
|
||||
})
|
||||
|
||||
// Both should be instances of OpenAiHandler
|
||||
expect(openaiHandler).toBeInstanceOf(OpenAiHandler)
|
||||
expect(openaiCompatibleHandler).toBeInstanceOf(OpenAiHandler)
|
||||
|
||||
// Both should have the same model configuration
|
||||
expect(openaiHandler.getModel()).toEqual(openaiCompatibleHandler.getModel())
|
||||
})
|
||||
})
|
||||
|
|
@ -56,6 +56,7 @@ export class ProfileValidator {
|
|||
private static getModelIdFromProfile(profile: ProviderSettings): string | undefined {
|
||||
switch (profile.apiProvider) {
|
||||
case "openai":
|
||||
case "openai-compatible":
|
||||
return profile.openAiModelId
|
||||
case "anthropic":
|
||||
case "openai-native":
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue