@@ -46,6 +46,7 @@ const (
4646 RouterPD types.RoutingAlgorithm = "pd"
4747 VLLMEngine string = "vllm"
4848 SGLangEngine string = "sglang"
49+ TensorRTLLM string = "trtllm"
4950 SGLangBootstrapPort int64 = 8998
5051 SGLangBootstrapPortIdentifier string = "model.aibrix.ai/sglang-bootstrap-port"
5152 LLMEngineIdentifier string = constants .ModelLabelEngine
@@ -62,11 +63,10 @@ const (
6263 defaultRequestRateHighLoadThreshold = 1.0
6364 defaultRequestRateLowLoadThreshold = 0.25
6465
65- pdRouteValidateLLMEngineFail = "pd-validate-llm-engine-fail"
66- pdRouteFilterPrefillDecodePodsFail = "pd-filter-prefill-decode-pods-fail"
67- pdRoutePrefillRequestError = "pd-do-prefill-request-error"
68- pdRoutePrefillRequestSuccess = "pd-prefill-request-success"
69- pdRoutePrefillEmptyKVTransferParams = "pd-prefill-empty-kv-transfer-params"
66+ pdRouteValidateLLMEngineFail = "pd-validate-llm-engine-fail"
67+ pdRouteFilterPrefillDecodePodsFail = "pd-filter-prefill-decode-pods-fail"
68+ pdRoutePrefillRequestError = "pd-do-prefill-request-error"
69+ pdRoutePrefillRequestSuccess = "pd-prefill-request-success"
7070)
7171
7272const (
@@ -522,6 +522,7 @@ func (r *pdRouter) finalPDScore(routingCtx *types.RoutingContext,
522522
523523 return targetPrefillPod , targetDecodePod , nil
524524}
525+
525526func (r * pdRouter ) doPrefillRequest (routingCtx * types.RoutingContext , prefillPod * v1.Pod , llmEngine string ) error {
526527 // Prepare prefill request payload
527528 payload , err := r .preparePrefillPayload (routingCtx , prefillPod , llmEngine )
@@ -577,53 +578,57 @@ func (r *pdRouter) doPrefillRequest(routingCtx *types.RoutingContext, prefillPod
577578 }()
578579
579580 case VLLMEngine :
580- defer r .prefillRequestTracker .RemovePrefillRequest (routingCtx .RequestID )
581-
582581 // For vLLM, wait synchronously to get KV transfer params from response
583- responseData , err := r .executeHTTPRequest (apiURL , routingCtx , payload )
584- if err != nil {
585- klog .ErrorS (err , "prefill_request_failed" ,
586- "request_id" , routingCtx .RequestID ,
587- "llm_engine" , llmEngine ,
588- "prefill_pod" , prefillPod .Name ,
589- "prefill_pod_ip" , prefillPod .Status .PodIP ,
590- "elapsed" , routingCtx .Elapsed (time .Now ()))
591- return fmt .Errorf ("prefill request failed for request %s, pod %s: %w" , routingCtx .RequestID , prefillPod .Name , err )
592- }
582+ return r .handleSyncPrefill (routingCtx , prefillPod , llmEngine , apiURL , payload , fields , r .updateRoutingContextWithKVTransferParams , "KV transfer params" )
593583
594- // Update routing context with KV transfer params from prefill response
595- if err := r .updateRoutingContextWithKVTransferParams (routingCtx , responseData , prefillPod ); err != nil {
596- return fmt .Errorf ("failed to update routing context with KV transfer params for request %s: %w" , routingCtx .RequestID , err )
597- }
598-
599- routingCtx .PrefillEndTime = time .Now ()
600- fields = append (fields ,
601- "routing_time_taken" , routingCtx .PrefillStartTime .Sub (routingCtx .RequestTime ),
602- "prefill_time_taken" , routingCtx .PrefillEndTime .Sub (routingCtx .PrefillStartTime ),
603- "outstanding_prefill_requests" , r .prefillRequestTracker .GetPrefillRequestCountsForPod (prefillPod .Name )- 1 )
604- klog .InfoS ("prefill_request_end" , fields ... )
584+ case TensorRTLLM :
585+ // For TensorRT-LLM, wait synchronously to get disaggregated_params from response.
586+ // The prefill response contains first_gen_tokens and opaque_state needed by the decode worker.
587+ return r .handleSyncPrefill (routingCtx , prefillPod , llmEngine , apiURL , payload , fields , r .updateRoutingContextWithTRTDisaggParams , "TRT disagg params" )
605588
606589 default :
607- defer r .prefillRequestTracker .RemovePrefillRequest (routingCtx .RequestID )
608-
609590 // For unknown engines, use synchronous approach as a safe default
610- if _ , err := r .executeHTTPRequest (apiURL , routingCtx , payload ); err != nil {
611- klog .ErrorS (err , "prefill_request_failed" ,
612- "request_id" , routingCtx .RequestID ,
613- "llm_engine" , llmEngine ,
614- "prefill_pod" , prefillPod .Name ,
615- "prefill_pod_ip" , prefillPod .Status .PodIP ,
616- "elapsed" , routingCtx .Elapsed (time .Now ()))
617- return fmt .Errorf ("prefill request failed for request %s, pod %s: %w" , routingCtx .RequestID , prefillPod .Name , err )
591+ return r .handleSyncPrefill (routingCtx , prefillPod , llmEngine , apiURL , payload , fields , nil , "" )
592+ }
593+
594+ return nil
595+ }
596+
597+ // handleSyncPrefill executes a synchronous prefill request, optionally calling updateCtxFunc
598+ // to process the response. Pass nil for updateCtxFunc when no response processing is needed.
599+ func (r * pdRouter ) handleSyncPrefill (
600+ routingCtx * types.RoutingContext ,
601+ prefillPod * v1.Pod ,
602+ llmEngine , apiURL string ,
603+ payload []byte ,
604+ fields []interface {},
605+ updateCtxFunc func (* types.RoutingContext , map [string ]any , * v1.Pod ) error ,
606+ errorContext string ) error {
607+ defer r .prefillRequestTracker .RemovePrefillRequest (routingCtx .RequestID )
608+
609+ responseData , err := r .executeHTTPRequest (apiURL , routingCtx , payload )
610+ if err != nil {
611+ klog .ErrorS (err , "prefill_request_failed" ,
612+ "request_id" , routingCtx .RequestID ,
613+ "llm_engine" , llmEngine ,
614+ "prefill_pod" , prefillPod .Name ,
615+ "prefill_pod_ip" , prefillPod .Status .PodIP ,
616+ "elapsed" , routingCtx .Elapsed (time .Now ()))
617+ return fmt .Errorf ("prefill request failed for request %s, pod %s: %w" , routingCtx .RequestID , prefillPod .Name , err )
618+ }
619+
620+ if updateCtxFunc != nil {
621+ if err := updateCtxFunc (routingCtx , responseData , prefillPod ); err != nil {
622+ return fmt .Errorf ("failed to update routing context with %s for request %s: %w" , errorContext , routingCtx .RequestID , err )
618623 }
619- routingCtx .PrefillEndTime = time .Now ()
620- fields = append (fields ,
621- "routing_time_taken" , routingCtx .PrefillStartTime .Sub (routingCtx .RequestTime ),
622- "prefill_time_taken" , routingCtx .PrefillEndTime .Sub (routingCtx .PrefillStartTime ),
623- "outstanding_prefill_requests" , r .prefillRequestTracker .GetPrefillRequestCountsForPod (prefillPod .Name )- 1 )
624- klog .InfoS ("prefill_request_end" , fields ... )
625624 }
626625
626+ routingCtx .PrefillEndTime = time .Now ()
627+ fields = append (fields ,
628+ "routing_time_taken" , routingCtx .PrefillStartTime .Sub (routingCtx .RequestTime ),
629+ "prefill_time_taken" , routingCtx .PrefillEndTime .Sub (routingCtx .PrefillStartTime ),
630+ "outstanding_prefill_requests" , r .prefillRequestTracker .GetPrefillRequestCountsForPod (prefillPod .Name )- 1 )
631+ klog .InfoS ("prefill_request_end" , fields ... )
627632 return nil
628633}
629634
@@ -662,9 +667,22 @@ func (r *pdRouter) preparePrefillPayload(routingCtx *types.RoutingContext, pod *
662667 }
663668 }
664669
670+ if llmEngine == TensorRTLLM {
671+ // Signal to TensorRT-LLM that this is a context-only (prefill) request.
672+ // The prefill response will return disaggregated_params containing
673+ // first_gen_tokens and opaque_state, which are injected into the decode request.
674+ completionRequest ["disaggregated_params" ] = map [string ]any {
675+ "request_type" : "context_only" ,
676+ }
677+ }
678+
665679 // Set prefill-specific parameters
666680 completionRequest ["max_tokens" ] = 1
667- completionRequest ["max_completion_tokens" ] = 1
681+ if llmEngine == TensorRTLLM {
682+ delete (completionRequest , "max_completion_tokens" )
683+ } else {
684+ completionRequest ["max_completion_tokens" ] = 1
685+ }
668686 completionRequest ["stream" ] = false
669687 delete (completionRequest , "stream_options" )
670688
@@ -788,6 +806,57 @@ func (r *pdRouter) updateRoutingContextWithKVTransferParams(routingCtx *types.Ro
788806 return nil
789807}
790808
809+ func (r * pdRouter ) updateRoutingContextWithTRTDisaggParams (routingCtx * types.RoutingContext , responseData map [string ]any , prefillPod * v1.Pod ) error {
810+ // Parse the original request body
811+ var originalRequest map [string ]any
812+ if err := sonic .Unmarshal (routingCtx .ReqBody , & originalRequest ); err != nil {
813+ return fmt .Errorf ("failed to unmarshal original request body: %w" , err )
814+ }
815+
816+ // Extract disaggregated_params from prefill response.
817+ // TRT-LLM may return it at the top level or inside choices[0].
818+ var disaggParams any
819+ var exists bool
820+
821+ disaggParams , exists = responseData ["disaggregated_params" ]
822+ if ! exists {
823+ // Fallback: check choices[0] (TRT-LLM serializes handler output as a choice)
824+ if choices , ok := responseData ["choices" ].([]any ); ok && len (choices ) > 0 {
825+ if choice , ok := choices [0 ].(map [string ]any ); ok {
826+ disaggParams , exists = choice ["disaggregated_params" ]
827+ }
828+ }
829+ }
830+
831+ if ! exists {
832+ klog .InfoS ("no disaggregated_params in TRT prefill response" , "request_id" , routingCtx .RequestID )
833+ return nil
834+ }
835+
836+ disaggParamsMap , ok := disaggParams .(map [string ]any )
837+ if ! ok {
838+ return fmt .Errorf ("disaggregated_params has unexpected type %T, expected map[string]any" , disaggParams )
839+ }
840+
841+ // Override request_type to generation_only for the decode request
842+ disaggParamsMap ["request_type" ] = "generation_only"
843+ originalRequest ["disaggregated_params" ] = disaggParamsMap
844+
845+ updatedReqBody , err := sonic .Marshal (originalRequest )
846+ if err != nil {
847+ return fmt .Errorf ("failed to marshal updated request body: %w" , err )
848+ }
849+
850+ routingCtx .ReqBody = updatedReqBody
851+
852+ klog .InfoS ("updated routing context with disaggregated_params (TensorRT-LLM)" ,
853+ "request_id" , routingCtx .RequestID ,
854+ "prefill_pod" , prefillPod .Name ,
855+ "prefill_host" , prefillPod .Status .PodIP )
856+
857+ return nil
858+ }
859+
791860func (r * pdRouter ) SubscribedMetrics () []string {
792861 return []string {}
793862}
0 commit comments