1use super::{
4 Job, UsageStats,
5 http::error,
6 output::{Frame, Output},
7 pacer::ChunkPacer,
8 registry::{Ticket, tool_call_id},
9 request::{PreparedRequest, parse_request},
10 sessions::{Store, Turn},
11 tool_call::{FallbackReason, ToolCallOutputDecoder, WireToolCall},
12};
13use crate::units::BYTES_PER_MIB;
14use crate::{
15 control::{State, state::ActiveRequest},
16 options::Options,
17 prefix_cache::{Prefill, PrefixCache},
18 qwen4_exp::gpu::{Gpu, PrefixState},
19 runner::Decode,
20 tok::ChatTokenizer,
21};
22use anyhow::{Result, ensure};
23use serde_json::json;
24use std::{
25 collections::VecDeque,
26 sync::{Arc, Mutex, mpsc},
27 time::{Duration, Instant},
28};
29
30struct Pending {
31 job: Job,
32 prepared: PreparedRequest,
33 reservation: usize,
34}
35
36enum Phase {
37 Ready,
38 Prefill(Prefill),
39 Decode(Box<Decode>),
40}
41
42struct Active<'a> {
43 prepared: PreparedRequest,
44 phase: Phase,
45 checkpoint: Option<PrefixState>,
46 reservation: usize,
47 ticket: Arc<Ticket>,
48 session: Option<Turn>,
49 output: Output,
50 text_decoder: Box<dyn FnMut(u32) -> Result<Option<String>> + 'a>,
51 text: GeneratedText,
52 usage: UsageStats,
53 config_generation: u64,
54 response_bytes: usize,
55 last_chunk: Option<(usize, f64)>,
57}
58
59#[derive(Default)]
62struct GeneratedText {
63 raw: String,
64 content: String,
65 tool_decoder: Option<ToolCallOutputDecoder>,
66}
67
68impl GeneratedText {
69 fn feed(&mut self, delta: &str, output: &Output, limit: usize) -> Result<()> {
70 ensure!(
71 self.raw.len() + delta.len() <= limit,
72 "response exceeds response_bytes"
73 );
74 self.raw.push_str(delta);
75
76 let visible = match self.tool_decoder.as_mut() {
77 Some(decoder) => decoder.feed(delta),
78 None => delta.to_owned(),
79 };
80
81 self.publish(visible, output)
82 }
83
84 fn publish(&mut self, text: String, output: &Output) -> Result<()> {
85 if !text.is_empty() {
86 self.content.push_str(&text);
87 output.send(Frame::Text(text))?;
88 }
89
90 Ok(())
91 }
92}
93
94pub(super) struct Worker<'a> {
95 activity_snapshot: crate::qwen4_exp::gpu::ExpertActivity,
96 gpu: Gpu<'a>,
97 tok: &'a ChatTokenizer,
98 options: Options,
99 cache: PrefixCache,
100 state: Arc<State>,
101 sessions: Arc<Mutex<Store>>,
102 active: VecDeque<Active<'a>>,
103 pending: VecDeque<Pending>,
104 loaded: Option<String>,
105 reserved: usize,
106 pacer: ChunkPacer,
107}
108
109impl<'a> Worker<'a> {
110 pub(super) fn new(
111 gpu: Gpu<'a>,
112 tok: &'a ChatTokenizer,
113 mut options: Options,
114 cache: PrefixCache,
115 state: Arc<State>,
116 sessions: Arc<Mutex<Store>>,
117 ) -> Self {
118 if !gpu.has_mtp() {
119 options.drafts = 0;
120 }
121
122 let limits = &state.config().config.limits;
123 let pacer = ChunkPacer::new(limits.prefill_quantum, limits.prefill_chunk_seconds);
124
125 Self {
126 activity_snapshot: Default::default(),
127 gpu,
128 tok,
129 options,
130 cache,
131 state,
132 sessions,
133 active: VecDeque::new(),
134 pending: VecDeque::new(),
135 loaded: None,
136 reserved: 0,
137 pacer,
138 }
139 }
140
141 pub(super) fn run(mut self, receiver: mpsc::Receiver<Job>) {
142 loop {
143 if self.active.is_empty() && self.pending.is_empty() {
144 match receiver.recv_timeout(Duration::from_secs(1)) {
145 Ok(job) => self.prepare(job),
146 Err(mpsc::RecvTimeoutError::Disconnected) => return,
147 Err(mpsc::RecvTimeoutError::Timeout) => {}
148 }
149 }
150
151 for job in receiver.try_iter() {
152 self.prepare(job);
153 }
154
155 self.admit();
156 self.turn();
157 self.gpu.clear_profile();
158 self.cache.expire();
159 self.observe();
160 }
161 }
162
163 fn prepare(&mut self, mut job: Job) {
164 let result = parse_request(
165 &job.body,
166 job.kind,
167 &job.settings.config.defaults,
168 job.session.as_ref().map(|turn| &turn.input),
169 self.tok.template.as_ref(),
170 )
171 .and_then(|r| {
172 r.prepare(
173 self.tok,
174 &self.options,
175 job.settings.config.limits.max_output_tokens,
176 )
177 });
178 let prepared = match result {
179 Ok(p) => p,
180 Err(e) => {
181 let _ = error(&mut job.stream, 400, &e.to_string());
182
183 self.state.update(|s| {
184 s.queued_requests -= 1;
185 s.failed_requests += 1;
186 });
187
188 return;
189 }
190 };
191 let limits = &job.settings.config.limits;
192 let context =
193 prepared.ids.len() + prepared.request.max_tokens + prepared.options.effective_drafts();
194 let draft_context = if self.gpu.has_mtp() { context } else { 0 };
197 let checkpoint_bytes = 2 * self.gpu.state_bytes_at(context, draft_context);
198 let sampler_bytes = self.gpu.logits().len() * 24;
199 let token_bytes = context * 16;
200 let message_bytes = 4 * (limits.request_bytes + limits.response_bytes);
201 let reservation = checkpoint_bytes + sampler_bytes + token_bytes + message_bytes;
202
203 if reservation > limits.active_state_mib * BYTES_PER_MIB {
204 let _ = error(
205 &mut job.stream,
206 503,
207 "request exceeds active_state_mib; reduce token limits or increase the state budget",
208 );
209
210 self.state.update(|s| {
211 s.queued_requests -= 1;
212 s.rejected_requests += 1;
213 });
214
215 return;
216 }
217
218 job.body = serde_json::Value::Null;
219
220 self.pending.push_back(Pending {
221 job,
222 prepared,
223 reservation,
224 });
225 }
226
227 fn admit(&mut self) {
228 let mut waiting = VecDeque::new();
229
230 while let Some(mut pending) = self.pending.pop_front() {
231 if pending.job.ticket.cancelled() {
232 let _ = error(
233 &mut pending.job.stream,
234 499,
235 "request cancelled before admission",
236 );
237
238 self.state.update(|s| {
239 s.queued_requests -= 1;
240 s.cancelled_requests += 1;
241 });
242
243 continue;
244 }
245
246 let limits = &pending.job.settings.config.limits;
247
248 if self.active.len() >= limits.active_requests
249 || self.reserved + pending.reservation > limits.active_state_mib * BYTES_PER_MIB
250 {
251 waiting.push_back(pending);
252
253 continue;
254 }
255
256 self.state.begin();
257
258 self.reserved += pending.reservation;
259 let tool_decoder = pending
260 .prepared
261 .request
262 .tool_contract
263 .clone()
264 .map(|contract| {
265 ToolCallOutputDecoder::new(contract, true)
268 });
269 let mut decoder = self.tok.inner.decode_stream(false);
270 let output = Output::new(
271 pending.job.stream,
272 pending.job.kind,
273 &pending.prepared.request,
274 pending.job.ticket.clone(),
275 self.state.clone(),
276 );
277
278 self.active.push_back(Active {
279 prepared: pending.prepared,
280 phase: Phase::Ready,
281 checkpoint: None,
282 reservation: pending.reservation,
283 ticket: pending.job.ticket,
284 session: pending.job.session,
285 output,
286 text_decoder: Box::new(move |token| {
287 decoder
288 .step(token)
289 .map_err(|e| anyhow::anyhow!("decode: {e}"))
290 }),
291 text: GeneratedText {
292 tool_decoder,
293 ..Default::default()
294 },
295 usage: UsageStats::default(),
296 config_generation: pending.job.settings.generation,
297 response_bytes: limits.response_bytes,
298 last_chunk: None,
299 });
300 }
301
302 self.pending = waiting;
303
304 self.state
305 .update(|s| s.active_state_reserved_bytes = self.reserved);
306 }
307
308 fn activate(&mut self, request: &Active<'_>) -> Result<()> {
309 if self.loaded.as_deref() == Some(&request.ticket.id) {
310 return Ok(());
311 }
312
313 if let Some(previous) = self
314 .active
315 .iter_mut()
316 .find(|a| Some(&a.ticket.id) == self.loaded.as_ref())
317 {
318 self.gpu.save_into(&mut previous.checkpoint);
319 }
320
321 self.gpu.reset_request();
322
323 if let Some(checkpoint) = &request.checkpoint {
324 self.gpu.restore_prefix(checkpoint)?;
325 }
326
327 self.loaded = Some(request.ticket.id.clone());
328
329 Ok(())
330 }
331
332 fn turn(&mut self) {
333 let Some(mut request) = self.active.pop_front() else {
334 return;
335 };
336 let contended = !self.active.is_empty() || !self.pending.is_empty();
338 let quantum = self.pacer.quantum(contended);
339 let result = if request.ticket.cancelled() {
340 Ok(true)
341 } else {
342 self.activate(&request).and_then(|()| {
343 request.step(
344 &mut self.gpu,
345 &mut self.cache,
346 self.tok,
347 &self.state,
348 quantum,
349 )
350 })
351 };
352
353 if let Some((tokens, seconds)) = request.last_chunk.take() {
354 self.pacer.observe(tokens, seconds);
355 }
356
357 if matches!(result, Ok(false)) && !request.ticket.cancelled() {
358 self.active.push_back(request);
359
360 return;
361 }
362
363 self.reserved -= request.reservation;
364
365 if self.loaded.as_deref() == Some(&request.ticket.id) {
366 self.loaded = None;
367 }
368
369 let cancelled = request.ticket.cancelled();
370
371 if !cancelled {
372 request.report(&self.gpu);
373 }
374
375 let _ = match result {
377 Err(e) if !cancelled => request.output.send(Frame::Error(e.to_string())),
378 _ => request.finish(self.tok),
379 };
380
381 self.state.update(|s| {
382 s.active_requests -= 1;
383 s.active_state_reserved_bytes = self.reserved;
384 s.current = None;
385 });
386 }
387
388 fn observe(&mut self) {
389 self.state
390 .observe(&self.gpu, &self.cache, &mut self.activity_snapshot);
391
392 let mut sessions = self.sessions.lock().unwrap();
393
394 sessions.expire();
395 self.state.update(|s| {
396 s.active_state_reserved_bytes = self.reserved;
397 s.prefill_chunk_tokens = self.pacer.tokens();
398 s.sessions = sessions.stats();
399 s.active = self.active.iter().map(Active::stats).collect();
400 });
401 }
402}
403
404impl Active<'_> {
405 fn stats(&self) -> ActiveRequest {
406 ActiveRequest {
407 id: self.ticket.id.clone(),
408 session_id: self.session.as_ref().map(|s| s.id.clone()),
409 phase: self.phase_name(),
410 usage: self.usage,
411 reserved_state_bytes: self.reservation,
412 }
413 }
414
415 fn phase_name(&self) -> &'static str {
416 match self.phase {
417 Phase::Decode(_) => "decode",
418 _ => "prefill",
419 }
420 }
421
422 fn generated(&self) -> &[u32] {
423 match &self.phase {
424 Phase::Decode(d) => &d.tokens,
425 _ => &[],
426 }
427 }
428
429 fn report(&self, gpu: &Gpu) {
432 let prompt_tokens = self.prepared.ids.len();
433 let decode_tokens = self.generated().len();
434 let prefill_tokens = self
435 .usage
436 .prompt_tokens
437 .saturating_sub(self.usage.cached_tokens);
438 let prefill_tps = prefill_tokens as f64 / self.usage.prefill_seconds.max(1e-9);
439 let decode_tps = decode_tokens as f64 / self.usage.decode_seconds.max(1e-9);
440 let prefill_s = self.usage.prefill_seconds;
441 let decode_s = self.usage.decode_seconds;
442 let ctx_pos = gpu.pos;
443 let max_ctx = gpu.max_t.max(1);
444 let ctx_pct = gpu.context_fullness() * 100.0;
445 let kv_pct = gpu.kv_cache_fullness() * 100.0;
446 let kv_used_mb = gpu.kv_cache_bytes() as f64 / BYTES_PER_MIB as f64;
447 let kv_total_mb = gpu.kv_cache_capacity() as f64 / BYTES_PER_MIB as f64;
448
449 eprintln!(
450 "serve {id}: prefill {prompt_tokens} tok in {prefill_s:.2}s ({prefill_tps:.0} tok/s) \
451 | decode {decode_tokens} tok in {decode_s:.2}s ({decode_tps:.1} tok/s) \
452 | ctx {ctx_pos}/{max_ctx} ({ctx_pct:.0}% full) \
453 | kv cache {kv_used_mb:.0}/{kv_total_mb:.0} MB ({kv_pct:.0}% full)",
454 id = self.ticket.id,
455 );
456 }
457
458 fn step(
459 &mut self,
460 gpu: &mut Gpu<'_>,
461 cache: &mut PrefixCache,
462 tok: &ChatTokenizer,
463 state: &State,
464 quantum: usize,
465 ) -> Result<bool> {
466 state.current(
467 &self.ticket.id,
468 self.config_generation,
469 self.phase_name(),
470 self.generated().len(),
471 );
472
473 if matches!(self.phase, Phase::Ready) {
474 let prefill = cache.begin(
475 gpu,
476 &self.prepared.ids,
477 &self.prepared.boundaries,
478 &self.prepared.options,
479 )?;
480 self.usage.prompt_tokens = self.prepared.ids.len() as u64;
481 self.usage.cached_tokens = prefill.cached as u64;
482
483 state.update(|s| {
484 s.prompt_tokens += self.prepared.ids.len() as u64;
485 s.cached_tokens += prefill.cached as u64;
486 });
487
488 self.phase = Phase::Prefill(prefill);
489 }
490
491 let started = Instant::now();
492
493 if let Phase::Prefill(prefill) = &mut self.phase {
494 if self.ticket.cancelled() {
495 return Ok(false);
496 }
497
498 let result = prefill.advance(
499 gpu,
500 cache,
501 &self.prepared.ids,
502 &self.prepared.options,
503 quantum,
504 );
505 let elapsed = started.elapsed().as_secs_f64();
506 self.usage.prefill_seconds += elapsed;
507
508 state.update(|s| s.prefill_seconds += elapsed);
509
510 let progress = result?;
511 self.last_chunk = progress.pacing_tokens.map(|tokens| (tokens, elapsed));
512
513 if !progress.done || self.ticket.cancelled() {
514 return Ok(false);
515 }
516
517 let seed = std::mem::take(&mut prefill.seed);
518 let rng = self.session.as_ref().and_then(|s| s.rng.clone());
519 self.phase = Phase::Decode(Box::new(Decode::new(
520 seed,
521 &self.prepared.ids,
522 tok,
523 &self.prepared.options,
524 self.prepared.request.max_tokens,
525 rng,
526 )?));
527
528 return Ok(false);
529 }
530
531 let Phase::Decode(decoder) = &mut self.phase else {
532 unreachable!()
533 };
534
535 let result = decoder.step(gpu, None, None, &mut |token| {
536 ensure!(!self.ticket.cancelled(), "request cancelled");
537
538 if let Some(delta) = (self.text_decoder)(token)? {
539 self.text.feed(&delta, &self.output, self.response_bytes)?;
540 }
541
542 state.token();
543
544 Ok(())
545 });
546 let elapsed = started.elapsed().as_secs_f64();
547 self.usage.decode_seconds += elapsed;
548 self.usage.generated_tokens = decoder.tokens.len() as u64;
549
550 state.update(|s| s.decode_seconds += elapsed);
551
552 result?;
553
554 Ok(decoder.finish_reason.is_some())
555 }
556
557 fn finish(mut self, tok: &ChatTokenizer) -> Result<()> {
558 let result = self.finish_response(tok);
559
560 if let Err(e) = &result {
561 let _ = self.output.send(Frame::Error(e.to_string()));
562 }
563
564 result
565 }
566
567 fn finish_response(&mut self, tok: &ChatTokenizer) -> Result<()> {
568 let cancelled = self.ticket.cancelled();
569 let full = tok.decode(self.generated())?;
570
571 ensure!(
572 full.len() <= self.response_bytes,
573 "response exceeds response_bytes"
574 );
575
576 let tail = full
579 .strip_prefix(&self.text.raw)
580 .ok_or_else(|| anyhow::anyhow!("final decoded text differs from streamed output"))?;
581
582 self.text.feed(tail, &self.output, self.response_bytes)?;
583
584 let terminal = self
585 .text
586 .tool_decoder
587 .take()
588 .map(|decoder| decoder.finish());
589
590 let tool_calls: Option<Vec<WireToolCall>> = terminal
591 .as_ref()
592 .filter(|terminal| !terminal.tool_calls.is_empty())
593 .map(|terminal| {
594 terminal
595 .tool_calls
596 .iter()
597 .map(|call| WireToolCall {
598 id: tool_call_id(),
599 name: call.name.clone(),
600 arguments: call.arguments_json.clone(),
601 })
602 .collect()
603 });
604
605 if let (Some(terminal), Some(calls)) = (&terminal, &tool_calls) {
606 let d = &terminal.diagnostics;
607
608 if d.fallback_reason != FallbackReason::None {
609 eprintln!(
610 "serve {}: tool calls {}x (recovered: {:?})",
611 self.ticket.id,
612 calls.len(),
613 d.fallback_reason
614 );
615 } else if d.schema_mismatch_arguments > 0 || d.empty_arguments_omitted > 0 {
616 eprintln!(
617 "serve {}: tool calls {}x ({} schema mismatches, {} empty arguments omitted)",
618 self.ticket.id,
619 calls.len(),
620 d.schema_mismatch_arguments,
621 d.empty_arguments_omitted
622 );
623 }
624 }
625
626 if let Some(terminal) = &terminal {
629 self.text.publish(terminal.content.clone(), &self.output)?;
630 }
631
632 let (reason, rng) = match &self.phase {
633 Phase::Decode(d) => (d.finish_reason.unwrap_or("cancelled"), d.rng()),
634 _ => ("cancelled", None),
635 };
636 let reason = if tool_calls.is_some() && reason == "stop" {
640 "tool_calls"
641 } else {
642 reason
643 };
644 let usage = json!({
645 "prompt_tokens": self.prepared.ids.len(),
646 "prompt_tokens_details": {"cached_tokens": self.usage.cached_tokens},
647 "completion_tokens": self.generated().len(),
648 "total_tokens": self.prepared.ids.len() + self.generated().len(),
649 });
650 let turn = if cancelled {
651 None
652 } else {
653 self.session
654 .take()
655 .map(|t| t.prepare(&self.text.content, rng, self.usage, tool_calls.as_deref()))
656 .transpose()?
657 };
658
659 self.output.send(Frame::Finish {
660 text: std::mem::take(&mut self.text.content),
661 reason,
662 usage,
663 turn: turn.map(Box::new),
664 tool_calls,
665 })
666 }
667}
668
669#[cfg(test)]
670#[path = "../../tests/unit/server/worker.rs"]
671mod tests;