Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
90 changes: 53 additions & 37 deletions pkg/epp/handlers/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -272,47 +272,31 @@ func (s *StreamingServer) Process(srv extProcPb.ExternalProcessor_ProcessServer)
reqCtx.respHeaderResp = s.generateResponseHeaderResponse(reqCtx)

case *extProcPb.ProcessingRequest_ResponseBody:
reqCtx.ResponseComplete = v.ResponseBody.EndOfStream
if reqCtx.ResponseComplete {
loggerTrace.Info("stream completed")
reqCtx.ResponseCompleteTimestamp = time.Now()
}
endOfStream := v.ResponseBody.EndOfStream
chunk := v.ResponseBody.Body

if reqCtx.modelServerStreaming {
// Currently we punt on response parsing if the modelServer is streaming, and we just passthrough.
s.HandleResponseBodyModelStreaming(ctx, reqCtx, v.ResponseBody.Body, v.ResponseBody.EndOfStream)
if v.ResponseBody.EndOfStream {
if _, err := s.director.HandleResponseBodyComplete(ctx, reqCtx); err != nil {
logger.Error(err, "error in HandleResponseBodyComplete")
}
metrics.RecordRequestLatencies(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp)
metrics.RecordResponseSizes(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.ResponseSize)
metrics.RecordNormalizedTimePerOutputToken(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp, reqCtx.Usage.CompletionTokens)
}
reqCtx.respBodyResp = generateResponseBodyResponses(v.ResponseBody.Body, v.ResponseBody.EndOfStream)
s.HandleResponseBodyModelStreaming(ctx, reqCtx, chunk, endOfStream)
reqCtx.respBodyResp = generateResponseBodyResponses(chunk, endOfStream)
} else {
body = append(body, v.ResponseBody.Body...)

// Message is buffered, we can read and decode.
if v.ResponseBody.EndOfStream {
reqCtx.ResponseSize = len(body)
reqCtx.respBodyResp = generateResponseBodyResponses(body, true)

var responseErr error
reqCtx, responseErr = s.HandleResponseBody(ctx, reqCtx, body)
if responseErr != nil {
break
}
metrics.RecordRequestLatencies(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp)
metrics.RecordResponseSizes(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.ResponseSize)
metrics.RecordInputTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.PromptTokens)
metrics.RecordOutputTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.CompletionTokens)
if reqCtx.Usage.PromptTokenDetails != nil {
metrics.RecordPromptCachedTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.PromptTokenDetails.CachedTokens)
}
}
body = append(body, chunk...)
}

if endOfStream {
err = s.finishResponse(ctx, reqCtx, body)
}
case *extProcPb.ProcessingRequest_ResponseTrailers:
// This is currently unused.
// For HTTP, the response trailer is not sent. Thus, it won't achieve this case.
Comment thread
zetxqx marked this conversation as resolved.
Outdated

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

even if it was sent in the http case and this logic gets executed, we should be ok, right?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

yes, because we only set reqCtx.ResponseComplete to true in the new finishResponse. and the first early return check for finishResponse makes it idempotent in the whole request lifecycle

image

// For gRPC(over HTTP2), the protocol relies on responseTrialers to determine whether a response is complete.
// More info: https://chromium.googlesource.com/external/github.com/grpc/grpc/+/HEAD/doc/PROTOCOL-HTTP2.md#responses
err = s.finishResponse(ctx, reqCtx, body)
if err == nil {
reqCtx.respTrailerResp = &extProcPb.ProcessingResponse{
Response: &extProcPb.ProcessingResponse_ResponseTrailers{
ResponseTrailers: &extProcPb.TrailersResponse{},
},
}
}
}

// Handle the err and fire an immediate response.
Expand All @@ -339,6 +323,38 @@ func (s *StreamingServer) Process(srv extProcPb.ExternalProcessor_ProcessServer)
}
}

func (s *StreamingServer) finishResponse(ctx context.Context, reqCtx *RequestContext, body []byte) error {
if reqCtx.ResponseComplete {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

add a comment pls on what this means, iiuc it means we already completed the response and so we don't want to execute finishResponse again, right?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

yes, added a comment in 4704366

return nil
}

reqCtx.ResponseComplete = true
reqCtx.ResponseCompleteTimestamp = time.Now()
reqCtx.ResponseSize = len(body)

if reqCtx.modelServerStreaming {
if _, err := s.director.HandleResponseBodyComplete(ctx, reqCtx); err != nil {
log.FromContext(ctx).Error(err, "error in HandleResponseBodyComplete")
}
metrics.RecordRequestLatencies(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp)
metrics.RecordResponseSizes(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.ResponseSize)
metrics.RecordNormalizedTimePerOutputToken(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp, reqCtx.Usage.CompletionTokens)
} else {
reqCtx.respBodyResp = generateResponseBodyResponses(body, true)
if _, err := s.HandleResponseBody(ctx, reqCtx, body); err != nil {
return err
}
metrics.RecordRequestLatencies(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp)
metrics.RecordResponseSizes(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.ResponseSize)
metrics.RecordInputTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.PromptTokens)
metrics.RecordOutputTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.CompletionTokens)
if reqCtx.Usage.PromptTokenDetails != nil {
metrics.RecordPromptCachedTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.PromptTokenDetails.CachedTokens)
}
}
return nil
}

// updateStateAndSendIfNeeded checks state and can send mutiple responses in a single pass, but only if ordered properly.
// Order of requests matter in FULL_DUPLEX_STREAMING. For both request and response, the order of response sent back MUST be: Header->Body->Trailer, with trailer being optional.
func (r *RequestContext) updateStateAndSendIfNeeded(srv extProcPb.ExternalProcessor_ProcessServer, logger logr.Logger) error {
Expand Down
Loading