mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-10-08 03:07:53 +00:00
feat: add dedicated Azure OpenAI provider using @ai-sdk/azure package
This commit is contained in:
parent
d7714e4e07
commit
d6dd55b1e3
8 changed files with 590 additions and 0 deletions
13
.changeset/azure-ai-sdk-migration.md
Normal file
13
.changeset/azure-ai-sdk-migration.md
Normal 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
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
371
src/api/providers/__tests__/azure.spec.ts
Normal file
371
src/api/providers/__tests__/azure.spec.ts
Normal 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
180
src/api/providers/azure.ts
Normal 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
|
||||
}
|
||||
}
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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: {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue