Skip to content

Commit d5b64d9

Browse files
committed
processing respTrailers for gRPC request
1 parent b32374e commit d5b64d9

2 files changed

Lines changed: 280 additions & 140 deletions

File tree

pkg/epp/handlers/server.go

Lines changed: 53 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -272,47 +272,31 @@ func (s *StreamingServer) Process(srv extProcPb.ExternalProcessor_ProcessServer)
272272
reqCtx.respHeaderResp = s.generateResponseHeaderResponse(reqCtx)
273273

274274
case *extProcPb.ProcessingRequest_ResponseBody:
275-
reqCtx.ResponseComplete = v.ResponseBody.EndOfStream
276-
if reqCtx.ResponseComplete {
277-
loggerTrace.Info("stream completed")
278-
reqCtx.ResponseCompleteTimestamp = time.Now()
279-
}
275+
endOfStream := v.ResponseBody.EndOfStream
276+
chunk := v.ResponseBody.Body
277+
280278
if reqCtx.modelServerStreaming {
281-
// Currently we punt on response parsing if the modelServer is streaming, and we just passthrough.
282-
s.HandleResponseBodyModelStreaming(ctx, reqCtx, v.ResponseBody.Body, v.ResponseBody.EndOfStream)
283-
if v.ResponseBody.EndOfStream {
284-
if _, err := s.director.HandleResponseBodyComplete(ctx, reqCtx); err != nil {
285-
logger.Error(err, "error in HandleResponseBodyComplete")
286-
}
287-
metrics.RecordRequestLatencies(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp)
288-
metrics.RecordResponseSizes(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.ResponseSize)
289-
metrics.RecordNormalizedTimePerOutputToken(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp, reqCtx.Usage.CompletionTokens)
290-
}
291-
reqCtx.respBodyResp = generateResponseBodyResponses(v.ResponseBody.Body, v.ResponseBody.EndOfStream)
279+
s.HandleResponseBodyModelStreaming(ctx, reqCtx, chunk, endOfStream)
280+
reqCtx.respBodyResp = generateResponseBodyResponses(chunk, endOfStream)
292281
} else {
293-
body = append(body, v.ResponseBody.Body...)
294-
295-
// Message is buffered, we can read and decode.
296-
if v.ResponseBody.EndOfStream {
297-
reqCtx.ResponseSize = len(body)
298-
reqCtx.respBodyResp = generateResponseBodyResponses(body, true)
299-
300-
var responseErr error
301-
reqCtx, responseErr = s.HandleResponseBody(ctx, reqCtx, body)
302-
if responseErr != nil {
303-
break
304-
}
305-
metrics.RecordRequestLatencies(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp)
306-
metrics.RecordResponseSizes(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.ResponseSize)
307-
metrics.RecordInputTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.PromptTokens)
308-
metrics.RecordOutputTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.CompletionTokens)
309-
if reqCtx.Usage.PromptTokenDetails != nil {
310-
metrics.RecordPromptCachedTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.PromptTokenDetails.CachedTokens)
311-
}
312-
}
282+
body = append(body, chunk...)
283+
}
284+
285+
if endOfStream {
286+
err = s.finishResponse(ctx, reqCtx, body)
313287
}
314288
case *extProcPb.ProcessingRequest_ResponseTrailers:
315-
// This is currently unused.
289+
// For HTTP, the response trailer is not sent. Thus, it won't achieve this case.
290+
// For gRPC(over HTTP2), the protocol relies on responseTrialers to determine whether a response is complete.
291+
// More info: https://chromium.googlesource.com/external/github.com/grpc/grpc/+/HEAD/doc/PROTOCOL-HTTP2.md#responses
292+
err = s.finishResponse(ctx, reqCtx, body)
293+
if err == nil {
294+
reqCtx.respTrailerResp = &extProcPb.ProcessingResponse{
295+
Response: &extProcPb.ProcessingResponse_ResponseTrailers{
296+
ResponseTrailers: &extProcPb.TrailersResponse{},
297+
},
298+
}
299+
}
316300
}
317301

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

326+
func (s *StreamingServer) finishResponse(ctx context.Context, reqCtx *RequestContext, body []byte) error {
327+
if reqCtx.ResponseComplete {
328+
return nil
329+
}
330+
331+
reqCtx.ResponseComplete = true
332+
reqCtx.ResponseCompleteTimestamp = time.Now()
333+
reqCtx.ResponseSize = len(body)
334+
335+
if reqCtx.modelServerStreaming {
336+
if _, err := s.director.HandleResponseBodyComplete(ctx, reqCtx); err != nil {
337+
log.FromContext(ctx).Error(err, "error in HandleResponseBodyComplete")
338+
}
339+
metrics.RecordRequestLatencies(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp)
340+
metrics.RecordResponseSizes(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.ResponseSize)
341+
metrics.RecordNormalizedTimePerOutputToken(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp, reqCtx.Usage.CompletionTokens)
342+
} else {
343+
reqCtx.respBodyResp = generateResponseBodyResponses(body, true)
344+
if _, err := s.HandleResponseBody(ctx, reqCtx, body); err != nil {
345+
return err
346+
}
347+
metrics.RecordRequestLatencies(ctx, reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.RequestReceivedTimestamp, reqCtx.ResponseCompleteTimestamp)
348+
metrics.RecordResponseSizes(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.ResponseSize)
349+
metrics.RecordInputTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.PromptTokens)
350+
metrics.RecordOutputTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.CompletionTokens)
351+
if reqCtx.Usage.PromptTokenDetails != nil {
352+
metrics.RecordPromptCachedTokens(reqCtx.IncomingModelName, reqCtx.TargetModelName, reqCtx.Usage.PromptTokenDetails.CachedTokens)
353+
}
354+
}
355+
return nil
356+
}
357+
342358
// updateStateAndSendIfNeeded checks state and can send mutiple responses in a single pass, but only if ordered properly.
343359
// 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.
344360
func (r *RequestContext) updateStateAndSendIfNeeded(srv extProcPb.ExternalProcessor_ProcessServer, logger logr.Logger) error {

0 commit comments

Comments
 (0)