mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-09-05 08:10:14 +00:00
fix: add AbortController support to OpenAI-compatible providers
Fixes #10779 - Stop button now properly cancels requests for OpenAI Compatible API providers (llama.cpp, LM Studio, etc.) Changes: - Add abortController property to OpenAiHandler and BaseOpenAiCompatibleProvider - Pass abort signal to all SDK calls in createMessage and completePrompt methods - Check for abort signal in stream loops to break early when cancelled - Update tests to expect abort signal in request options
This commit is contained in:
parent
9bf7173725
commit
5d606ba57a
4 changed files with 346 additions and 265 deletions
|
|
@ -354,7 +354,9 @@ describe("BaseOpenAiCompatibleProvider", () => {
|
|||
stream: true,
|
||||
stream_options: { include_usage: true },
|
||||
}),
|
||||
undefined,
|
||||
expect.objectContaining({
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
|
|
|
|||
|
|
@ -549,7 +549,9 @@ describe("OpenAiHandler", () => {
|
|||
model: mockOptions.openAiModelId,
|
||||
messages: [{ role: "user", content: "Test prompt" }],
|
||||
},
|
||||
{},
|
||||
expect.objectContaining({
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
|
|
@ -634,7 +636,10 @@ describe("OpenAiHandler", () => {
|
|||
stream_options: { include_usage: true },
|
||||
temperature: 0,
|
||||
},
|
||||
{ path: "/models/chat/completions" },
|
||||
expect.objectContaining({
|
||||
path: "/models/chat/completions",
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
|
||||
// Verify max_tokens is NOT included when includeMaxTokens is not set
|
||||
|
|
@ -680,7 +685,10 @@ describe("OpenAiHandler", () => {
|
|||
{ role: "user", content: "Hello!" },
|
||||
],
|
||||
},
|
||||
{ path: "/models/chat/completions" },
|
||||
expect.objectContaining({
|
||||
path: "/models/chat/completions",
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
|
||||
// Verify max_tokens is NOT included when includeMaxTokens is not set
|
||||
|
|
@ -697,7 +705,10 @@ describe("OpenAiHandler", () => {
|
|||
model: azureOptions.openAiModelId,
|
||||
messages: [{ role: "user", content: "Test prompt" }],
|
||||
},
|
||||
{ path: "/models/chat/completions" },
|
||||
expect.objectContaining({
|
||||
path: "/models/chat/completions",
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
|
||||
// Verify max_tokens is NOT included when includeMaxTokens is not set
|
||||
|
|
@ -737,7 +748,9 @@ describe("OpenAiHandler", () => {
|
|||
model: grokOptions.openAiModelId,
|
||||
stream: true,
|
||||
}),
|
||||
{},
|
||||
expect.objectContaining({
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
|
||||
const mockCalls = mockCreate.mock.calls
|
||||
|
|
@ -796,7 +809,9 @@ describe("OpenAiHandler", () => {
|
|||
// O3 models do not support deprecated max_tokens but do support max_completion_tokens
|
||||
max_completion_tokens: 32000,
|
||||
}),
|
||||
{},
|
||||
expect.objectContaining({
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
|
|
@ -953,7 +968,9 @@ describe("OpenAiHandler", () => {
|
|||
reasoning_effort: "medium",
|
||||
temperature: undefined,
|
||||
}),
|
||||
{},
|
||||
expect.objectContaining({
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
|
||||
// Verify max_tokens is NOT included
|
||||
|
|
@ -997,7 +1014,9 @@ describe("OpenAiHandler", () => {
|
|||
// O3 models do not support deprecated max_tokens but do support max_completion_tokens
|
||||
max_completion_tokens: 65536, // Using default maxTokens from o3Options
|
||||
}),
|
||||
{},
|
||||
expect.objectContaining({
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
|
||||
// Verify stream is not set
|
||||
|
|
@ -1074,7 +1093,9 @@ describe("OpenAiHandler", () => {
|
|||
expect.objectContaining({
|
||||
temperature: undefined, // Temperature is not supported for O3 models
|
||||
}),
|
||||
{},
|
||||
expect.objectContaining({
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
})
|
||||
|
||||
|
|
@ -1099,7 +1120,10 @@ describe("OpenAiHandler", () => {
|
|||
expect.objectContaining({
|
||||
model: "o3-mini",
|
||||
}),
|
||||
{ path: "/models/chat/completions" },
|
||||
expect.objectContaining({
|
||||
path: "/models/chat/completions",
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
|
||||
// Verify max_tokens is NOT included when includeMaxTokens is false
|
||||
|
|
@ -1129,7 +1153,10 @@ describe("OpenAiHandler", () => {
|
|||
model: "o3-mini",
|
||||
// O3 models do not support max_tokens
|
||||
}),
|
||||
{ path: "/models/chat/completions" },
|
||||
expect.objectContaining({
|
||||
path: "/models/chat/completions",
|
||||
signal: expect.any(AbortSignal),
|
||||
}),
|
||||
)
|
||||
})
|
||||
})
|
||||
|
|
|
|||
|
|
@ -37,6 +37,9 @@ export abstract class BaseOpenAiCompatibleProvider<ModelName extends string>
|
|||
|
||||
protected client: OpenAI
|
||||
|
||||
// Abort controller for cancelling ongoing requests
|
||||
private abortController?: AbortController
|
||||
|
||||
constructor({
|
||||
providerName,
|
||||
baseURL,
|
||||
|
|
@ -106,7 +109,12 @@ export abstract class BaseOpenAiCompatibleProvider<ModelName extends string>
|
|||
}
|
||||
|
||||
try {
|
||||
return this.client.chat.completions.create(params, requestOptions)
|
||||
// Merge abort signal with any existing request options
|
||||
const mergedOptions: OpenAI.RequestOptions = {
|
||||
...requestOptions,
|
||||
signal: this.abortController?.signal,
|
||||
}
|
||||
return this.client.chat.completions.create(params, mergedOptions)
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
|
|
@ -117,87 +125,99 @@ export abstract class BaseOpenAiCompatibleProvider<ModelName extends string>
|
|||
messages: Anthropic.Messages.MessageParam[],
|
||||
metadata?: ApiHandlerCreateMessageMetadata,
|
||||
): ApiStream {
|
||||
const stream = await this.createStream(systemPrompt, messages, metadata)
|
||||
// Create AbortController for cancellation
|
||||
this.abortController = new AbortController()
|
||||
|
||||
const matcher = new XmlMatcher(
|
||||
"think",
|
||||
(chunk) =>
|
||||
({
|
||||
type: chunk.matched ? "reasoning" : "text",
|
||||
text: chunk.data,
|
||||
}) as const,
|
||||
)
|
||||
try {
|
||||
const stream = await this.createStream(systemPrompt, messages, metadata)
|
||||
|
||||
let lastUsage: OpenAI.CompletionUsage | undefined
|
||||
const activeToolCallIds = new Set<string>()
|
||||
const matcher = new XmlMatcher(
|
||||
"think",
|
||||
(chunk) =>
|
||||
({
|
||||
type: chunk.matched ? "reasoning" : "text",
|
||||
text: chunk.data,
|
||||
}) as const,
|
||||
)
|
||||
|
||||
for await (const chunk of stream) {
|
||||
// Check for provider-specific error responses (e.g., MiniMax base_resp)
|
||||
const chunkAny = chunk as any
|
||||
if (chunkAny.base_resp?.status_code && chunkAny.base_resp.status_code !== 0) {
|
||||
throw new Error(
|
||||
`${this.providerName} API Error (${chunkAny.base_resp.status_code}): ${chunkAny.base_resp.status_msg || "Unknown error"}`,
|
||||
)
|
||||
}
|
||||
let lastUsage: OpenAI.CompletionUsage | undefined
|
||||
const activeToolCallIds = new Set<string>()
|
||||
|
||||
const delta = chunk.choices?.[0]?.delta
|
||||
const finishReason = chunk.choices?.[0]?.finish_reason
|
||||
|
||||
if (delta?.content) {
|
||||
for (const processedChunk of matcher.update(delta.content)) {
|
||||
yield processedChunk
|
||||
for await (const chunk of stream) {
|
||||
// Check if request was aborted
|
||||
if (this.abortController?.signal.aborted) {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if (delta) {
|
||||
for (const key of ["reasoning_content", "reasoning"] as const) {
|
||||
if (key in delta) {
|
||||
const reasoning_content = ((delta as any)[key] as string | undefined) || ""
|
||||
if (reasoning_content?.trim()) {
|
||||
yield { type: "reasoning", text: reasoning_content }
|
||||
// Check for provider-specific error responses (e.g., MiniMax base_resp)
|
||||
const chunkAny = chunk as any
|
||||
if (chunkAny.base_resp?.status_code && chunkAny.base_resp.status_code !== 0) {
|
||||
throw new Error(
|
||||
`${this.providerName} API Error (${chunkAny.base_resp.status_code}): ${chunkAny.base_resp.status_msg || "Unknown error"}`,
|
||||
)
|
||||
}
|
||||
|
||||
const delta = chunk.choices?.[0]?.delta
|
||||
const finishReason = chunk.choices?.[0]?.finish_reason
|
||||
|
||||
if (delta?.content) {
|
||||
for (const processedChunk of matcher.update(delta.content)) {
|
||||
yield processedChunk
|
||||
}
|
||||
}
|
||||
|
||||
if (delta) {
|
||||
for (const key of ["reasoning_content", "reasoning"] as const) {
|
||||
if (key in delta) {
|
||||
const reasoning_content = ((delta as any)[key] as string | undefined) || ""
|
||||
if (reasoning_content?.trim()) {
|
||||
yield { type: "reasoning", text: reasoning_content }
|
||||
}
|
||||
break
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// Emit raw tool call chunks - NativeToolCallParser handles state management
|
||||
if (delta?.tool_calls) {
|
||||
for (const toolCall of delta.tool_calls) {
|
||||
if (toolCall.id) {
|
||||
activeToolCallIds.add(toolCall.id)
|
||||
}
|
||||
yield {
|
||||
type: "tool_call_partial",
|
||||
index: toolCall.index,
|
||||
id: toolCall.id,
|
||||
name: toolCall.function?.name,
|
||||
arguments: toolCall.function?.arguments,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Emit tool_call_end events when finish_reason is "tool_calls"
|
||||
// This ensures tool calls are finalized even if the stream doesn't properly close
|
||||
if (finishReason === "tool_calls" && activeToolCallIds.size > 0) {
|
||||
for (const id of activeToolCallIds) {
|
||||
yield { type: "tool_call_end", id }
|
||||
}
|
||||
activeToolCallIds.clear()
|
||||
}
|
||||
|
||||
if (chunk.usage) {
|
||||
lastUsage = chunk.usage
|
||||
}
|
||||
}
|
||||
|
||||
// Emit raw tool call chunks - NativeToolCallParser handles state management
|
||||
if (delta?.tool_calls) {
|
||||
for (const toolCall of delta.tool_calls) {
|
||||
if (toolCall.id) {
|
||||
activeToolCallIds.add(toolCall.id)
|
||||
}
|
||||
yield {
|
||||
type: "tool_call_partial",
|
||||
index: toolCall.index,
|
||||
id: toolCall.id,
|
||||
name: toolCall.function?.name,
|
||||
arguments: toolCall.function?.arguments,
|
||||
}
|
||||
}
|
||||
if (lastUsage) {
|
||||
yield this.processUsageMetrics(lastUsage, this.getModel().info)
|
||||
}
|
||||
|
||||
// Emit tool_call_end events when finish_reason is "tool_calls"
|
||||
// This ensures tool calls are finalized even if the stream doesn't properly close
|
||||
if (finishReason === "tool_calls" && activeToolCallIds.size > 0) {
|
||||
for (const id of activeToolCallIds) {
|
||||
yield { type: "tool_call_end", id }
|
||||
}
|
||||
activeToolCallIds.clear()
|
||||
// Process any remaining content
|
||||
for (const processedChunk of matcher.final()) {
|
||||
yield processedChunk
|
||||
}
|
||||
|
||||
if (chunk.usage) {
|
||||
lastUsage = chunk.usage
|
||||
}
|
||||
}
|
||||
|
||||
if (lastUsage) {
|
||||
yield this.processUsageMetrics(lastUsage, this.getModel().info)
|
||||
}
|
||||
|
||||
// Process any remaining content
|
||||
for (const processedChunk of matcher.final()) {
|
||||
yield processedChunk
|
||||
} finally {
|
||||
this.abortController = undefined
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -222,20 +242,25 @@ export abstract class BaseOpenAiCompatibleProvider<ModelName extends string>
|
|||
}
|
||||
|
||||
async completePrompt(prompt: string): Promise<string> {
|
||||
const { id: modelId, info: modelInfo } = this.getModel()
|
||||
|
||||
const params: OpenAI.Chat.Completions.ChatCompletionCreateParams = {
|
||||
model: modelId,
|
||||
messages: [{ role: "user", content: prompt }],
|
||||
}
|
||||
|
||||
// Add thinking parameter if reasoning is enabled and model supports it
|
||||
if (this.options.enableReasoningEffort && modelInfo.supportsReasoningBinary) {
|
||||
;(params as any).thinking = { type: "enabled" }
|
||||
}
|
||||
// Create AbortController for cancellation
|
||||
this.abortController = new AbortController()
|
||||
|
||||
try {
|
||||
const response = await this.client.chat.completions.create(params)
|
||||
const { id: modelId, info: modelInfo } = this.getModel()
|
||||
|
||||
const params: OpenAI.Chat.Completions.ChatCompletionCreateParams = {
|
||||
model: modelId,
|
||||
messages: [{ role: "user", content: prompt }],
|
||||
}
|
||||
|
||||
// Add thinking parameter if reasoning is enabled and model supports it
|
||||
if (this.options.enableReasoningEffort && modelInfo.supportsReasoningBinary) {
|
||||
;(params as any).thinking = { type: "enabled" }
|
||||
}
|
||||
|
||||
const response = await this.client.chat.completions.create(params, {
|
||||
signal: this.abortController.signal,
|
||||
})
|
||||
|
||||
// Check for provider-specific error responses (e.g., MiniMax base_resp)
|
||||
const responseAny = response as any
|
||||
|
|
@ -248,6 +273,8 @@ export abstract class BaseOpenAiCompatibleProvider<ModelName extends string>
|
|||
return response.choices?.[0]?.message.content || ""
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
} finally {
|
||||
this.abortController = undefined
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -33,6 +33,8 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
protected options: ApiHandlerOptions
|
||||
protected client: OpenAI
|
||||
private readonly providerName = "OpenAI"
|
||||
// Abort controller for cancelling ongoing requests
|
||||
private abortController?: AbortController
|
||||
|
||||
constructor(options: ApiHandlerOptions) {
|
||||
super()
|
||||
|
|
@ -85,193 +87,206 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
messages: Anthropic.Messages.MessageParam[],
|
||||
metadata?: ApiHandlerCreateMessageMetadata,
|
||||
): ApiStream {
|
||||
const { info: modelInfo, reasoning } = this.getModel()
|
||||
const modelUrl = this.options.openAiBaseUrl ?? ""
|
||||
const modelId = this.options.openAiModelId ?? ""
|
||||
const enabledR1Format = this.options.openAiR1FormatEnabled ?? false
|
||||
const isAzureAiInference = this._isAzureAiInference(modelUrl)
|
||||
const deepseekReasoner = modelId.includes("deepseek-reasoner") || enabledR1Format
|
||||
// Create AbortController for cancellation
|
||||
this.abortController = new AbortController()
|
||||
|
||||
if (modelId.includes("o1") || modelId.includes("o3") || modelId.includes("o4")) {
|
||||
yield* this.handleO3FamilyMessage(modelId, systemPrompt, messages, metadata)
|
||||
return
|
||||
}
|
||||
try {
|
||||
const { info: modelInfo, reasoning } = this.getModel()
|
||||
const modelUrl = this.options.openAiBaseUrl ?? ""
|
||||
const modelId = this.options.openAiModelId ?? ""
|
||||
const enabledR1Format = this.options.openAiR1FormatEnabled ?? false
|
||||
const isAzureAiInference = this._isAzureAiInference(modelUrl)
|
||||
const deepseekReasoner = modelId.includes("deepseek-reasoner") || enabledR1Format
|
||||
|
||||
let systemMessage: OpenAI.Chat.ChatCompletionSystemMessageParam = {
|
||||
role: "system",
|
||||
content: systemPrompt,
|
||||
}
|
||||
if (modelId.includes("o1") || modelId.includes("o3") || modelId.includes("o4")) {
|
||||
yield* this.handleO3FamilyMessage(modelId, systemPrompt, messages, metadata)
|
||||
return
|
||||
}
|
||||
|
||||
if (this.options.openAiStreamingEnabled ?? true) {
|
||||
let convertedMessages
|
||||
let systemMessage: OpenAI.Chat.ChatCompletionSystemMessageParam = {
|
||||
role: "system",
|
||||
content: systemPrompt,
|
||||
}
|
||||
|
||||
if (deepseekReasoner) {
|
||||
convertedMessages = convertToR1Format([{ role: "user", content: systemPrompt }, ...messages])
|
||||
} else {
|
||||
if (modelInfo.supportsPromptCache) {
|
||||
systemMessage = {
|
||||
role: "system",
|
||||
content: [
|
||||
{
|
||||
type: "text",
|
||||
text: systemPrompt,
|
||||
// @ts-ignore-next-line
|
||||
cache_control: { type: "ephemeral" },
|
||||
},
|
||||
],
|
||||
}
|
||||
}
|
||||
if (this.options.openAiStreamingEnabled ?? true) {
|
||||
let convertedMessages
|
||||
|
||||
convertedMessages = [systemMessage, ...convertToOpenAiMessages(messages)]
|
||||
|
||||
if (modelInfo.supportsPromptCache) {
|
||||
// Note: the following logic is copied from openrouter:
|
||||
// Add cache_control to the last two user messages
|
||||
// (note: this works because we only ever add one user message at a time, but if we added multiple we'd need to mark the user message before the last assistant message)
|
||||
const lastTwoUserMessages = convertedMessages.filter((msg) => msg.role === "user").slice(-2)
|
||||
|
||||
lastTwoUserMessages.forEach((msg) => {
|
||||
if (typeof msg.content === "string") {
|
||||
msg.content = [{ type: "text", text: msg.content }]
|
||||
if (deepseekReasoner) {
|
||||
convertedMessages = convertToR1Format([{ role: "user", content: systemPrompt }, ...messages])
|
||||
} else {
|
||||
if (modelInfo.supportsPromptCache) {
|
||||
systemMessage = {
|
||||
role: "system",
|
||||
content: [
|
||||
{
|
||||
type: "text",
|
||||
text: systemPrompt,
|
||||
// @ts-ignore-next-line
|
||||
cache_control: { type: "ephemeral" },
|
||||
},
|
||||
],
|
||||
}
|
||||
}
|
||||
|
||||
if (Array.isArray(msg.content)) {
|
||||
// NOTE: this is fine since env details will always be added at the end. but if it weren't there, and the user added a image_url type message, it would pop a text part before it and then move it after to the end.
|
||||
let lastTextPart = msg.content.filter((part) => part.type === "text").pop()
|
||||
convertedMessages = [systemMessage, ...convertToOpenAiMessages(messages)]
|
||||
|
||||
if (!lastTextPart) {
|
||||
lastTextPart = { type: "text", text: "..." }
|
||||
msg.content.push(lastTextPart)
|
||||
if (modelInfo.supportsPromptCache) {
|
||||
// Note: the following logic is copied from openrouter:
|
||||
// Add cache_control to the last two user messages
|
||||
// (note: this works because we only ever add one user message at a time, but if we added multiple we'd need to mark the user message before the last assistant message)
|
||||
const lastTwoUserMessages = convertedMessages.filter((msg) => msg.role === "user").slice(-2)
|
||||
|
||||
lastTwoUserMessages.forEach((msg) => {
|
||||
if (typeof msg.content === "string") {
|
||||
msg.content = [{ type: "text", text: msg.content }]
|
||||
}
|
||||
|
||||
// @ts-ignore-next-line
|
||||
lastTextPart["cache_control"] = { type: "ephemeral" }
|
||||
}
|
||||
if (Array.isArray(msg.content)) {
|
||||
// NOTE: this is fine since env details will always be added at the end. but if it weren't there, and the user added a image_url type message, it would pop a text part before it and then move it after to the end.
|
||||
let lastTextPart = msg.content.filter((part) => part.type === "text").pop()
|
||||
|
||||
if (!lastTextPart) {
|
||||
lastTextPart = { type: "text", text: "..." }
|
||||
msg.content.push(lastTextPart)
|
||||
}
|
||||
|
||||
// @ts-ignore-next-line
|
||||
lastTextPart["cache_control"] = { type: "ephemeral" }
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
const isGrokXAI = this._isGrokXAI(this.options.openAiBaseUrl)
|
||||
|
||||
const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = {
|
||||
model: modelId,
|
||||
temperature:
|
||||
this.options.modelTemperature ?? (deepseekReasoner ? DEEP_SEEK_DEFAULT_TEMPERATURE : 0),
|
||||
messages: convertedMessages,
|
||||
stream: true as const,
|
||||
...(isGrokXAI ? {} : { stream_options: { include_usage: true } }),
|
||||
...(reasoning && reasoning),
|
||||
...(metadata?.tools && { tools: this.convertToolsForOpenAI(metadata.tools) }),
|
||||
...(metadata?.tool_choice && { tool_choice: metadata.tool_choice }),
|
||||
...(metadata?.toolProtocol === "native" &&
|
||||
metadata.parallelToolCalls === true && {
|
||||
parallel_tool_calls: true,
|
||||
}),
|
||||
}
|
||||
|
||||
// Add max_tokens if needed
|
||||
this.addMaxTokensIfNeeded(requestOptions, modelInfo)
|
||||
|
||||
let stream
|
||||
try {
|
||||
stream = await this.client.chat.completions.create(requestOptions, {
|
||||
...(isAzureAiInference ? { path: OPENAI_AZURE_AI_INFERENCE_PATH } : {}),
|
||||
signal: this.abortController.signal,
|
||||
})
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
}
|
||||
|
||||
const isGrokXAI = this._isGrokXAI(this.options.openAiBaseUrl)
|
||||
|
||||
const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = {
|
||||
model: modelId,
|
||||
temperature: this.options.modelTemperature ?? (deepseekReasoner ? DEEP_SEEK_DEFAULT_TEMPERATURE : 0),
|
||||
messages: convertedMessages,
|
||||
stream: true as const,
|
||||
...(isGrokXAI ? {} : { stream_options: { include_usage: true } }),
|
||||
...(reasoning && reasoning),
|
||||
...(metadata?.tools && { tools: this.convertToolsForOpenAI(metadata.tools) }),
|
||||
...(metadata?.tool_choice && { tool_choice: metadata.tool_choice }),
|
||||
...(metadata?.toolProtocol === "native" &&
|
||||
metadata.parallelToolCalls === true && {
|
||||
parallel_tool_calls: true,
|
||||
}),
|
||||
}
|
||||
|
||||
// Add max_tokens if needed
|
||||
this.addMaxTokensIfNeeded(requestOptions, modelInfo)
|
||||
|
||||
let stream
|
||||
try {
|
||||
stream = await this.client.chat.completions.create(
|
||||
requestOptions,
|
||||
isAzureAiInference ? { path: OPENAI_AZURE_AI_INFERENCE_PATH } : {},
|
||||
const matcher = new XmlMatcher(
|
||||
"think",
|
||||
(chunk) =>
|
||||
({
|
||||
type: chunk.matched ? "reasoning" : "text",
|
||||
text: chunk.data,
|
||||
}) as const,
|
||||
)
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
|
||||
const matcher = new XmlMatcher(
|
||||
"think",
|
||||
(chunk) =>
|
||||
({
|
||||
type: chunk.matched ? "reasoning" : "text",
|
||||
text: chunk.data,
|
||||
}) as const,
|
||||
)
|
||||
let lastUsage
|
||||
const activeToolCallIds = new Set<string>()
|
||||
|
||||
let lastUsage
|
||||
const activeToolCallIds = new Set<string>()
|
||||
|
||||
for await (const chunk of stream) {
|
||||
const delta = chunk.choices?.[0]?.delta ?? {}
|
||||
const finishReason = chunk.choices?.[0]?.finish_reason
|
||||
|
||||
if (delta.content) {
|
||||
for (const chunk of matcher.update(delta.content)) {
|
||||
yield chunk
|
||||
for await (const chunk of stream) {
|
||||
// Check if request was aborted
|
||||
if (this.abortController?.signal.aborted) {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if ("reasoning_content" in delta && delta.reasoning_content) {
|
||||
yield {
|
||||
type: "reasoning",
|
||||
text: (delta.reasoning_content as string | undefined) || "",
|
||||
const delta = chunk.choices?.[0]?.delta ?? {}
|
||||
const finishReason = chunk.choices?.[0]?.finish_reason
|
||||
|
||||
if (delta.content) {
|
||||
for (const chunk of matcher.update(delta.content)) {
|
||||
yield chunk
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
yield* this.processToolCalls(delta, finishReason, activeToolCallIds)
|
||||
|
||||
if (chunk.usage) {
|
||||
lastUsage = chunk.usage
|
||||
}
|
||||
}
|
||||
|
||||
for (const chunk of matcher.final()) {
|
||||
yield chunk
|
||||
}
|
||||
|
||||
if (lastUsage) {
|
||||
yield this.processUsageMetrics(lastUsage, modelInfo)
|
||||
}
|
||||
} else {
|
||||
const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = {
|
||||
model: modelId,
|
||||
messages: deepseekReasoner
|
||||
? convertToR1Format([{ role: "user", content: systemPrompt }, ...messages])
|
||||
: [systemMessage, ...convertToOpenAiMessages(messages)],
|
||||
...(metadata?.tools && { tools: this.convertToolsForOpenAI(metadata.tools) }),
|
||||
...(metadata?.tool_choice && { tool_choice: metadata.tool_choice }),
|
||||
...(metadata?.toolProtocol === "native" &&
|
||||
metadata.parallelToolCalls === true && {
|
||||
parallel_tool_calls: true,
|
||||
}),
|
||||
}
|
||||
|
||||
// Add max_tokens if needed
|
||||
this.addMaxTokensIfNeeded(requestOptions, modelInfo)
|
||||
|
||||
let response
|
||||
try {
|
||||
response = await this.client.chat.completions.create(
|
||||
requestOptions,
|
||||
this._isAzureAiInference(modelUrl) ? { path: OPENAI_AZURE_AI_INFERENCE_PATH } : {},
|
||||
)
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
|
||||
const message = response.choices?.[0]?.message
|
||||
|
||||
if (message?.tool_calls) {
|
||||
for (const toolCall of message.tool_calls) {
|
||||
if (toolCall.type === "function") {
|
||||
if ("reasoning_content" in delta && delta.reasoning_content) {
|
||||
yield {
|
||||
type: "tool_call",
|
||||
id: toolCall.id,
|
||||
name: toolCall.function.name,
|
||||
arguments: toolCall.function.arguments,
|
||||
type: "reasoning",
|
||||
text: (delta.reasoning_content as string | undefined) || "",
|
||||
}
|
||||
}
|
||||
|
||||
yield* this.processToolCalls(delta, finishReason, activeToolCallIds)
|
||||
|
||||
if (chunk.usage) {
|
||||
lastUsage = chunk.usage
|
||||
}
|
||||
}
|
||||
|
||||
for (const chunk of matcher.final()) {
|
||||
yield chunk
|
||||
}
|
||||
|
||||
if (lastUsage) {
|
||||
yield this.processUsageMetrics(lastUsage, modelInfo)
|
||||
}
|
||||
} else {
|
||||
const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = {
|
||||
model: modelId,
|
||||
messages: deepseekReasoner
|
||||
? convertToR1Format([{ role: "user", content: systemPrompt }, ...messages])
|
||||
: [systemMessage, ...convertToOpenAiMessages(messages)],
|
||||
...(metadata?.tools && { tools: this.convertToolsForOpenAI(metadata.tools) }),
|
||||
...(metadata?.tool_choice && { tool_choice: metadata.tool_choice }),
|
||||
...(metadata?.toolProtocol === "native" &&
|
||||
metadata.parallelToolCalls === true && {
|
||||
parallel_tool_calls: true,
|
||||
}),
|
||||
}
|
||||
|
||||
// Add max_tokens if needed
|
||||
this.addMaxTokensIfNeeded(requestOptions, modelInfo)
|
||||
|
||||
let response
|
||||
try {
|
||||
response = await this.client.chat.completions.create(requestOptions, {
|
||||
...(this._isAzureAiInference(modelUrl) ? { path: OPENAI_AZURE_AI_INFERENCE_PATH } : {}),
|
||||
signal: this.abortController.signal,
|
||||
})
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
|
||||
const message = response.choices?.[0]?.message
|
||||
|
||||
if (message?.tool_calls) {
|
||||
for (const toolCall of message.tool_calls) {
|
||||
if (toolCall.type === "function") {
|
||||
yield {
|
||||
type: "tool_call",
|
||||
id: toolCall.id,
|
||||
name: toolCall.function.name,
|
||||
arguments: toolCall.function.arguments,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
yield {
|
||||
type: "text",
|
||||
text: message?.content || "",
|
||||
}
|
||||
yield {
|
||||
type: "text",
|
||||
text: message?.content || "",
|
||||
}
|
||||
|
||||
yield this.processUsageMetrics(response.usage, modelInfo)
|
||||
yield this.processUsageMetrics(response.usage, modelInfo)
|
||||
}
|
||||
} finally {
|
||||
this.abortController = undefined
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -299,6 +314,9 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
}
|
||||
|
||||
async completePrompt(prompt: string): Promise<string> {
|
||||
// Create AbortController for cancellation
|
||||
this.abortController = new AbortController()
|
||||
|
||||
try {
|
||||
const isAzureAiInference = this._isAzureAiInference(this.options.openAiBaseUrl)
|
||||
const model = this.getModel()
|
||||
|
|
@ -314,10 +332,10 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
|
||||
let response
|
||||
try {
|
||||
response = await this.client.chat.completions.create(
|
||||
requestOptions,
|
||||
isAzureAiInference ? { path: OPENAI_AZURE_AI_INFERENCE_PATH } : {},
|
||||
)
|
||||
response = await this.client.chat.completions.create(requestOptions, {
|
||||
...(isAzureAiInference ? { path: OPENAI_AZURE_AI_INFERENCE_PATH } : {}),
|
||||
signal: this.abortController.signal,
|
||||
})
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
|
|
@ -329,6 +347,8 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
}
|
||||
|
||||
throw error
|
||||
} finally {
|
||||
this.abortController = undefined
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -372,10 +392,10 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
|
||||
let stream
|
||||
try {
|
||||
stream = await this.client.chat.completions.create(
|
||||
requestOptions,
|
||||
methodIsAzureAiInference ? { path: OPENAI_AZURE_AI_INFERENCE_PATH } : {},
|
||||
)
|
||||
stream = await this.client.chat.completions.create(requestOptions, {
|
||||
...(methodIsAzureAiInference ? { path: OPENAI_AZURE_AI_INFERENCE_PATH } : {}),
|
||||
signal: this.abortController!.signal,
|
||||
})
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
|
|
@ -408,10 +428,10 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
|
||||
let response
|
||||
try {
|
||||
response = await this.client.chat.completions.create(
|
||||
requestOptions,
|
||||
methodIsAzureAiInference ? { path: OPENAI_AZURE_AI_INFERENCE_PATH } : {},
|
||||
)
|
||||
response = await this.client.chat.completions.create(requestOptions, {
|
||||
...(methodIsAzureAiInference ? { path: OPENAI_AZURE_AI_INFERENCE_PATH } : {}),
|
||||
signal: this.abortController!.signal,
|
||||
})
|
||||
} catch (error) {
|
||||
throw handleOpenAIError(error, this.providerName)
|
||||
}
|
||||
|
|
@ -442,6 +462,11 @@ export class OpenAiHandler extends BaseProvider implements SingleCompletionHandl
|
|||
const activeToolCallIds = new Set<string>()
|
||||
|
||||
for await (const chunk of stream) {
|
||||
// Check if request was aborted
|
||||
if (this.abortController?.signal.aborted) {
|
||||
break
|
||||
}
|
||||
|
||||
const delta = chunk.choices?.[0]?.delta
|
||||
const finishReason = chunk.choices?.[0]?.finish_reason
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue