1use serde::{Deserialize, Serialize};
8use serde_json::Value;
9use uuid::Uuid;
10
11#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
18#[serde(rename_all = "snake_case")]
19pub enum ReplayPromptTokenSource {
20 #[default]
23 Materialized,
24 LengthOnlySynthetic,
27}
28
29#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
32pub struct ReplayRequestContext {
33 pub authored_id: String,
34 #[serde(default, skip_serializing_if = "Option::is_none")]
35 pub session_id: Option<String>,
36 #[serde(default, skip_serializing_if = "Option::is_none")]
37 pub turn_index: Option<usize>,
38 #[serde(default, skip_serializing_if = "Value::is_null")]
39 pub metadata: Value,
40 #[serde(default)]
41 pub prompt_token_source: ReplayPromptTokenSource,
42}
43
44#[doc(hidden)]
46#[derive(Debug, Clone, Default, Serialize, Deserialize)]
47pub struct DirectRequest {
48 pub tokens: Vec<u32>,
49 pub max_output_tokens: usize,
50 #[serde(default, skip_serializing_if = "Option::is_none")]
51 pub output_token_ids: Option<Vec<u32>>,
52 pub uuid: Option<Uuid>,
53 pub dp_rank: u32,
54 #[serde(default, skip_serializing_if = "Option::is_none")]
58 pub preferred_dp_rank: Option<u32>,
59 pub arrival_timestamp_ms: Option<f64>,
60 #[serde(default, skip_serializing_if = "is_zero_i32")]
61 pub priority: i32,
62 #[serde(default, skip_serializing_if = "is_zero_u32")]
63 pub strict_priority: u32,
64 #[serde(default, skip_serializing_if = "Option::is_none")]
65 pub policy_class: Option<String>,
66 #[serde(default, skip_serializing_if = "Option::is_none")]
69 pub replay_context: Option<ReplayRequestContext>,
70}
71
72fn is_zero_i32(value: &i32) -> bool {
73 *value == 0
74}
75
76fn is_zero_u32(value: &u32) -> bool {
77 *value == 0
78}
79
80impl DirectRequest {
81 pub fn request_id(&self) -> Option<Uuid> {
82 self.uuid
83 }
84
85 pub fn input_tokens(&self) -> &[u32] {
86 &self.tokens
87 }
88
89 pub fn max_output_tokens(&self) -> usize {
90 self.max_output_tokens
91 }
92
93 #[inline]
95 pub fn effective_max_output_tokens(&self) -> usize {
96 self.output_token_ids
97 .as_ref()
98 .map_or(self.max_output_tokens, Vec::len)
99 }
100
101 pub(crate) fn clone_with_output_limit(&self, limit: usize) -> Self {
102 let max_output_tokens = self.effective_max_output_tokens().min(limit);
103 let mut request = self.clone();
104 request.max_output_tokens = max_output_tokens;
105 request.output_token_ids = self
106 .output_token_ids
107 .as_ref()
108 .map(|ids| ids[..max_output_tokens].to_vec());
109 request
110 }
111
112 pub fn arrival_time_ms(&self) -> Option<f64> {
113 self.arrival_timestamp_ms
114 }
115
116 pub fn preferred_dp_rank(&self) -> Option<u32> {
117 self.preferred_dp_rank
118 }
119
120 pub fn router_priorities(&self) -> (f64, u32) {
121 (f64::from(self.priority.max(0)), self.strict_priority)
122 }
123
124 pub fn policy_class(&self) -> Option<&str> {
125 self.policy_class.as_deref()
126 }
127
128 pub fn prompt_tokens_are_placement_safe(&self) -> bool {
129 self.replay_context.as_ref().is_none_or(|context| {
130 context.prompt_token_source != ReplayPromptTokenSource::LengthOnlySynthetic
131 })
132 }
133}
134
135#[cfg(test)]
136mod tests {
137 use super::DirectRequest;
138
139 #[test]
140 fn output_plan_is_authoritative_when_limited() {
141 let request = DirectRequest {
142 tokens: vec![1, 2],
143 max_output_tokens: 0,
144 output_token_ids: Some(vec![7, 8]),
145 ..Default::default()
146 };
147
148 assert_eq!(request.effective_max_output_tokens(), 2);
149 let limited = request.clone_with_output_limit(1);
150 assert_eq!(limited.max_output_tokens, 1);
151 assert_eq!(limited.output_token_ids.as_deref(), Some(&[7][..]));
152 assert_eq!(limited.output_token_ids.unwrap().capacity(), 1);
153 assert_eq!(request.output_token_ids.as_deref(), Some(&[7, 8][..]));
154
155 let empty_plan = DirectRequest {
156 max_output_tokens: 4,
157 output_token_ids: Some(Vec::new()),
158 ..Default::default()
159 };
160 assert_eq!(empty_plan.effective_max_output_tokens(), 0);
161 let limited = empty_plan.clone_with_output_limit(1);
162 assert_eq!(limited.max_output_tokens, 0);
163 assert_eq!(limited.output_token_ids.as_deref(), Some(&[][..]));
164
165 let unplanned = DirectRequest::default();
166 assert_eq!(unplanned.effective_max_output_tokens(), 0);
167 let limited = unplanned.clone_with_output_limit(1);
168 assert_eq!(limited.max_output_tokens, 0);
169 assert!(limited.output_token_ids.is_none());
170 }
171}
172
173#[derive(Debug, Clone, Serialize, Deserialize)]
175pub(crate) struct OutputSignal {
176 pub(crate) uuid: Uuid,
177 #[serde(default, skip_serializing_if = "Option::is_none")]
178 pub(crate) token_id: Option<u32>,
179 pub(crate) completed: bool,
180 #[serde(default)]
181 pub(crate) rejected: bool,
182 #[serde(default, skip_serializing_if = "Option::is_none")]
183 pub(crate) handoff_delay_ms: Option<f64>,
184 #[serde(default, skip_serializing_if = "Option::is_none")]
187 pub(crate) cached_tokens: Option<usize>,
188}
189
190#[derive(Debug, Clone, Default, PartialEq)]
192pub struct ForwardPassSnapshot {
193 pub version: u32,
194 pub worker_id: String,
195 pub dp_rank: u32,
196 pub counter_id: u64,
197 pub num_prefill_requests: u32,
198 pub sum_prefill_tokens: u64,
199 pub var_prefill_length: f64,
200 pub sum_prefill_kv_tokens: u64,
201 pub num_decode_requests: u32,
202 pub sum_decode_kv_tokens: u64,
203 pub var_decode_kv_tokens: f64,
204 pub num_queued_prefill: u32,
205 pub sum_queued_prefill_tokens: u64,
206 pub var_queued_prefill_length: f64,
207 pub num_queued_decode: u32,
208 pub sum_queued_decode_kv_tokens: u64,
209 pub var_queued_decode_kv_tokens: f64,
210 pub wall_time_secs: f64,
211}