Skip to content

Commit 29b6b0a

Browse files
committed
fix(pylon): bind request state to individual lifetimes
Fence queue guards and observations with an internal request identity so a reused correlation ID cannot let earlier work change the replacement. Own queue cleanup across body validation and rejection. Refs: #1817
1 parent 344a6e4 commit 29b6b0a

8 files changed

Lines changed: 257 additions & 59 deletions

File tree

‎src/libraries/rust/stargate/crates/pylon-lib/src/bringup/upstream.rs‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -212,6 +212,7 @@ pub(super) async fn send_completion_request(
212212
TunnelRequestObserver::accepted(
213213
RequestObservationEndpoint::ChatCompletions,
214214
RequiredTunnelHeaders {
215+
request_instance: Default::default(),
215216
request_id: request_id.clone(),
216217
routing_key: None,
217218
model_id: model_id.to_string(),

‎src/libraries/rust/stargate/crates/pylon-lib/src/queue_admission.rs‎

Lines changed: 152 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,7 @@ use stargate_protocol::common::{
2424
use stargate_protocol::tunnel_contract::HEADER_STARGATE_EXPECTED_QUEUE_MS;
2525

2626
use crate::request_observer::{RequestObservation, RequestObservationState, RequiredTunnelHeaders};
27-
use crate::runtime_state::ModelGeneration;
27+
use crate::runtime_state::{ModelGeneration, RequestInstance};
2828

2929
pub(crate) const RETRY_REASON_QUEUE_ESTIMATE_MISMATCH: &str = "queue_estimate_mismatch";
3030

@@ -55,10 +55,16 @@ pub(crate) struct LiveRequestState {
5555

5656
#[derive(Debug, Default)]
5757
struct QueueAdmissionState {
58-
requests: HashMap<String, LiveRequest>,
58+
requests: HashMap<String, LiveRequestRecord>,
5959
models: HashMap<ModelGeneration, QueueModelState>,
6060
}
6161

62+
#[derive(Debug)]
63+
struct LiveRequestRecord {
64+
instance: Option<RequestInstance>,
65+
request: LiveRequest,
66+
}
67+
6268
#[derive(Debug, Default)]
6369
struct QueueModelState {
6470
last_mean_input_tps: Option<f64>,
@@ -155,6 +161,7 @@ pub(crate) struct QueueModelSnapshot {
155161
pub(crate) struct QueueTrackedRequestGuard {
156162
live_requests: LiveRequestState,
157163
request_id: String,
164+
instance: RequestInstance,
158165
finished: bool,
159166
}
160167

@@ -221,7 +228,7 @@ impl LiveRequestState {
221228
.lock()
222229
.requests
223230
.get(request_id)
224-
.map(|request| request.generation().clone())
231+
.map(|record| record.request.generation().clone())
225232
}
226233

227234
pub(crate) fn update_generation_throughput(
@@ -247,7 +254,7 @@ impl LiveRequestState {
247254
let request_ids = state
248255
.requests
249256
.iter()
250-
.filter(|(_, request)| request.generation() == generation)
257+
.filter(|(_, record)| record.request.generation() == generation)
251258
.map(|(request_id, _)| request_id.clone())
252259
.collect::<Vec<_>>();
253260
for request_id in request_ids {
@@ -302,7 +309,7 @@ impl LiveRequestState {
302309
let excluded_request = state
303310
.requests
304311
.get(&required.request_id)
305-
.and_then(|request| match request {
312+
.and_then(|record| match &record.request {
306313
LiveRequest::Queue(queue, _) => Some(queue),
307314
LiveRequest::Observed(_) => None,
308315
})
@@ -343,11 +350,11 @@ impl LiveRequestState {
343350
self.track_generation_request(required, ModelGeneration::new(required.model_id.clone(), 0))
344351
}
345352

346-
pub(crate) fn track_generation_request(
353+
pub(crate) fn begin_request(
347354
&self,
348355
required: &RequiredTunnelHeaders,
349356
generation: ModelGeneration,
350-
) -> QueueTrackedRequestGuard {
357+
) {
351358
let request_id = required.request_id.clone();
352359
let request = TrackedPromptRequest {
353360
generation,
@@ -360,12 +367,40 @@ impl LiveRequestState {
360367
let mut state = self.inner.lock();
361368
let observed = state
362369
.remove_request(&request_id)
363-
.and_then(|(_, request)| request.into_observed());
364-
state.insert_request(request_id.clone(), LiveRequest::Queue(request, observed));
370+
.and_then(|(_, request)| request.request.into_observed());
371+
state.insert_request(
372+
request_id,
373+
LiveRequest::Queue(request, observed),
374+
Some(required.request_instance.clone()),
375+
);
365376
}
377+
}
378+
379+
#[cfg(test)]
380+
pub(crate) fn track_generation_request(
381+
&self,
382+
required: &RequiredTunnelHeaders,
383+
generation: ModelGeneration,
384+
) -> QueueTrackedRequestGuard {
385+
self.begin_request(required, generation);
386+
self.request_guard(required)
387+
}
388+
389+
pub(crate) fn track_existing_request(
390+
&self,
391+
required: &RequiredTunnelHeaders,
392+
) -> Option<QueueTrackedRequestGuard> {
393+
self.inner
394+
.lock()
395+
.owns_instance(&required.request_id, &required.request_instance)
396+
.then(|| self.request_guard(required))
397+
}
398+
399+
fn request_guard(&self, required: &RequiredTunnelHeaders) -> QueueTrackedRequestGuard {
366400
QueueTrackedRequestGuard {
367401
live_requests: self.clone(),
368-
request_id,
402+
request_id: required.request_id.clone(),
403+
instance: required.request_instance.clone(),
369404
finished: false,
370405
}
371406
}
@@ -390,22 +425,30 @@ impl LiveRequestState {
390425
observe: impl FnOnce(&RequestObservationTransition),
391426
) -> RequestObservationTransition {
392427
let generation = ModelGeneration::new(observation.model_id.clone(), 0);
393-
self.transition_generation_observation_with(observation, Some(&generation), observe)
428+
self.transition_generation_observation_with(observation, Some(&generation), None, observe)
429+
.unwrap()
394430
}
395431

396432
pub(crate) fn transition_generation_observation_with(
397433
&self,
398434
observation: &RequestObservation,
399435
generation: Option<&ModelGeneration>,
436+
instance: Option<&RequestInstance>,
400437
observe: impl FnOnce(&RequestObservationTransition),
401-
) -> RequestObservationTransition {
438+
) -> Option<RequestObservationTransition> {
402439
let _order = self.observation_order.lock();
403-
let transition = self
404-
.inner
405-
.lock()
406-
.transition_observation(observation, generation);
440+
let transition = {
441+
let mut state = self.inner.lock();
442+
if let Some(instance) = instance
443+
&& generation.is_some()
444+
&& !state.owns_instance(&observation.request_id, instance)
445+
{
446+
return None;
447+
}
448+
state.transition_observation(observation, generation, instance)
449+
};
407450
observe(&transition);
408-
transition
451+
Some(transition)
409452
}
410453

411454
pub(crate) fn update_active_output_tps(
@@ -418,24 +461,44 @@ impl LiveRequestState {
418461
.update_active_output_tps(request_id, active_chat_output_tps)
419462
}
420463

421-
pub(crate) fn finish_queue_request(&self, request_id: &str) {
464+
pub(crate) fn finish_queue_request(
465+
&self,
466+
request_id: &str,
467+
instance: Option<&RequestInstance>,
468+
) {
422469
let mut state = self.inner.lock();
423-
if let Some((request_id, request)) = state.remove_request(request_id)
424-
&& let Some(observed) = request.into_observed()
470+
if instance.is_some_and(|instance| !state.owns_instance(request_id, instance)) {
471+
return;
472+
}
473+
if let Some((request_id, record)) = state.remove_request(request_id)
474+
&& let Some(observed) = record.request.into_observed()
425475
{
426-
state.insert_request(request_id, LiveRequest::Observed(observed));
476+
state.insert_request(request_id, LiveRequest::Observed(observed), record.instance);
427477
}
428478
}
429479
}
430480

431481
impl QueueAdmissionState {
482+
fn owns_instance(&self, request_id: &str, instance: &RequestInstance) -> bool {
483+
self.requests
484+
.get(request_id)
485+
.is_some_and(|record| record.instance.as_ref() == Some(instance))
486+
}
487+
432488
fn transition_observation(
433489
&mut self,
434490
observation: &RequestObservation,
435491
generation: Option<&ModelGeneration>,
492+
instance: Option<&RequestInstance>,
436493
) -> RequestObservationTransition {
494+
let prior = self.remove_request(&observation.request_id);
495+
let instance = instance.cloned().or_else(|| {
496+
prior
497+
.as_ref()
498+
.and_then(|(_, record)| record.instance.clone())
499+
});
437500
let (request_id, prior_queue, prior_observed) =
438-
match self.remove_request(&observation.request_id) {
501+
match prior.map(|(id, record)| (id, record.request)) {
439502
Some((request_id, LiveRequest::Queue(queue, observed))) => {
440503
(request_id, Some(queue), observed)
441504
}
@@ -493,6 +556,7 @@ impl QueueAdmissionState {
493556
},
494557
Some(current.clone()),
495558
),
559+
instance,
496560
);
497561
Some(current)
498562
}
@@ -513,14 +577,28 @@ impl QueueAdmissionState {
513577
request_id: &str,
514578
active_chat_output_tps: Option<f64>,
515579
) -> Option<String> {
516-
let (request_id, mut request) = self.remove_request(request_id)?;
517-
let model_id = request.update_active_output_tps(active_chat_output_tps);
518-
self.insert_request(request_id, request);
580+
let (request_id, mut record) = self.remove_request(request_id)?;
581+
let model_id = record
582+
.request
583+
.update_active_output_tps(active_chat_output_tps);
584+
self.insert_request(request_id, record.request, record.instance);
519585
model_id
520586
}
521587

522-
fn advance_request_phase(&mut self, request_id: &str, next_phase: TrackedPromptPhase) {
523-
let Some(LiveRequest::Queue(request, _)) = self.requests.get_mut(request_id) else {
588+
fn advance_request_phase(
589+
&mut self,
590+
request_id: &str,
591+
instance: &RequestInstance,
592+
next_phase: TrackedPromptPhase,
593+
) {
594+
if !self.owns_instance(request_id, instance) {
595+
return;
596+
}
597+
let Some(LiveRequest::Queue(request, _)) = self
598+
.requests
599+
.get_mut(request_id)
600+
.map(|record| &mut record.request)
601+
else {
524602
return;
525603
};
526604
if next_phase <= request.phase {
@@ -535,15 +613,21 @@ impl QueueAdmissionState {
535613
request.phase = next_phase;
536614
}
537615

538-
fn remove_request(&mut self, request_id: &str) -> Option<(String, LiveRequest)> {
539-
let (request_id, request) = self.requests.remove_entry(request_id)?;
540-
self.adjust_live_request(&request, -1);
541-
Some((request_id, request))
616+
fn remove_request(&mut self, request_id: &str) -> Option<(String, LiveRequestRecord)> {
617+
let (request_id, record) = self.requests.remove_entry(request_id)?;
618+
self.adjust_live_request(&record.request, -1);
619+
Some((request_id, record))
542620
}
543621

544-
fn insert_request(&mut self, request_id: String, request: LiveRequest) {
622+
fn insert_request(
623+
&mut self,
624+
request_id: String,
625+
request: LiveRequest,
626+
instance: Option<RequestInstance>,
627+
) {
545628
self.adjust_live_request(&request, 1);
546-
self.requests.insert(request_id, request);
629+
self.requests
630+
.insert(request_id, LiveRequestRecord { instance, request });
547631
}
548632

549633
fn adjust_live_request(&mut self, request: &LiveRequest, delta: i8) {
@@ -786,17 +870,26 @@ impl TrackedPromptPhase {
786870
impl QueueTrackedRequestGuard {
787871
pub(crate) fn on_backend_submission(&mut self) {
788872
let mut state = self.live_requests.inner.lock();
789-
state.advance_request_phase(&self.request_id, TrackedPromptPhase::InputProcessing);
873+
state.advance_request_phase(
874+
&self.request_id,
875+
&self.instance,
876+
TrackedPromptPhase::InputProcessing,
877+
);
790878
}
791879

792880
pub(crate) fn observe_output(&mut self) {
793881
let mut state = self.live_requests.inner.lock();
794-
state.advance_request_phase(&self.request_id, TrackedPromptPhase::OutputGeneration);
882+
state.advance_request_phase(
883+
&self.request_id,
884+
&self.instance,
885+
TrackedPromptPhase::OutputGeneration,
886+
);
795887
}
796888

797889
pub(crate) fn finish(&mut self) {
798890
if !self.finished {
799-
self.live_requests.finish_queue_request(&self.request_id);
891+
self.live_requests
892+
.finish_queue_request(&self.request_id, Some(&self.instance));
800893
self.finished = true;
801894
}
802895
}
@@ -840,7 +933,7 @@ mod tests {
840933
.lock()
841934
.requests
842935
.values()
843-
.filter(|request| matches!(request, LiveRequest::Queue(..)))
936+
.filter(|record| matches!(record.request, LiveRequest::Queue(..)))
844937
.count()
845938
}
846939
}
@@ -856,6 +949,7 @@ mod tests {
856949
input_tokens: u64,
857950
) -> RequiredTunnelHeaders {
858951
RequiredTunnelHeaders {
952+
request_instance: Default::default(),
859953
request_id: request_id.to_string(),
860954
routing_key: None,
861955
model_id: model_id.to_string(),
@@ -910,6 +1004,27 @@ mod tests {
9101004
}
9111005
}
9121006

1007+
#[test]
1008+
fn stale_guards_cannot_advance_or_remove_reused_request_ids() {
1009+
for replacement_model in ["model-a", "model-b"] {
1010+
let live = LiveRequestState::default();
1011+
let mut first = live.track_request(&required_for_model("same-id", "model-a", 0, 100));
1012+
let replacement =
1013+
live.track_request(&required_for_model("same-id", replacement_model, 1, 20));
1014+
let expected = live.snapshot_model(replacement_model);
1015+
first.on_backend_submission();
1016+
first.observe_output();
1017+
drop(first);
1018+
assert_eq!(live.snapshot_model(replacement_model), expected);
1019+
assert_eq!(expected.num_running_queries, 1);
1020+
drop(replacement);
1021+
assert_eq!(
1022+
live.snapshot_model(replacement_model).num_running_queries,
1023+
0
1024+
);
1025+
}
1026+
}
1027+
9131028
#[test]
9141029
fn one_live_request_transition_updates_queue_and_active_output_load() {
9151030
let live_requests = LiveRequestState::default();

0 commit comments

Comments
 (0)