diff --git a/rust/sglang-server/src/detokenizer.rs b/rust/sglang-server/src/detokenizer.rs new file mode 100644 index 000000000..d0a800994 --- /dev/null +++ b/rust/sglang-server/src/detokenizer.rs @@ -0,0 +1,597 @@ +//! Detokenizer shards — CPU-bound, one pinned thread per shard. +//! +//! Each shard owns a *local* `rid -> DetokState` map. There is no lock: a given +//! rid is routed to exactly one shard (by `Rid::shard`) for both its +//! `Register` and all its `Chunk`s, so the map has a single accessor. +//! +//! The hash PARTITIONS, the rid IDENTIFIES. Keying the map by the hash meant two +//! distinct rids that happened to collide became one entry: `Register` evicted the +//! first client's sink and their tokens were then written to the second client's +//! connection. Chunks carry the rid string (moved out of the frame header, which +//! owns it and would otherwise drop it), so a collision now only co-locates. +//! +//! Real detokenization uses dynamo-tokenizers' `DecodeStream`, a stateful +//! incremental decoder (TGI/vLLM-style: it buffers partial UTF-8 / byte-fallback +//! tokens and only emits text once a valid boundary is reached). Each request +//! gets its own `DecodeStream`. When no tokenizer is configured (or +//! `skip_tokenizer_init` is set) the backend is `Skip`: no decoding, the raw +//! `output_ids` are emitted instead of text. +//! +//! Per-chunk egress flow (no FSM state change inside Streaming): +//! ChunkEvent{finish:None} -> step ids -> delta -> Server frame +//! ChunkEvent{finish:Some} -> step ids -> delta -> final frame + +use std::collections::HashMap; + +use crate::error::Error; +use crate::fsm::{Event, RequestState}; +use crate::ids::Rid; +use crate::message::DetokMsg; +use crate::message::{ChunkEvent, EgressItem, EgressSink, Matched, SinkError, TokenIds}; +use crate::runtime::Runnable; +use crate::tokenizer_manager::AbortSource; + +/// Default for `skip_special_tokens` (SGLang's SamplingParams default). The +/// per-request value isn't available on the egress side yet; see the note in +/// `DetokenizerBackend::new_decoder`. +const SKIP_SPECIAL_TOKENS: bool = true; + +/// Per-request incremental decoder. `step` feeds the new token ids for one chunk +/// and returns the newly decoded text delta (empty if the ids only produced a +/// partial/incomplete multi-byte sequence that needs more tokens). +pub trait StreamDecoder: Send { + fn step(&mut self, token_ids: &[i32]) -> Result; +} + +/// Real decoder wrapping a dynamo-tokenizers `DecodeStream`. +struct DynamoDecoder { + stream: dynamo_tokenizers::DecodeStream, +} + +impl StreamDecoder for DynamoDecoder { + fn step(&mut self, token_ids: &[i32]) -> Result { + let mut out = String::new(); + for &id in token_ids { + if let Some(chunk) = self + .stream + .step(id as u32) + .map_err(|e| Error::Detokenize(e.to_string()))? + { + out.push_str(&chunk); + } + } + Ok(out) + } +} + +/// Shard-wide detok backend. Cloned per shard; mints a fresh per-request decoder +/// on each `Register`. +#[derive(Clone)] +pub enum DetokenizerBackend { + Dynamo(dynamo_tokenizers::Tokenizer), + /// No decoding at all — the shard emits each chunk's raw output token ids as + /// `output_ids` (no `DecodeStream`, no accumulation). Used for + /// `skip_tokenizer_init` and when no tokenizer is configured. + Skip, +} + +impl DetokenizerBackend { + /// Mint a per-request decoder, or `None` in skip mode (the shard passes the + /// token ids through untouched instead of decoding text). + fn new_decoder(&self) -> Option> { + match self { + // NOTE: the stream is seeded with an empty prompt context, which is + // correct for the common case. Seeding with the prompt's trailing + // tokens (for perfect first-token spacing) would require Register to + // carry input_ids — deferred. + DetokenizerBackend::Dynamo(t) => Some(Box::new(DynamoDecoder { + stream: t.decode_stream(&[], SKIP_SPECIAL_TOKENS), + })), + DetokenizerBackend::Skip => None, + } + } + + /// Decode each logprob token id to its own text (one id at a time, matching + /// Python's `batch_decode([[id] for id in ids])`). Runs on this CPU-bound + /// shard, not the api-server I/O threads. `Skip` mode (no tokenizer) yields + /// no text, so the `[logprob, token_id, text]` tuple's text slot stays null. + fn decode_logprob_texts(&self, idxs: &[i32]) -> Vec { + match self { + DetokenizerBackend::Dynamo(t) => idxs + .iter() + .map(|&id| { + t.decode(&[id as u32], false) + .map(String::from) + .unwrap_or_default() + }) + .collect(), + DetokenizerBackend::Skip => Vec::new(), + } + } +} + +struct DetokState { + sink: EgressSink, + /// `return_text_in_logprobs`: whether to decode this request's logprob token + /// ids to text (in this shard) for the `[logprob, token_id, text]` tuples. + decode_logprob_text: bool, + /// `SamplingParams.no_stop_trim`: keep the matched stop in the output. Default + /// (`false`) trims it off the final chunk (see [`trim_stop_str`]). + no_stop_trim: bool, + /// Per-request incremental decoder; `None` in `skip_tokenizer_init` mode. + /// This is the *only* per-request accumulation the shard keeps: the decoder's + /// internal byte/UTF-8 buffer. Decoded **text deltas** are emitted per chunk + /// (no cumulative buffer here) — the api-server's drain loop reassembles the + /// cumulative view where a consumer needs it (every unary response and the + /// cumulative SGLang `/generate` stream); OpenAI streaming forwards deltas. + decoder: Option>, + /// Egress half of the lifecycle FSM. Lives here because the ingress + /// `Request` (and its FSM) was handed to the scheduler when queued; the + /// shard is the sole owner of the request's egress state, so no lock. + fsm: RequestState, +} + +/// One detokenizer shard: owns a *local* `rid -> DetokState` map (single accessor, +/// no lock) and the egress backend. Spawned (pinned) per shard as a [`Runnable`]; +/// a given rid is routed to exactly one shard. +pub struct DetokenizerWorker { + shard: usize, + rx: flume::Receiver, + backend: DetokenizerBackend, + /// Unbounded abort lane, used to abort a request the shard had to drop + /// (client backpressure) so the scheduler stops generating for it. + abort: flume::Sender, +} + +impl DetokenizerWorker { + pub fn new( + shard: usize, + rx: flume::Receiver, + backend: DetokenizerBackend, + abort: flume::Sender, + ) -> Self { + Self { + shard, + rx, + backend, + abort, + } + } +} + +impl Runnable for DetokenizerWorker { + fn run(self) { + let mut table: HashMap = HashMap::new(); + tracing::debug!(shard = self.shard, "detokenizer worker started"); + + // Plain `recv`: exits when the `DetokMsg` channel closes (every `Senders` + // clone gone). On shutdown that happens once the API runtime drop cancels + // in-flight handlers (their `AbortGuard`s release the last clones) and + // tm-ingress/tm-egress exit — no shutdown signal needed here. + while let Ok(msg) = self.rx.recv() { + match msg { + DetokMsg::Register { + rid, + sink, + decode_logprob_text, + no_stop_trim, + } => { + table.insert( + rid.clone(), + DetokState { + sink, + decode_logprob_text, + no_stop_trim, + decoder: self.backend.new_decoder(), + // Registered == handed to the scheduler == Queued. + fsm: RequestState::Queued, + }, + ); + } + // One decode step's chunks for this shard, batched by tm-egress. + DetokMsg::Chunks(evs) => { + for ev in evs { + handle_chunk(&mut table, ev, &self.backend, &self.abort); + } + } + DetokMsg::Result { rid, payload } => handle_result(&mut table, &rid, payload), + DetokMsg::Fail { rid, message } => { + handle_fail(&mut table, &rid, message, &self.abort) + } + DetokMsg::Deregister { rid } => { + table.remove(&rid); + } + } + } + } +} + +/// Control-request result: deliver the JSON payload to the sink verbatim as a +/// single `Done` frame — no detokenization, no streaming. +fn handle_result(table: &mut HashMap, rid: &Rid, payload: bytes::Bytes) { + if let Some(mut st) = table.remove(rid) { + let _ = st.sink.try_send(EgressItem::Control(payload)); + // Egress FSM: a control request goes straight to Completed (no Streaming + // / Finalizing states — single response, never streamed). + st.fsm = RequestState::Completed; + } +} + +/// Terminal per-request failure (bad request header): send an `Error` to the sink +/// (the api-server turns it into an HTTP 400) and drop the request. +/// Terminal per-request failure. `Internal` (500), not `Validation` (400): the +/// producers of this message are server faults — a malformed scheduler output +/// frame — not bad client input. Also aborts the request on the scheduler, which +/// otherwise keeps generating tokens for a connection that will never read them. +fn handle_fail( + table: &mut HashMap, + rid: &Rid, + message: String, + abort: &flume::Sender, +) { + if let Some(mut st) = table.remove(rid) { + // Abort first: `try_send` on the sink can release the handler, which frees + // the rid for reuse (same ordering hazard as the disconnect path). + let _ = abort.send(AbortSource::Detok(rid.clone())); + let _ = st + .sink + .try_send(EgressItem::Error(Error::Internal(message))); + st.fsm = RequestState::Completed; + } +} + +fn handle_chunk( + table: &mut HashMap, + mut ev: ChunkEvent, + backend: &DetokenizerBackend, + abort: &flume::Sender, +) { + // Copied once: `ev` is moved into the sink below, but the rid is still + // needed to look the request up and to remove it. + let rid = ev.rid.clone(); + + let Some(st) = table.get_mut(&rid) else { + // Late chunk after completion/abort — drop. + return; + }; + let decode_logprob_text = st.decode_logprob_text; + let no_stop_trim = st.no_stop_trim; + + // Queued → Streaming on the first chunk (the scheduler picked it). + if matches!(st.fsm, RequestState::Queued) { + let _ = st.fsm.apply(Event::SchedulerPicked); + } + + let finished = ev.finish_reason.is_some(); + // Matched-stop trim (Python `trim_matched_stop`): the final chunk's finish + // reason names the stop it matched — a stop STRING or a stop TOKEN id. By + // default that stop is removed from the output; `no_stop_trim` keeps it. + let matched = finished + .then(|| ev.finish_reason.as_ref().and_then(|fr| fr.matched())) + .flatten() + .cloned(); + // Count generated tokens (incl. a matched stop token) *before* trimming. + let n_tok = ev.token_ids.len() as u64; + // Stop TOKEN: drop it before decode, so it reaches neither `text` nor `output_ids`. + trim_stop_token(&mut ev.token_ids, &matched, no_stop_trim); + + // Fully incremental: decode just this chunk's delta. `token_ids` stays in the + // event — it's ALSO surfaced as the `/generate` response's `output_ids` (the + // Python server returns them by default alongside `text`), in both normal and + // `skip_tokenizer_init` mode. Nothing cumulative is kept here — the api-server's + // drain loop reassembles it where needed. + let mut delta_text = match &mut st.decoder { + Some(decoder) => match decoder.step(&ev.token_ids) { + Ok(delta) => delta, + Err(e) => { + // Abort too: this is terminal for the request, and without it the + // scheduler keeps generating for a connection that is already gone + // — the other two terminal paths (disconnect, fail) both abort. + let _ = st.fsm.apply(Event::Error(e.clone())); + let _ = abort.send(AbortSource::Detok(rid.clone())); + let _ = st.sink.try_send(EgressItem::Error(e)); + table.remove(&rid); + return; + } + }, + // skip_tokenizer_init: no decode; the token ids pass through in `ev`. + None => String::new(), + }; + // Stop STRING: trim it (and anything after) from the decoded delta's tail. + if let Some(Matched::Str(stop)) = &matched { + trim_stop_str(&mut delta_text, stop, no_stop_trim); + } + + // Streaming → Streaming (finish:false) or Streaming → Finalizing (finish:true). + let _ = st.fsm.apply(Event::Chunk { finish: finished }); + + // `return_text_in_logprobs`: decode each logprob token id to text HERE (this + // CPU-bound shard) rather than on the api-server I/O threads. Flat text columns + // stay parallel to the `idx` buffers, so `sglang_frame` just reads them. Only the + // logprob-carrying frames have an `extras` box; a plain token frame skips this. + if decode_logprob_text && let Some(ex) = ev.extras.as_deref_mut() { + ex.out_lp_txt = backend.decode_logprob_texts(&ex.out_lp_idx); + ex.in_lp_txt = backend.decode_logprob_texts(&ex.in_lp_idx); + ex.out_top_txt = backend.decode_logprob_texts(&ex.out_top_idx); + ex.in_top_txt = backend.decode_logprob_texts(&ex.in_top_idx); + ex.out_tid_txt = backend.decode_logprob_texts(&ex.out_tid_idx); + ex.in_tid_txt = backend.decode_logprob_texts(&ex.in_tid_idx); + } + + // Fill the decode outputs in place; the pre-decode columns (boxed logprobs/hidden, + // token_ids, prompt_tokens, finish_reason) already ride in `ev`. The API handler + // formats this delta (and accumulates for the cumulative view). + ev.text = delta_text; + ev.completion_tokens = n_tok; + + if finished { + // The Done frame *is* the final frame: Finalizing → Completed. + let sent = st.sink.try_send(EgressItem::Done(ev)).is_ok(); + let _ = st.fsm.apply(if sent { + Event::FinalFrameSent + } else { + Event::Disconnect + }); + table.remove(&rid); + } else { + // Every intermediate chunk emits its delta frame. A failed send means the + // client can't receive it — `Closed` (gone) or `Full` (backpressure: not + // reading fast enough). Either way we can't buffer unboundedly, and + // silently dropping the frame would truncate the response and still look + // like success at EOS. So treat both as terminal: drop the request AND + // abort scheduler work for it. + if let Err(e) = st.sink.try_send(EgressItem::Frame(ev)) { + match e { + SinkError::Full => { + tracing::warn!( + rid = %rid, + "detok: sink full; aborting (client backpressure)" + ) + } + SinkError::Closed => { + tracing::debug!(rid = %rid, "detok: sink closed; aborting (client gone)") + } + } + let _ = st.fsm.apply(Event::Disconnect); + // Abort ONLY when the sink is full. `Closed` means the handler future is + // already gone, so its `AbortGuard` has run: it aborted and released the + // rid. A second abort from here is unordered with respect to that + // release, so it lands after a resubmit of the same rid has registered + // and deregisters the NEW request — the cross-wiring the rid registry + // exists to prevent, reached through the one abort producer that + // bypasses the guard's ordering. + if matches!(e, SinkError::Full) { + let _ = abort.send(AbortSource::Detok(rid.clone())); + } + table.remove(&rid); + } + } +} + +/// Drop a matched stop TOKEN from the final chunk (Python `trim_matched_stop`, +/// token branch); `no_stop_trim` / non-token match keeps it. +fn trim_stop_token(token_ids: &mut TokenIds, matched: &Option, no_stop_trim: bool) { + // Token id 0 is NOT a match: Python guards with `if not matched`, and 0 is + // falsy there, so it trims nothing. Trimming on 0 drops a real generated token + // for any model whose stop id happens to be 0. + if !no_stop_trim && matches!(matched, Some(Matched::Token(t)) if *t != 0) { + token_ids.pop(); + } +} + +/// Remove the matched stop string from the decoded final chunk (Python +/// `trim_matched_stop`, string branch); `no_stop_trim` keeps it. Truncates at +/// the FIRST occurrence. +fn trim_stop_str(text: &mut String, stop: &str, no_stop_trim: bool) { + if stop.is_empty() { + return; + } + if let Some(pos) = text.find(stop) { + text.truncate(if no_stop_trim { pos + stop.len() } else { pos }); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use tokio::sync::mpsc; + + /// A non-terminal chunk that can't be delivered (sink full → client + /// backpressure) drops the request AND aborts scheduler work — it does not + /// silently keep state, which would later read as a clean completion at EOS. + #[test] + fn full_sink_drops_request_and_aborts_scheduler() { + // Capacity-1 sink, pre-filled so the next send hits `Full`. + let (tx, _rx) = mpsc::channel::(1); + tx.try_send(EgressItem::Frame(ChunkEvent::default())) + .unwrap(); + + let mut table = HashMap::new(); + table.insert( + Rid::from("1"), + DetokState { + sink: EgressSink::Local(tx), + decode_logprob_text: false, + no_stop_trim: false, + decoder: None, + fsm: RequestState::Queued, + }, + ); + + let (tm_tx, tm_rx) = flume::unbounded::(); + let ev = ChunkEvent { + rid: Rid::from("1"), + token_ids: vec![5], + ..Default::default() // finish_reason None → non-terminal + }; + handle_chunk(&mut table, ev, &DetokenizerBackend::Skip, &tm_tx); + + // Request removed (no lingering state to be mistaken for success)... + assert!(!table.contains_key(&Rid::from("1"))); + // ...and the scheduler was told to abort it. + assert!(matches!( + tm_rx.try_recv(), + Ok(AbortSource::Detok(rid)) if rid == Rid::from("1") + )); + } + + /// `trim_stop_str` reproduces the base's stop-string semantics: `stop: "3"` on + /// output " 1, 2, 3" yields " 1, 2, " by default and " 1, 2, 3" with + /// `no_stop_trim`. + #[test] + fn trim_stop_str_matches_base() { + let mut t = " 1, 2, 3".to_string(); + trim_stop_str(&mut t, "3", false); + assert_eq!(t, " 1, 2, "); + + let mut t = " 1, 2, 3".to_string(); + trim_stop_str(&mut t, "3", true); + assert_eq!(t, " 1, 2, 3"); + + // Empty / absent stop is a no-op. + let mut t = "abc".to_string(); + trim_stop_str(&mut t, "", false); + assert_eq!(t, "abc"); + + // The stop can occur twice in the final chunk + let mut t = "a STOP b STOP".to_string(); + trim_stop_str(&mut t, "STOP", false); + assert_eq!(t, "a "); + + let mut t = "a STOP b STOP".to_string(); + trim_stop_str(&mut t, "STOP", true); + assert_eq!(t, "a STOP"); + } + + /// Two requests on the SAME shard keep separate entries. This is what a + /// A shard-hash collision now degrades to: the hash partitions, the rid + /// identifies. Keying the table by the hash made colliding rids one entry, so + /// `Register` evicted the first client's sink and their tokens were written to + /// the second client's connection. A single shard forces co-location + /// deterministically, without needing to find a real 64-bit collision. + #[test] + fn co_located_requests_keep_their_own_sinks() { + let (tx_a, mut rx_a) = mpsc::channel::(4); + let (tx_b, mut rx_b) = mpsc::channel::(4); + let mut table = HashMap::new(); + let state = |tx| DetokState { + sink: EgressSink::Local(tx), + decode_logprob_text: false, + no_stop_trim: false, + decoder: None, + fsm: RequestState::Queued, + }; + table.insert(Rid::from("alice"), state(tx_a)); + table.insert(Rid::from("bob"), state(tx_b)); + let (tm_tx, _tm_rx) = flume::unbounded::(); + + let chunk = |rid: &str, id: i32| ChunkEvent { + rid: Rid::from(rid.to_string()), + token_ids: vec![id], + ..Default::default() + }; + handle_chunk( + &mut table, + chunk("alice", 11), + &DetokenizerBackend::Skip, + &tm_tx, + ); + handle_chunk( + &mut table, + chunk("bob", 22), + &DetokenizerBackend::Skip, + &tm_tx, + ); + + let ids = |rx: &mut mpsc::Receiver| match rx.try_recv() { + Ok(EgressItem::Frame(ev)) => ev.token_ids, + other => panic!("expected a frame, got {other:?}"), + }; + assert_eq!( + ids(&mut rx_a), + vec![11], + "alice must not receive bob's tokens" + ); + assert_eq!( + ids(&mut rx_b), + vec![22], + "bob must not receive alice's tokens" + ); + assert_eq!(table.len(), 2, "neither registration evicted the other"); + } + + /// Drive a final (`finish_reason`) chunk through `handle_chunk` in skip mode and + /// return the emitted `Done` event. + fn final_chunk( + no_stop_trim: bool, + finish_reason: serde_json::Value, + ids: Vec, + ) -> ChunkEvent { + let (tx, mut rx) = mpsc::channel::(4); + let mut table = HashMap::new(); + table.insert( + Rid::from("1"), + DetokState { + sink: EgressSink::Local(tx), + decode_logprob_text: false, + no_stop_trim, + decoder: None, // skip mode → output_ids passthrough + fsm: RequestState::Queued, + }, + ); + let (tm_tx, _tm_rx) = flume::unbounded::(); + let ev = ChunkEvent { + rid: Rid::from("1"), + token_ids: ids, + // Parsed from the wire map, so the trim paths are driven by the same + // shape Python emits rather than a hand-built enum. + finish_reason: Some( + serde_json::from_value(finish_reason).expect("finish reason must parse"), + ), + ..Default::default() + }; + handle_chunk(&mut table, ev, &DetokenizerBackend::Skip, &tm_tx); + match rx.try_recv() { + Ok(EgressItem::Done(out)) => out, + other => panic!("expected Done, got {other:?}"), + } + } + + /// Token id 0 is not a match: Python's `trim_matched_stop` guards with + /// `if not matched`, and 0 is falsy there. Trimming on it drops a real + /// generated token for any model whose stop id is 0. + #[test] + fn matched_token_zero_does_not_trim() { + let mut ids = vec![1, 2, 0]; + trim_stop_token(&mut ids, &Some(Matched::Token(0)), false); + assert_eq!(ids, vec![1, 2, 0], "id 0 is not a matched stop"); + // A real stop id still trims. + let mut ids = vec![1, 2, 3]; + trim_stop_token(&mut ids, &Some(Matched::Token(3)), false); + assert_eq!(ids, vec![1, 2]); + } + + /// A matched stop TOKEN is dropped from the surfaced `output_ids` by default + /// (but still counted in `completion_tokens`); `no_stop_trim` keeps it. + #[test] + fn stop_token_trimmed_from_output_ids() { + let fr = serde_json::json!({ "type": "stop", "matched": 3 }); + let out = final_chunk(false, fr.clone(), vec![1, 2, 3]); + assert_eq!(out.token_ids, vec![1, 2], "matched stop token dropped"); + assert_eq!( + out.completion_tokens, 3, + "generated count still includes it" + ); + + let out = final_chunk(true, fr, vec![1, 2, 3]); + assert_eq!(out.token_ids, vec![1, 2, 3], "no_stop_trim keeps it"); + } + + /// A non-stop finish (`length`, no `matched`) never trims. + #[test] + fn length_finish_keeps_all_tokens() { + let fr = serde_json::json!({ "type": "length", "length": 3 }); + let out = final_chunk(false, fr, vec![1, 2, 3]); + assert_eq!(out.token_ids, vec![1, 2, 3]); + } +} diff --git a/rust/sglang-server/src/tokenizer.rs b/rust/sglang-server/src/tokenizer.rs new file mode 100644 index 000000000..ee90d605a --- /dev/null +++ b/rust/sglang-server/src/tokenizer.rs @@ -0,0 +1,162 @@ +//! Tokenizer pool — CPU-bound, runs on pinned OS threads (off the async +//! executor). Each worker pulls a `Request` from the shared `flume` receiver, +//! fills `input_ids`, and moves the request back to the TokenizerManager inbox. +//! +//! The text→ids step is behind [`TextTokenizer`], implemented by +//! [`DynamoTokenizer`] (dynamo-tokenizers: HuggingFace / tiktoken / fastokens). +//! A non-skip server requires a real tokenizer (enforced at startup); under +//! `skip_tokenizer_init` the pool isn't spawned at all. +//! +//! Mirrors the Python `_tokenize_one_request` text path: when the request +//! already carries `input_ids` it skips tokenization (handled upstream in the +//! TokenizerManager `classify`); otherwise the prompt text is encoded here. + +use std::path::Path; +use std::sync::Arc; + +use crate::error::Error; +use crate::fsm::Event; +use crate::message::{Request, RequestKind, TokenIds}; +use crate::runtime::Runnable; +use crate::tokenizer_manager::TmEvent; + +/// Pluggable text→token-ids backend. `Send + Sync` so one instance is shared +/// (read-only) across all pinned workers. +pub trait TextTokenizer: Send + Sync { + fn encode(&self, text: &str) -> Result; +} + +/// Load the tokenizer shared (Arc-backed) by the encode pool and detok shards. +/// `None` under `skip_tokenizer_init`, else required (missing/failed load → `Err`). +/// `tokenizer_path` is a tokenizer file, a model dir, or an HF Hub repo id +/// (resolved from the local cache — no network). +pub fn load_tokenizer( + tokenizer_path: Option<&str>, + revision: Option<&str>, + skip_tokenizer_init: bool, +) -> Result, String> { + if skip_tokenizer_init { + tracing::info!("skip_tokenizer_init: token ids in and out; no tokenizer/detokenizer"); + return Ok(None); + } + let path = tokenizer_path.ok_or_else(|| { + "no tokenizer configured: set tokenizer_path or enable skip_tokenizer_init".to_string() + })?; + + let file = resolve_model_file(path, revision, "tokenizer.json") + .ok_or_else(|| format!("tokenizer.json not found for '{path}'"))?; + let tokenizer = dynamo_tokenizers::Tokenizer::from_file_with_options( + &file, + dynamo_tokenizers::TokenizerOptions { + add_special_tokens: true, + }, + ) + .map_err(|e| format!("tokenizer load failed ({file}): {e}"))?; + tracing::info!(%path, "loaded tokenizer"); + Ok(Some(tokenizer)) +} + +/// Resolve a model file from the tokenizer source: a dir → `dir/`, a file → +/// its sibling, else an HF Hub repo id → the local cache. `None` if not found. +pub fn resolve_model_file(path: &str, revision: Option<&str>, filename: &str) -> Option { + let p = Path::new(path); + if p.is_dir() { + let f = p.join(filename); + return f.is_file().then(|| f.to_string_lossy().into_owned()); + } + if p.is_file() { + // `path` is a file (e.g. `tokenizer.json`); look for the sibling. + let f = p.parent()?.join(filename); + return f.is_file().then(|| f.to_string_lossy().into_owned()); + } + // Not a local path → HF Hub repo id (offline cache lookup). + resolve_from_hub_cache(path, revision, filename) +} + +/// Locate a file for an HF Hub repo id in the local cache (`HF_HOME`). Offline — +/// the scheduler pre-downloads the model. `None` if not cached. +fn resolve_from_hub_cache(repo_id: &str, revision: Option<&str>, filename: &str) -> Option { + use hf_hub::{Cache, Repo, RepoType}; + + let rev = revision.unwrap_or("main"); + Cache::from_env() + .repo(Repo::with_revision( + repo_id.to_string(), + RepoType::Model, + rev.to_string(), + )) + .get(filename) + .map(|p| p.to_string_lossy().into_owned()) +} + +/// Real tokenizer over an already-loaded dynamo `Tokenizer` (Arc inside). +pub struct DynamoTokenizer { + inner: dynamo_tokenizers::Tokenizer, +} + +impl DynamoTokenizer { + pub fn new(inner: dynamo_tokenizers::Tokenizer) -> Self { + Self { inner } + } +} + +impl TextTokenizer for DynamoTokenizer { + fn encode(&self, text: &str) -> Result { + if text.is_empty() { + // Match Python sglang: reject an empty prompt as a 400 (`Validation`), + // not the misleading 500 a tokenize error would give. + return Err(Error::Validation("prompt cannot be empty".into())); + } + let encoding = self + .inner + .encode(text) + .map_err(|e| Error::Tokenize(e.to_string()))?; + // Vocab ids are non-negative and fit in i32. + Ok(encoding.token_ids().iter().map(|&id| id as i32).collect()) + } +} + +/// One tokenizer worker: pulls a `Request` off the shared inbox, fills +/// `input_ids`, returns it to the TokenizerManager. Pinned; backend shared. +pub struct TokenizerWorker { + rx: flume::Receiver, + tm: flume::Sender, + tokenizer: Arc, +} + +impl TokenizerWorker { + pub fn new( + rx: flume::Receiver, + tm: flume::Sender, + tokenizer: Arc, + ) -> Self { + Self { rx, tm, tokenizer } + } +} + +impl Runnable for TokenizerWorker { + fn run(self) { + while let Ok(mut req) = self.rx.recv() { + // The tokenizer pool only ever receives generate requests. Encode, + // then advance the FSM: `TokenizeDone` on success (→ PreSendValidating). + let event = { + let RequestKind::Generate(g) = &mut req.kind else { + tracing::error!("tokenizer pool received a non-generate request"); + continue; + }; + match self.tokenizer.encode(g.text.as_deref().unwrap_or("")) { + Ok(ids) => { + g.input_ids = Some(ids); + Event::TokenizeDone + } + Err(err) => Event::Error(err), + } + }; + let _ = req.state.apply(event); + if self.tm.send(TmEvent::Tokenized(req)).is_err() { + tracing::error!("tm inbox closed; dropping request"); + break; + } + } + } +} diff --git a/rust/sglang-server/src/tokenizer_manager/egress.rs b/rust/sglang-server/src/tokenizer_manager/egress.rs new file mode 100644 index 000000000..1ad6df0ec --- /dev/null +++ b/rust/sglang-server/src/tokenizer_manager/egress.rs @@ -0,0 +1,217 @@ +//! TokenizerManager egress thread — drains the egress ring (scheduler output +//! pushed from Python) and routes each message to the detok shard that owns its +//! `Rid::shard`. Routing is a pure function of the rid, so it matches the shard +//! the request registered with on ingress — no shared map, no lock. +//! +//! The ring carries a 1-byte frame tag: `BATCH` (a whole decode batch, fanned +//! out here into per-request chunks), `RESULT` (a single control-request JSON +//! payload, e.g. `/server_info`), or `ERROR` (a terminal per-request failure the +//! scheduler ingress couldn't decode, routed back as a 400). + +use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; + +use bytes::Bytes; + +use crate::ids::Rid; +use crate::message::DetokMsg; +use crate::message::{ + ChunkEvent, EGRESS_TAG_BATCH, EGRESS_TAG_ERROR, EGRESS_TAG_RESULT, for_each_chunk, +}; +use crate::ring::EgressConsumer; +use crate::runtime::Runnable; +use crate::tokenizer_manager::{Senders, recv}; + +/// A monotonic counter bumped once per egress-ring frame the dispatcher drains. +/// It's the rust-native equivalent of the Python `TokenizerManager`'s +/// `last_receive_tstamp`: `/health_generate` watches it advance to confirm the +/// scheduler → detok path is alive (the value itself is meaningless). +pub type ActivityCounter = Arc; + +/// Egress dispatcher stage. Owns the egress-ring consumer + the detok-shard +/// senders, so the runtime spawns it as a [`Runnable`]. +pub struct Egress { + egress: EgressConsumer, + senders: Senders, + activity: ActivityCounter, + shutdown: flume::Receiver<()>, +} + +impl Egress { + pub fn new( + egress: EgressConsumer, + senders: Senders, + activity: ActivityCounter, + shutdown: flume::Receiver<()>, + ) -> Self { + Self { + egress, + senders, + activity, + shutdown, + } + } +} + +impl Runnable for Egress { + fn run(self) { + // Reused across frames (`clear` keeps capacity) — steady state allocates nothing. + let shards = self.senders.detok.len(); + let mut buckets: Vec> = (0..shards).map(|_| Vec::new()).collect(); + + while let Some(bytes) = recv(self.egress.receiver(), &self.shutdown) { + let Some((&tag, body)) = bytes.split_first() else { + continue; + }; + match tag { + // A whole decode batch: bucket each request by the shard owning its + // rid, then hand each shard its chunks in one send. + EGRESS_TAG_BATCH => { + for b in buckets.iter_mut() { + b.clear(); + } + let decoded = for_each_chunk(body, |ev| { + // The rid picks the shard; a hash collision only co-locates two + // requests now, it no longer merges them. + buckets[ev.rid.shard(shards)].push(ev); + }); + // Routing only fills the buckets; nothing is delivered until the + // sends below. So dropping them here makes rejection atomic for + // free: a frame whose columns drifted would otherwise deliver the + // requests decoded before the bad one, carrying another request's + // logprobs — the corruption the decoder's bounds checks exist to + // prevent. Better a lost frame than a silently wrong one. + if !decoded.ok { + // Dropping the frame keeps wrong data off the wire, but a + // request whose chunk was in it would otherwise wait forever: + // mid-stream it gets a hole, and if its FINAL chunk was here + // it never sees `Done` and the connection hangs — there is no + // server-side timeout. Fail from the HEADER's rids, not the + // buckets: a frame that fails at request 0 buckets nothing, + // so bucket-driven cleanup would leave every request in it + // hanging. + // Distinguish "failed N requests" from "named nobody": a + // frame whose header would not decode at all yields no rids, + // so nothing downstream fails and every request in it waits + // forever. That is the case worth paging on, and it used to + // log the same line as the recoverable one. + if decoded.rids.is_empty() { + tracing::error!( + "egress: bad batch frame named NO rids; any request in \ + it will hang (header undecodable, or empty rid column)" + ); + } else { + tracing::warn!( + rids = decoded.rids.len(), + "egress: bad batch frame; failing its requests" + ); + } + for b in buckets.iter_mut() { + b.clear(); + } + for rid in decoded.rids { + // 500, not 400: the client's request was fine — the + // scheduler's own output frame was not. + let shard = rid.shard(shards); + let _ = self.senders.detok[shard].send(DetokMsg::Fail { + rid, + message: "internal error: malformed scheduler output frame".into(), + }); + } + continue; + } + for (i, b) in buckets.iter_mut().enumerate() { + if b.is_empty() { + continue; + } + let chunks = DetokMsg::Chunks(std::mem::take(b)); + if self.senders.detok[i].send(chunks).is_err() { + tracing::error!("egress: detok shard closed"); + } + } + // Any frame off the ring = the scheduler produced output → alive. + self.activity.fetch_add(1, Ordering::Relaxed); + } + EGRESS_TAG_RESULT => { + if let Some((rid, msg)) = decode_result(body) { + self.route(&rid, msg); + } + } + EGRESS_TAG_ERROR => { + if let Some((rid, msg)) = decode_error(body) { + self.route(&rid, msg); + } + } + other => tracing::warn!(tag = other, "egress: unknown frame tag"), + } + } + } +} + +impl Egress { + /// Route one message to the shard owning `rid`. HOL ceiling: a slow shard stalls + /// this thread; the fix is a per-shard egress ring (see `threads::TM_CORES`). + #[inline] + fn route(&self, rid: &Rid, msg: DetokMsg) { + if self.senders.detok_for(rid).send(msg).is_err() { + tracing::error!("egress: detok shard closed"); + } + } +} + +/// Control result: `[rid, payload]` → single non-streamed delivery to the sink. +fn decode_result(body: &[u8]) -> Option<(Rid, DetokMsg)> { + let val = rmpv::decode::read_value(&mut &body[..]).ok()?; + let rmpv::Value::Array(arr) = val else { + return None; + }; + let mut items = arr.into_iter(); + let rid = Rid::from(items.next()?.as_str()?); + // The decode already owns the payload buffer — move it out. + let payload = match items.next()? { + rmpv::Value::Binary(b) => Bytes::from(b), + rmpv::Value::String(s) => Bytes::from(s.into_bytes()), + _ => return None, + }; + Some((rid.clone(), DetokMsg::Result { rid, payload })) +} + +/// Per-request failure: `[rid, message]` → terminal `Error` to the sink (→ 400). +fn decode_error(body: &[u8]) -> Option<(Rid, DetokMsg)> { + let val = rmpv::decode::read_value(&mut &body[..]).ok()?; + let rmpv::Value::Array(arr) = val else { + return None; + }; + let mut items = arr.into_iter(); + let rid = Rid::from(items.next()?.as_str()?); + let message = match items.next()? { + rmpv::Value::String(s) => s.into_str()?, + _ => return None, + }; + Some((rid.clone(), DetokMsg::Fail { rid, message })) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::message::DetokMsg; + use crate::message::frame_egress_error; + + /// A framed error round-trips: `frame_egress_error` → tag stripped → + /// `decode_error` yields the rid + a `Fail` carrying the message. + #[test] + fn error_frame_roundtrips_to_fail() { + let framed = frame_egress_error("42", "invalid request: bad field"); + assert_eq!(framed[0], EGRESS_TAG_ERROR); + let (rid, msg) = decode_error(&framed[1..]).expect("decodes"); + let want = Rid::from("42"); + assert_eq!(rid, want); + match msg { + DetokMsg::Fail { rid, message } => { + assert_eq!(rid.clone(), want); + assert_eq!(message, "invalid request: bad field"); + } + _ => panic!("expected Fail"), + } + } +}