diff --git a/.changeset/azure-ai-sdk-migration.md b/.changeset/azure-ai-sdk-migration.md new file mode 100644 index 0000000000..9b6fe14452 --- /dev/null +++ b/.changeset/azure-ai-sdk-migration.md @@ -0,0 +1,13 @@ +--- +"roo-cline": minor +"@roo-code/types": minor +--- + +Add dedicated Azure OpenAI provider using @ai-sdk/azure package + +- Add new "azure" provider type to support Azure OpenAI deployments via the AI SDK +- Implement AzureHandler following the established pattern from DeepSeek, Groq, and Fireworks migrations +- Add azureSchema with Azure-specific options: azureApiKey, azureResourceName, azureDeploymentName, azureApiVersion +- Use streamText/generateText from the AI SDK for cleaner streaming implementation +- Support tool calling via tool-input-start/delta/end events +- Include cache metrics extraction from providerMetadata diff --git a/packages/types/src/provider-settings.ts b/packages/types/src/provider-settings.ts index 555513500b..f8d0a5556b 100644 --- a/packages/types/src/provider-settings.ts +++ b/packages/types/src/provider-settings.ts @@ -119,6 +119,7 @@ export const providerNames = [ ...customProviders, ...fauxProviders, "anthropic", + "azure", "bedrock", "baseten", "cerebras", @@ -413,12 +414,20 @@ const basetenSchema = apiModelIdProviderModelSchema.extend({ basetenApiKey: z.string().optional(), }) +const azureSchema = apiModelIdProviderModelSchema.extend({ + azureApiKey: z.string().optional(), + azureResourceName: z.string().optional(), + azureDeploymentName: z.string().optional(), + azureApiVersion: z.string().optional(), +}) + const defaultSchema = z.object({ apiProvider: z.undefined(), }) export const providerSettingsSchemaDiscriminated = z.discriminatedUnion("apiProvider", [ anthropicSchema.merge(z.object({ apiProvider: z.literal("anthropic") })), + azureSchema.merge(z.object({ apiProvider: z.literal("azure") })), openRouterSchema.merge(z.object({ apiProvider: z.literal("openrouter") })), bedrockSchema.merge(z.object({ apiProvider: z.literal("bedrock") })), vertexSchema.merge(z.object({ apiProvider: z.literal("vertex") })), @@ -460,6 +469,7 @@ export const providerSettingsSchemaDiscriminated = z.discriminatedUnion("apiProv export const providerSettingsSchema = z.object({ apiProvider: providerNamesSchema.optional(), ...anthropicSchema.shape, + ...azureSchema.shape, ...openRouterSchema.shape, ...bedrockSchema.shape, ...vertexSchema.shape, @@ -548,6 +558,7 @@ export const isTypicalProvider = (key: unknown): key is TypicalProvider => export const modelIdKeysByProvider: Record = { anthropic: "apiModelId", + azure: "apiModelId", openrouter: "openRouterModelId", bedrock: "apiModelId", vertex: "apiModelId", @@ -624,6 +635,11 @@ export const MODELS_BY_PROVIDER: Record< label: "Anthropic", models: Object.keys(anthropicModels), }, + azure: { + id: "azure", + label: "Azure OpenAI", + models: [], // Azure uses deployment names configured by the user + }, bedrock: { id: "bedrock", label: "Amazon Bedrock", diff --git a/src/api/index.ts b/src/api/index.ts index 0e25a739a6..1b0aaf4480 100644 --- a/src/api/index.ts +++ b/src/api/index.ts @@ -8,6 +8,7 @@ import { ApiStream } from "./transform/stream" import { AnthropicHandler, AwsBedrockHandler, + AzureHandler, CerebrasHandler, OpenRouterHandler, VertexHandler, @@ -134,6 +135,8 @@ export function buildApiHandler(configuration: ProviderSettings): ApiHandler { switch (apiProvider) { case "anthropic": return new AnthropicHandler(options) + case "azure": + return new AzureHandler(options) case "openrouter": return new OpenRouterHandler(options) case "bedrock": diff --git a/src/api/providers/__tests__/azure.spec.ts b/src/api/providers/__tests__/azure.spec.ts new file mode 100644 index 0000000000..c93fe6ca59 --- /dev/null +++ b/src/api/providers/__tests__/azure.spec.ts @@ -0,0 +1,371 @@ +// Use vi.hoisted to define mock functions that can be referenced in hoisted vi.mock() calls +const { mockStreamText, mockGenerateText } = vi.hoisted(() => ({ + mockStreamText: vi.fn(), + mockGenerateText: vi.fn(), +})) + +vi.mock("ai", async (importOriginal) => { + const actual = await importOriginal() + return { + ...actual, + streamText: mockStreamText, + generateText: mockGenerateText, + } +}) + +vi.mock("@ai-sdk/azure", () => ({ + createAzure: vi.fn(() => { + // Return a function that returns a mock language model + return vi.fn(() => ({ + modelId: "gpt-4o", + provider: "azure", + })) + }), +})) + +import type { Anthropic } from "@anthropic-ai/sdk" + +import type { ApiHandlerOptions } from "../../../shared/api" + +import { AzureHandler } from "../azure" + +describe("AzureHandler", () => { + let handler: AzureHandler + let mockOptions: ApiHandlerOptions + + beforeEach(() => { + mockOptions = { + azureApiKey: "test-api-key", + azureResourceName: "test-resource", + azureDeploymentName: "gpt-4o", + azureApiVersion: "2024-08-01-preview", + } + handler = new AzureHandler(mockOptions) + vi.clearAllMocks() + }) + + describe("constructor", () => { + it("should initialize with provided options", () => { + expect(handler).toBeInstanceOf(AzureHandler) + expect(handler.getModel().id).toBe(mockOptions.azureDeploymentName) + }) + + it("should use apiModelId if azureDeploymentName is not provided", () => { + const handlerWithModelId = new AzureHandler({ + ...mockOptions, + azureDeploymentName: undefined, + apiModelId: "gpt-35-turbo", + }) + expect(handlerWithModelId.getModel().id).toBe("gpt-35-turbo") + }) + + it("should use empty string if neither azureDeploymentName nor apiModelId is provided", () => { + const handlerWithoutModel = new AzureHandler({ + ...mockOptions, + azureDeploymentName: undefined, + apiModelId: undefined, + }) + expect(handlerWithoutModel.getModel().id).toBe("") + }) + + it("should use default API version if not provided", () => { + const handlerWithoutVersion = new AzureHandler({ + ...mockOptions, + azureApiVersion: undefined, + }) + expect(handlerWithoutVersion).toBeInstanceOf(AzureHandler) + }) + }) + + describe("getModel", () => { + it("should return model info with deployment name as ID", () => { + const model = handler.getModel() + expect(model.id).toBe(mockOptions.azureDeploymentName) + expect(model.info).toBeDefined() + }) + + it("should include model parameters from getModelParams", () => { + const model = handler.getModel() + expect(model).toHaveProperty("temperature") + expect(model).toHaveProperty("maxTokens") + }) + }) + + describe("createMessage", () => { + const systemPrompt = "You are a helpful assistant." + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: [ + { + type: "text" as const, + text: "Hello!", + }, + ], + }, + ] + + it("should handle streaming responses", async () => { + // Mock the fullStream async generator + async function* mockFullStream() { + yield { type: "text-delta", text: "Test response" } + } + + // Mock usage and providerMetadata promises + const mockUsage = Promise.resolve({ + inputTokens: 10, + outputTokens: 5, + }) + + const mockProviderMetadata = Promise.resolve({ + azure: { + promptCacheHitTokens: 2, + promptCacheMissTokens: 8, + }, + }) + + mockStreamText.mockReturnValue({ + fullStream: mockFullStream(), + usage: mockUsage, + providerMetadata: mockProviderMetadata, + }) + + const stream = handler.createMessage(systemPrompt, messages) + const chunks: any[] = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + + expect(chunks.length).toBeGreaterThan(0) + const textChunks = chunks.filter((chunk) => chunk.type === "text") + expect(textChunks).toHaveLength(1) + expect(textChunks[0].text).toBe("Test response") + }) + + it("should include usage information", async () => { + async function* mockFullStream() { + yield { type: "text-delta", text: "Test response" } + } + + const mockUsage = Promise.resolve({ + inputTokens: 10, + outputTokens: 5, + }) + + const mockProviderMetadata = Promise.resolve({ + azure: { + promptCacheHitTokens: 2, + promptCacheMissTokens: 8, + }, + }) + + mockStreamText.mockReturnValue({ + fullStream: mockFullStream(), + usage: mockUsage, + providerMetadata: mockProviderMetadata, + }) + + const stream = handler.createMessage(systemPrompt, messages) + const chunks: any[] = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + + const usageChunks = chunks.filter((chunk) => chunk.type === "usage") + expect(usageChunks.length).toBeGreaterThan(0) + expect(usageChunks[0].inputTokens).toBe(10) + expect(usageChunks[0].outputTokens).toBe(5) + }) + + it("should include cache metrics in usage information from providerMetadata", async () => { + async function* mockFullStream() { + yield { type: "text-delta", text: "Test response" } + } + + const mockUsage = Promise.resolve({ + inputTokens: 10, + outputTokens: 5, + }) + + // Azure provides cache metrics via providerMetadata + const mockProviderMetadata = Promise.resolve({ + azure: { + promptCacheHitTokens: 2, + promptCacheMissTokens: 8, + }, + }) + + mockStreamText.mockReturnValue({ + fullStream: mockFullStream(), + usage: mockUsage, + providerMetadata: mockProviderMetadata, + }) + + const stream = handler.createMessage(systemPrompt, messages) + const chunks: any[] = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + + const usageChunks = chunks.filter((chunk) => chunk.type === "usage") + expect(usageChunks.length).toBeGreaterThan(0) + expect(usageChunks[0].cacheWriteTokens).toBe(8) // promptCacheMissTokens + expect(usageChunks[0].cacheReadTokens).toBe(2) // promptCacheHitTokens + }) + + it("should handle tool calls via tool-input-start/delta/end events", async () => { + async function* mockFullStream() { + yield { type: "tool-input-start", id: "tool-1", toolName: "test_tool" } + yield { type: "tool-input-delta", id: "tool-1", delta: '{"arg":' } + yield { type: "tool-input-delta", id: "tool-1", delta: '"value"}' } + yield { type: "tool-input-end", id: "tool-1" } + } + + const mockUsage = Promise.resolve({ + inputTokens: 10, + outputTokens: 5, + }) + + const mockProviderMetadata = Promise.resolve({}) + + mockStreamText.mockReturnValue({ + fullStream: mockFullStream(), + usage: mockUsage, + providerMetadata: mockProviderMetadata, + }) + + const stream = handler.createMessage(systemPrompt, messages) + const chunks: any[] = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + + const toolStartChunks = chunks.filter((chunk) => chunk.type === "tool_call_start") + expect(toolStartChunks).toHaveLength(1) + expect(toolStartChunks[0].id).toBe("tool-1") + expect(toolStartChunks[0].name).toBe("test_tool") + + const toolDeltaChunks = chunks.filter((chunk) => chunk.type === "tool_call_delta") + expect(toolDeltaChunks).toHaveLength(2) + + const toolEndChunks = chunks.filter((chunk) => chunk.type === "tool_call_end") + expect(toolEndChunks).toHaveLength(1) + }) + + it("should handle errors from AI SDK", async () => { + const mockError = new Error("API Error") + ;(mockError as any).name = "AI_APICallError" + ;(mockError as any).status = 500 + + async function* mockFullStream(): AsyncGenerator { + yield { type: "text-delta", text: "" } + throw mockError + } + + mockStreamText.mockReturnValue({ + fullStream: mockFullStream(), + usage: Promise.resolve({}), + providerMetadata: Promise.resolve({}), + }) + + const stream = handler.createMessage(systemPrompt, messages) + await expect(async () => { + const chunks: any[] = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + }).rejects.toThrow("Azure OpenAI") + }) + }) + + describe("completePrompt", () => { + it("should complete a prompt using generateText", async () => { + mockGenerateText.mockResolvedValue({ + text: "Test completion", + }) + + const result = await handler.completePrompt("Test prompt") + + expect(result).toBe("Test completion") + expect(mockGenerateText).toHaveBeenCalledWith( + expect.objectContaining({ + prompt: "Test prompt", + }), + ) + }) + + it("should use configured temperature", async () => { + const handlerWithTemp = new AzureHandler({ + ...mockOptions, + modelTemperature: 0.7, + }) + + mockGenerateText.mockResolvedValue({ + text: "Test completion", + }) + + await handlerWithTemp.completePrompt("Test prompt") + + expect(mockGenerateText).toHaveBeenCalledWith( + expect.objectContaining({ + temperature: 0.7, + }), + ) + }) + }) + + describe("tools", () => { + const systemPrompt = "You are a helpful assistant." + const messages: Anthropic.Messages.MessageParam[] = [ + { + role: "user", + content: [{ type: "text" as const, text: "Use a tool" }], + }, + ] + + it("should pass tools to streamText", async () => { + async function* mockFullStream() { + yield { type: "text-delta", text: "Using tool" } + } + + mockStreamText.mockReturnValue({ + fullStream: mockFullStream(), + usage: Promise.resolve({ inputTokens: 10, outputTokens: 5 }), + providerMetadata: Promise.resolve({}), + }) + + const tools = [ + { + type: "function" as const, + function: { + name: "test_tool", + description: "A test tool", + parameters: { + type: "object", + properties: { + arg: { type: "string" }, + }, + required: ["arg"], + }, + }, + }, + ] + + const stream = handler.createMessage(systemPrompt, messages, { + taskId: "test-task", + tools, + }) + + const chunks: any[] = [] + for await (const chunk of stream) { + chunks.push(chunk) + } + + expect(mockStreamText).toHaveBeenCalledWith( + expect.objectContaining({ + tools: expect.any(Object), + }), + ) + }) + }) +}) diff --git a/src/api/providers/azure.ts b/src/api/providers/azure.ts new file mode 100644 index 0000000000..ac8da0d941 --- /dev/null +++ b/src/api/providers/azure.ts @@ -0,0 +1,180 @@ +import { Anthropic } from "@anthropic-ai/sdk" +import { createAzure } from "@ai-sdk/azure" +import { streamText, generateText, ToolSet } from "ai" + +import { azureOpenAiDefaultApiVersion, openAiModelInfoSaneDefaults, type ModelInfo } from "@roo-code/types" + +import type { ApiHandlerOptions } from "../../shared/api" + +import { + convertToAiSdkMessages, + convertToolsForAiSdk, + processAiSdkStreamPart, + mapToolChoice, + handleAiSdkError, +} from "../transform/ai-sdk" +import { ApiStream, ApiStreamUsageChunk } from "../transform/stream" +import { getModelParams } from "../transform/model-params" + +import { DEFAULT_HEADERS } from "./constants" +import { BaseProvider } from "./base-provider" +import type { SingleCompletionHandler, ApiHandlerCreateMessageMetadata } from "../index" + +const AZURE_DEFAULT_TEMPERATURE = 0 + +/** + * Azure OpenAI provider using the dedicated @ai-sdk/azure package. + * Provides native support for Azure OpenAI deployments with proper resource-based routing. + */ +export class AzureHandler extends BaseProvider implements SingleCompletionHandler { + protected options: ApiHandlerOptions + protected provider: ReturnType + + constructor(options: ApiHandlerOptions) { + super() + this.options = options + + // Create the Azure provider using AI SDK + // The @ai-sdk/azure package uses resourceName-based routing + this.provider = createAzure({ + resourceName: options.azureResourceName ?? "", + apiKey: options.azureApiKey ?? "not-provided", + apiVersion: options.azureApiVersion ?? azureOpenAiDefaultApiVersion, + headers: DEFAULT_HEADERS, + }) + } + + override getModel(): { id: string; info: ModelInfo; maxTokens?: number; temperature?: number } { + // Azure uses deployment names as model IDs + // Use azureDeploymentName if provided, otherwise fall back to apiModelId + const id = this.options.azureDeploymentName ?? this.options.apiModelId ?? "" + const info: ModelInfo = openAiModelInfoSaneDefaults + const params = getModelParams({ + format: "openai", + modelId: id, + model: info, + settings: this.options, + defaultTemperature: AZURE_DEFAULT_TEMPERATURE, + }) + return { id, info, ...params } + } + + /** + * Get the language model for the configured deployment name. + */ + protected getLanguageModel() { + const { id } = this.getModel() + return this.provider(id) + } + + /** + * Process usage metrics from the AI SDK response. + * Azure OpenAI provides standard OpenAI-compatible usage metrics. + */ + protected processUsageMetrics( + usage: { + inputTokens?: number + outputTokens?: number + details?: { + cachedInputTokens?: number + reasoningTokens?: number + } + }, + providerMetadata?: { + azure?: { + promptCacheHitTokens?: number + promptCacheMissTokens?: number + } + }, + ): ApiStreamUsageChunk { + // Extract cache metrics from Azure's providerMetadata if available + const cacheReadTokens = providerMetadata?.azure?.promptCacheHitTokens ?? usage.details?.cachedInputTokens + const cacheWriteTokens = providerMetadata?.azure?.promptCacheMissTokens + + return { + type: "usage", + inputTokens: usage.inputTokens || 0, + outputTokens: usage.outputTokens || 0, + cacheReadTokens, + cacheWriteTokens, + reasoningTokens: usage.details?.reasoningTokens, + } + } + + /** + * Get the max tokens parameter to include in the request. + */ + protected getMaxOutputTokens(): number | undefined { + const { info } = this.getModel() + return this.options.modelMaxTokens || info.maxTokens || undefined + } + + /** + * Create a message stream using the AI SDK. + */ + override async *createMessage( + systemPrompt: string, + messages: Anthropic.Messages.MessageParam[], + metadata?: ApiHandlerCreateMessageMetadata, + ): ApiStream { + const { temperature } = this.getModel() + const languageModel = this.getLanguageModel() + + // Convert messages to AI SDK format + const aiSdkMessages = convertToAiSdkMessages(messages) + + // Convert tools to OpenAI format first, then to AI SDK format + const openAiTools = this.convertToolsForOpenAI(metadata?.tools) + const aiSdkTools = convertToolsForAiSdk(openAiTools) as ToolSet | undefined + + // Build the request options + const requestOptions: Parameters[0] = { + model: languageModel, + system: systemPrompt, + messages: aiSdkMessages, + temperature: this.options.modelTemperature ?? temperature ?? AZURE_DEFAULT_TEMPERATURE, + maxOutputTokens: this.getMaxOutputTokens(), + tools: aiSdkTools, + toolChoice: mapToolChoice(metadata?.tool_choice), + } + + // Use streamText for streaming responses + const result = streamText(requestOptions) + + try { + // Process the full stream to get all events including reasoning + for await (const part of result.fullStream) { + for (const chunk of processAiSdkStreamPart(part)) { + yield chunk + } + } + + // Yield usage metrics at the end, including cache metrics from providerMetadata + const usage = await result.usage + const providerMetadata = await result.providerMetadata + if (usage) { + yield this.processUsageMetrics(usage, providerMetadata as any) + } + } catch (error) { + // Handle AI SDK errors (AI_RetryError, AI_APICallError, etc.) + throw handleAiSdkError(error, "Azure OpenAI") + } + } + + /** + * Complete a prompt using the AI SDK generateText. + */ + async completePrompt(prompt: string): Promise { + const { temperature } = this.getModel() + const languageModel = this.getLanguageModel() + + const { text } = await generateText({ + model: languageModel, + prompt, + maxOutputTokens: this.getMaxOutputTokens(), + temperature: this.options.modelTemperature ?? temperature ?? AZURE_DEFAULT_TEMPERATURE, + }) + + return text + } +} diff --git a/src/api/providers/index.ts b/src/api/providers/index.ts index cf49f75f18..3212e267cb 100644 --- a/src/api/providers/index.ts +++ b/src/api/providers/index.ts @@ -1,5 +1,6 @@ export { AnthropicVertexHandler } from "./anthropic-vertex" export { AnthropicHandler } from "./anthropic" +export { AzureHandler } from "./azure" export { AwsBedrockHandler } from "./bedrock" export { CerebrasHandler } from "./cerebras" export { ChutesHandler } from "./chutes" diff --git a/src/package.json b/src/package.json index eb67c0d1d7..89bacc099c 100644 --- a/src/package.json +++ b/src/package.json @@ -451,6 +451,7 @@ }, "dependencies": { "@ai-sdk/amazon-bedrock": "^4.0.50", + "@ai-sdk/azure": "^2.0.6", "@ai-sdk/baseten": "^1.0.31", "@ai-sdk/cerebras": "^1.0.0", "@ai-sdk/deepseek": "^2.0.14", diff --git a/webview-ui/src/components/ui/hooks/useSelectedModel.ts b/webview-ui/src/components/ui/hooks/useSelectedModel.ts index 5336b63583..8f39bdbac2 100644 --- a/webview-ui/src/components/ui/hooks/useSelectedModel.ts +++ b/webview-ui/src/components/ui/hooks/useSelectedModel.ts @@ -387,6 +387,11 @@ function getSelectedModel({ const info = routerModels["vercel-ai-gateway"]?.[id] return { id, info } } + case "azure": { + // Azure uses deployment names configured by the user + const id = apiConfiguration.azureDeploymentName ?? apiConfiguration.apiModelId ?? "" + return { id, info: undefined } + } // case "anthropic": // case "fake-ai": default: {