Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 17 additions & 9 deletions core/providers/anthropic/anthropic.go
Original file line number Diff line number Diff line change
Expand Up @@ -637,7 +637,7 @@ func HandleAnthropicChatCompletionStreaming(
providerUtils.DrainLargePayloadRemainder(ctx)
}
if err != nil {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
if errors.Is(err, context.Canceled) {
return nil, providerUtils.EnrichError(ctx, &schemas.BifrostError{
IsBifrostError: false,
Expand All @@ -659,7 +659,7 @@ func HandleAnthropicChatCompletionStreaming(

// Check for HTTP errors
if resp.StatusCode() != fasthttp.StatusOK {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
return nil, providerUtils.EnrichError(ctx, parseAnthropicError(resp), jsonBody, nil, sendBackRawRequest, sendBackRawResponse)
}

Expand All @@ -684,7 +684,7 @@ func HandleAnthropicChatCompletionStreaming(
}
close(responseChan)
}()
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)

if resp.BodyStream() == nil {
bifrostErr := providerUtils.NewBifrostOperationError(
Expand Down Expand Up @@ -740,6 +740,10 @@ func HandleAnthropicChatCompletionStreaming(
}
eventType, eventDataBytes, readErr := sseReader.ReadEvent()
if readErr != nil {
// Recheck context cancellation
if ctx.Err() != nil {
return
}
if readErr != io.EOF {
ctx.SetValue(schemas.BifrostContextKeyStreamEndIndicator, true)
logger.Warn("Error reading %s stream: %v", providerName, readErr)
Expand Down Expand Up @@ -1105,7 +1109,7 @@ func HandleAnthropicResponsesStream(
providerUtils.DrainLargePayloadRemainder(ctx)
}
if err != nil {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
if errors.Is(err, context.Canceled) {
return nil, providerUtils.EnrichError(ctx, &schemas.BifrostError{
IsBifrostError: false,
Expand All @@ -1127,7 +1131,7 @@ func HandleAnthropicResponsesStream(

// Check for HTTP errors
if resp.StatusCode() != fasthttp.StatusOK {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
return nil, providerUtils.EnrichError(ctx, parseAnthropicError(resp), jsonBody, nil, sendBackRawRequest, sendBackRawResponse)
}

Expand All @@ -1152,7 +1156,7 @@ func HandleAnthropicResponsesStream(
}
close(responseChan)
}()
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
// If body stream is nil, return an error
if resp.BodyStream() == nil {
bifrostErr := providerUtils.NewBifrostOperationError(
Expand Down Expand Up @@ -1206,6 +1210,10 @@ func HandleAnthropicResponsesStream(
}
eventType, eventDataBytes, readErr := sseReader.ReadEvent()
if readErr != nil {
// Recheck context cancellation
if ctx.Err() != nil {
return
}
if readErr != io.EOF {
ctx.SetValue(schemas.BifrostContextKeyStreamEndIndicator, true)
logger.Warn("Error reading %s stream: %v", providerName, readErr)
Expand Down Expand Up @@ -2649,7 +2657,7 @@ func (provider *AnthropicProvider) PassthroughStream(

activeClient := providerUtils.PrepareResponseStreaming(ctx, provider.streamingClient, resp)
if err := activeClient.Do(fasthttpReq, resp); err != nil {
providerUtils.ReleaseStreamingResponse(resp)
providerUtils.ReleaseStreamingResponse(ctx, resp)
if errors.Is(err, context.Canceled) {
return nil, &schemas.BifrostError{
IsBifrostError: false,
Expand All @@ -2671,7 +2679,7 @@ func (provider *AnthropicProvider) PassthroughStream(

bodyStream := resp.BodyStream()
if bodyStream == nil {
providerUtils.ReleaseStreamingResponse(resp)
providerUtils.ReleaseStreamingResponse(ctx, resp)
return nil, providerUtils.NewBifrostOperationError(
"provider returned an empty stream body",
fmt.Errorf("provider returned an empty stream body"),
Expand Down Expand Up @@ -2702,7 +2710,7 @@ func (provider *AnthropicProvider) PassthroughStream(
}
close(ch)
}()
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
defer stopIdleTimeout()
defer stopCancellation()

Expand Down
12 changes: 6 additions & 6 deletions core/providers/azure/azure.go
Original file line number Diff line number Diff line change
Expand Up @@ -1014,7 +1014,7 @@ func (provider *AzureProvider) SpeechStream(ctx *schemas.BifrostContext, postHoo
// Make the request
requestErr := provider.client.Do(req, resp)
if requestErr != nil {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
if errors.Is(requestErr, context.Canceled) {
return nil, providerUtils.EnrichError(ctx, &schemas.BifrostError{
IsBifrostError: false,
Expand All @@ -1036,7 +1036,7 @@ func (provider *AzureProvider) SpeechStream(ctx *schemas.BifrostContext, postHoo

// Check for HTTP errors
if resp.StatusCode() != fasthttp.StatusOK {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
return nil, providerUtils.EnrichError(ctx, openai.ParseOpenAIError(resp), jsonBody, nil, sendBackRawRequest, sendBackRawResponse)
}

Expand All @@ -1057,7 +1057,7 @@ func (provider *AzureProvider) SpeechStream(ctx *schemas.BifrostContext, postHoo
close(responseChan)
}()
// Always release response on exit; bodyStream close should prevent indefinite blocking.
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)

reader, releaseGzip := providerUtils.DecompressStreamBody(resp)
defer releaseGzip()
Expand Down Expand Up @@ -3736,7 +3736,7 @@ func (provider *AzureProvider) PassthroughStream(
startTime := time.Now()

if err := activeClient.Do(fasthttpReq, resp); err != nil {
providerUtils.ReleaseStreamingResponse(resp)
providerUtils.ReleaseStreamingResponse(ctx, resp)
if errors.Is(err, context.Canceled) {
return nil, &schemas.BifrostError{
IsBifrostError: false,
Expand All @@ -3758,7 +3758,7 @@ func (provider *AzureProvider) PassthroughStream(

rawBodyStream := resp.BodyStream()
if rawBodyStream == nil {
providerUtils.ReleaseStreamingResponse(resp)
providerUtils.ReleaseStreamingResponse(ctx, resp)
return nil, providerUtils.NewBifrostOperationError("provider returned an empty stream body", fmt.Errorf("provider returned an empty stream body"))
}

Expand All @@ -3781,7 +3781,7 @@ func (provider *AzureProvider) PassthroughStream(
}
close(ch)
}()
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
defer stopIdleTimeout()
defer stopCancellation()

Expand Down
26 changes: 14 additions & 12 deletions core/providers/cohere/cohere.go
Original file line number Diff line number Diff line change
Expand Up @@ -460,7 +460,7 @@ func (provider *CohereProvider) ChatCompletionStream(ctx *schemas.BifrostContext
providerUtils.DrainLargePayloadRemainder(ctx)
}
if err != nil {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
if errors.Is(err, context.Canceled) {
return nil, providerUtils.EnrichError(ctx, &schemas.BifrostError{
IsBifrostError: false,
Expand All @@ -482,7 +482,7 @@ func (provider *CohereProvider) ChatCompletionStream(ctx *schemas.BifrostContext

// Check for HTTP errors
if resp.StatusCode() != fasthttp.StatusOK {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
return nil, providerUtils.EnrichError(ctx, parseCohereError(resp), jsonBody, nil, sendBackRawRequest, sendBackRawResponse)
}

Expand All @@ -509,7 +509,7 @@ func (provider *CohereProvider) ChatCompletionStream(ctx *schemas.BifrostContext
}
close(responseChan)
}()
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
// Decompress gzip-encoded streams transparently (no-op for non-gzip)
reader, releaseGzip := providerUtils.DecompressStreamBody(resp)
defer releaseGzip()
Expand Down Expand Up @@ -537,10 +537,11 @@ func (provider *CohereProvider) ChatCompletionStream(ctx *schemas.BifrostContext
}
data, readErr := sseReader.ReadDataLine()
if readErr != nil {
// Recheck context cancellation
if ctx.Err() != nil {
return
}
if readErr != io.EOF {
if ctx.Err() != nil {
return
}
ctx.SetValue(schemas.BifrostContextKeyStreamEndIndicator, true)
provider.logger.Warn("Error reading stream: %v", readErr)
providerUtils.ProcessAndSendError(ctx, postHookRunner, readErr, responseChan, provider.logger, postHookSpanFinalizer)
Expand Down Expand Up @@ -724,7 +725,7 @@ func (provider *CohereProvider) ResponsesStream(ctx *schemas.BifrostContext, pos
providerUtils.DrainLargePayloadRemainder(ctx)
}
if err != nil {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
if errors.Is(err, context.Canceled) {
return nil, providerUtils.EnrichError(ctx, &schemas.BifrostError{
IsBifrostError: false,
Expand All @@ -746,7 +747,7 @@ func (provider *CohereProvider) ResponsesStream(ctx *schemas.BifrostContext, pos

// Check for HTTP errors
if resp.StatusCode() != fasthttp.StatusOK {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
return nil, providerUtils.EnrichError(ctx, parseCohereError(resp), jsonBody, nil, sendBackRawRequest, sendBackRawResponse)
}

Expand All @@ -773,7 +774,7 @@ func (provider *CohereProvider) ResponsesStream(ctx *schemas.BifrostContext, pos
}
close(responseChan)
}()
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
// Decompress gzip-encoded streams transparently (no-op for non-gzip)
reader, releaseGzip := providerUtils.DecompressStreamBody(resp)
defer releaseGzip()
Expand Down Expand Up @@ -806,10 +807,11 @@ func (provider *CohereProvider) ResponsesStream(ctx *schemas.BifrostContext, pos
}
data, readErr := sseReader.ReadDataLine()
if readErr != nil {
// Recheck context cancellation
if ctx.Err() != nil {
return
}
if readErr != io.EOF {
if ctx.Err() != nil {
return
}
ctx.SetValue(schemas.BifrostContextKeyStreamEndIndicator, true)
provider.logger.Warn("Error reading stream: %v", readErr)
providerUtils.ProcessAndSendError(ctx, postHookRunner, readErr, responseChan, provider.logger, postHookSpanFinalizer)
Expand Down
6 changes: 3 additions & 3 deletions core/providers/elevenlabs/elevenlabs.go
Original file line number Diff line number Diff line change
Expand Up @@ -352,7 +352,7 @@ func (provider *ElevenlabsProvider) SpeechStream(ctx *schemas.BifrostContext, po
startTime := time.Now()
err := provider.streamingClient.Do(req, resp)
if err != nil {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
if errors.Is(err, context.Canceled) {
return nil, providerUtils.EnrichError(ctx, &schemas.BifrostError{
IsBifrostError: false,
Expand All @@ -374,7 +374,7 @@ func (provider *ElevenlabsProvider) SpeechStream(ctx *schemas.BifrostContext, po

// Check for HTTP errors
if resp.StatusCode() != fasthttp.StatusOK {
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
return nil, providerUtils.EnrichError(ctx, parseElevenlabsError(resp), jsonBody, nil, provider.sendBackRawRequest, provider.sendBackRawResponse)
}

Expand All @@ -392,7 +392,7 @@ func (provider *ElevenlabsProvider) SpeechStream(ctx *schemas.BifrostContext, po
}
close(responseChan)
}()
defer providerUtils.ReleaseStreamingResponse(resp)
defer providerUtils.ReleaseStreamingResponse(ctx, resp)
// Decompress gzip-encoded streams transparently (no-op for non-gzip)
reader, releaseGzip := providerUtils.DecompressStreamBody(resp)
defer releaseGzip()
Expand Down
Loading
Loading