feat: add dedicated Azure OpenAI provider using @ai-sdk/azure package

This commit is contained in:
Roo Code 2026-02-01 18:08:30 +00:00 • committed by Hannes Rudolph
parent d7714e4e07
commit d6dd55b1e3
8 changed files with 590 additions and 0 deletions

View file

@ -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

View file

@ -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<TypicalProvider, ModelIdKey> = {
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",

View file

@ -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":

View file

@ -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<typeof import("ai")>()
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<any> {
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),
}),
)
})
})
})

180
src/api/providers/azure.ts Normal file
View file

@ -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<typeof createAzure>
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<typeof streamText>[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<string> {
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
}
}

View file

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

View file

@ -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",

View file

@ -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: {