@@ -87,8 +87,24 @@ class RequestTracker {
8787 {
8888 }
8989
90+ // Accessed without additional synchronization while protected by
91+ // EnsembleContext::mutex_.
9092 std::unique_ptr<InferenceRequest>& Request () { return request_; }
9193
94+ // Used from paths where request_ may be released concurrently.
95+ bool IsCancelled ()
96+ {
97+ std::lock_guard<std::mutex> lk (mtx_);
98+ return (request_ == nullptr ) || request_->IsCancelled ();
99+ }
100+
101+ std::string LogRequest ()
102+ {
103+ std::lock_guard<std::mutex> lk (mtx_);
104+ return (request_ != nullptr ) ? request_->LogRequest ()
105+ : std::string (" [request released] " );
106+ }
107+
92108 InferenceStatsAggregator* StatsAggregator () { return stats_aggregator_; }
93109
94110 MetricModelReporter* MetricReporter () { return metric_reporter_; }
@@ -141,6 +157,23 @@ class RequestTracker {
141157 status_ = status;
142158 }
143159
160+ void RespondIfError (const Status& status, FailureReason reason)
161+ {
162+ std::lock_guard<std::mutex> lk (mtx_);
163+ if (request_ != nullptr ) {
164+ InferenceRequest::RespondIfError (
165+ request_, status, false /* release_request */ , reason);
166+ }
167+ }
168+
169+ void SendFlags (const uint32_t flags)
170+ {
171+ std::lock_guard<std::mutex> lk (mtx_);
172+ if (request_ != nullptr ) {
173+ request_->ResponseFactory ()->SendFlags (flags);
174+ }
175+ }
176+
144177 private:
145178 std::mutex mtx_;
146179 uint32_t inflight_request_counter_;
@@ -153,6 +186,7 @@ class RequestTracker {
153186 triton::common::ThreadPool* const callback_pool_;
154187};
155188
189+ using RequestTrackerReference = std::shared_ptr<RequestTracker>;
156190// Step is used as 'userp' and keeps ensemble context alive
157191// until no more internal requests are inflight.
158192// Step contains metadata, and status for the
@@ -188,6 +222,11 @@ struct Step {
188222 const bool preserve_responses_order_;
189223
190224 size_t step_idx_;
225+
226+ // Heap allocation passed as the release-callback userp. The allocation
227+ // stores a shared_ptr<RequestTracker> so the tracker stays alive until the
228+ // release callback or local failure cleanup drops this reference.
229+ RequestTrackerReference* callback_tracker_ref_{nullptr };
191230};
192231
193232struct TensorData {
@@ -396,7 +435,7 @@ class EnsembleContext {
396435
397436 // Objects related to the ensemble infer request
398437 Status ensemble_status_;
399- RequestTracker* request_tracker_;
438+ std::shared_ptr< RequestTracker> request_tracker_;
400439 // Use in conjunction with 'is_decoupled_' in EnsembleInfo to
401440 // better distinguish ensemble ending behavior (see annotation in
402441 // FinishEnsemble for details).
@@ -429,7 +468,7 @@ EnsembleContext::EnsembleContext(
429468{
430469 uint64_t compute_start_ns = 0 ;
431470 INFER_STATS_SET_TIMESTAMP (compute_start_ns);
432- request_tracker_ = new RequestTracker (
471+ request_tracker_ = std::make_shared< RequestTracker> (
433472 std::move (request), compute_start_ns, metric_reporter, stats_aggregator,
434473 callback_pool);
435474
@@ -646,16 +685,17 @@ void
646685EnsembleContext::RequestComplete (
647686 TRITONSERVER_InferenceRequest* request, const uint32_t flags, void * userp)
648687{
649- auto request_tracker = reinterpret_cast <RequestTracker*>(userp);
688+ auto callback_tracker_ref = reinterpret_cast <RequestTrackerReference*>(userp);
689+ auto request_tracker = *callback_tracker_ref;
650690 auto pool = request_tracker->CallbackPool ();
651- auto fn = [request, flags, request_tracker]() {
691+ auto fn = [request, flags, request_tracker, callback_tracker_ref ]() {
652692 if ((flags & TRITONSERVER_REQUEST_RELEASE_ALL ) != 0 ) {
693+ std::unique_ptr<RequestTrackerReference> managed_callback_tracker_ref (
694+ callback_tracker_ref);
653695 LOG_TRITONSERVER_ERROR (
654696 TRITONSERVER_InferenceRequestDelete (request),
655697 " deleting ensemble inference request" );
656- if (request_tracker->DecrementCounter ()) {
657- delete request_tracker;
658- }
698+ request_tracker->DecrementCounter ();
659699 }
660700 };
661701
@@ -1070,12 +1110,19 @@ EnsembleContext::InitStep(
10701110 irequest->SetSecondaryStatsAggregator (
10711111 &request_tracker_->ContextStatsAggregator ());
10721112#endif
1113+ // Heap-allocate the release-callback userp because the C callback API only
1114+ // stores a raw void*. The heap object itself is single-owner here, while the
1115+ // object stored inside it is a shared_ptr<RequestTracker> that keeps the
1116+ // tracker alive until the callback or local cleanup drops this reference.
1117+ auto callback_tracker_ref =
1118+ std::make_unique<RequestTrackerReference>(request_tracker_);
10731119 irequest->SetResponseCallback (
10741120 reinterpret_cast <ResponseAllocator*>(allocator_.get ()), step->get (),
10751121 ResponseComplete, step->get ());
1076- irequest->SetReleaseCallback (RequestComplete, request_tracker_ );
1122+ irequest->SetReleaseCallback (RequestComplete, callback_tracker_ref. get () );
10771123
10781124 RETURN_IF_ERROR (irequest->PrepareForInference ());
1125+ (*step)->callback_tracker_ref_ = callback_tracker_ref.release ();
10791126
10801127#ifdef TRITON_ENABLE_TRACING
10811128 auto & parent_trace = request_tracker_->Request ()->TraceProxy ();
@@ -1220,16 +1267,14 @@ EnsembleContext::FinishEnsemble(std::unique_ptr<InferenceResponse>&& response)
12201267 ensemble_status_ = Status (
12211268 Status::Code::INVALID_ARG ,
12221269 " in ensemble '" + info_->ensemble_name_ + " ', " +
1223- request_tracker_->Request ()-> LogRequest () +
1270+ request_tracker_->LogRequest () +
12241271 " unexpected deadlock, at least one output is not set while no "
12251272 " more "
12261273 " ensemble steps can be made" );
1227- InferenceRequest::RespondIfError (
1228- request_tracker_->Request (), ensemble_status_,
1229- false /* release_requests */ , FailureReason::OTHER );
1274+ request_tracker_->RespondIfError (
1275+ ensemble_status_, FailureReason::OTHER );
12301276 } else {
1231- request_tracker_->Request ()->ResponseFactory ()->SendFlags (
1232- TRITONSERVER_RESPONSE_COMPLETE_FINAL );
1277+ request_tracker_->SendFlags (TRITONSERVER_RESPONSE_COMPLETE_FINAL );
12331278 }
12341279 }
12351280 } else {
@@ -1239,9 +1284,8 @@ EnsembleContext::FinishEnsemble(std::unique_ptr<InferenceResponse>&& response)
12391284 std::move (response), TRITONSERVER_RESPONSE_COMPLETE_FINAL ,
12401285 ensemble_status_);
12411286 } else {
1242- InferenceRequest::RespondIfError (
1243- request_tracker_->Request (), ensemble_status_,
1244- false /* release_requests */ , FailureReason::OTHER );
1287+ request_tracker_->RespondIfError (
1288+ ensemble_status_, FailureReason::OTHER );
12451289 }
12461290 error_response_sent_ = true ;
12471291 }
@@ -1251,10 +1295,8 @@ EnsembleContext::FinishEnsemble(std::unique_ptr<InferenceResponse>&& response)
12511295 // Reach here when the ensemble execution comes to the end,
12521296 // 'ensemble_status_' at this point is representative.
12531297 request_tracker_->SetStatus (ensemble_status_);
1254- if (request_tracker_->DecrementCounter ()) {
1255- delete request_tracker_;
1256- }
1257- request_tracker_ = nullptr ;
1298+ request_tracker_->DecrementCounter ();
1299+ request_tracker_.reset ();
12581300 }
12591301 return ensemble_status_;
12601302}
@@ -1436,7 +1478,7 @@ EnsembleContext::ScheduleSteps(
14361478 if (should_schedule) {
14371479 // If the ensemble request is cancelled, propagate the cancellation to the
14381480 // next request step.
1439- if (context->request_tracker_ ->Request ()-> IsCancelled ()) {
1481+ if (context->request_tracker_ ->IsCancelled ()) {
14401482 step->request_ ->Cancel ();
14411483 }
14421484 // Acquire a slot from the per-step shared limiter only for steps that
@@ -1469,16 +1511,20 @@ EnsembleContext::ScheduleSteps(
14691511 // Reaching here means the step is not being scheduled, update corresponding
14701512 // counters and attempt to finish ensemble if it is the last step.
14711513
1472-
14731514 // Release the limiter slot if one was acquired, and update counters.
14741515 if (should_schedule &&
14751516 !context->info_ ->step_inflight_request_limiters_ .empty ()) {
14761517 context->info_ ->step_inflight_request_limiters_ [this_step_idx]->Release ();
14771518 }
14781519
14791520 std::lock_guard<std::mutex> lock (context->mutex_ );
1480- // Decrement only when IncrementCounter was called. An unconditional
1481- // decrement would underflow the counter and cause a use-after-free.
1521+ // The request never reaches the callback-owned release path, so drop the
1522+ // heap-allocated callback userp here.
1523+ delete step->callback_tracker_ref_ ;
1524+ step->callback_tracker_ref_ = nullptr ;
1525+ // Only undo IncrementCounter() for steps that actually reached the
1526+ // scheduling path. Otherwise the counter can underflow and release the
1527+ // top-level request while FinishEnsemble is still using it.
14821528 if (should_schedule) {
14831529 context->request_tracker_ ->DecrementCounter ();
14841530 }
0 commit comments