Skip to content

Commit c5ec7ac

Browse files
mattwittwernightflight-dk
authored andcommitted
fix: RequestTracker counter mismatch (triton-inference-server#483)
* Fix RequestTracker counter mismatch in ScheduleSteps with parallel failures
1 parent 6c42f72 commit c5ec7ac

2 files changed

Lines changed: 73 additions & 25 deletions

File tree

src/ensemble_scheduler/ensemble_scheduler.cc

Lines changed: 71 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -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

193232
struct 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
646685
EnsembleContext::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
}

src/ensemble_scheduler/ensemble_scheduler.h

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,9 @@
2727

2828
#ifdef TRITON_ENABLE_ENSEMBLE
2929

30+
#include <condition_variable>
3031
#include <memory>
32+
#include <mutex>
3133

3234
#include "metric_model_reporter.h"
3335
#include "model.h"

0 commit comments

Comments
 (0)