Skip to main content

cherenkov/server/
worker.rs

1//! Round-robin GPU scheduling. Each active request owns its decoding state.
2
3use 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    /// Tokens and elapsed seconds from the last step, if eligible for pacing.
56    last_chunk: Option<(usize, f64)>,
57}
58
59/// Raw decoded bytes are used for reconciliation and limits. Only parsed
60/// content is published to clients and retained in session history.
61#[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        // Vec growth can double initialized snapshot bytes. Reserve that upper
195        // bound, sampler workspace, tokens, and bounded request/response copies.
196        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                    // Recovery is always on: complete calls stand, and truncated
266                    // or suffixed output never poisons the visible text.
267                    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        // Both paced and lone-request chunks stay within the reserved quantum.
337        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        // The output writer accounts for completion and delivery failures.
376        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    /// Per-request throughput and context/KV-cache fullness, reported to
430    /// the server console when the request completes.
431    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        // Reconcile against raw bytes before closing the parser: the token
577        // stream may have withheld a final Unicode fragment or marker suffix.
578        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        // Fallback bytes and trailing whitespace take the same publication
627        // path as incremental content, including for streaming clients.
628        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        // Recovered calls get the tool_calls finish reason only when the model
637        // stopped on its own; a generation cut off by the token budget keeps
638        // its "length" reason so the client learns output was truncated.
639        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;