From 690724a064358e6381c9f500d037a5c9b218a456 Mon Sep 17 00:00:00 2001 From: ahrazzle Date: Tue, 15 Sep 2026 08:25:42 -0400 Subject: [PATCH] fix(server): decode streamed SSE text from the cumulative id list Rebased onto main: upstream moved the tree under python/; the fix applies unchanged to python/src/edge0/server/app.py. --- python/src/edge0/server/app.py | 49 +++++++++++++++- python/tests/test_server.py | 103 +++++++++++++++++++++++++++++++++ 2 files changed, 150 insertions(+), 2 deletions(-) diff --git a/python/src/edge0/server/app.py b/python/src/edge0/server/app.py index 9d94b1b..35b2718 100644 --- a/python/src/edge0/server/app.py +++ b/python/src/edge0/server/app.py @@ -72,6 +72,30 @@ def _chat_once(server: QueueServer, payload: dict): } +def _incremental_suffix(prev: str, text: str) -> str: + """Suffix of a cumulative decode not yet emitted, holding back a + trailing incomplete character. + + A byte-level BPE token is a *byte fragment*: decoding one id in + isolation turns any multi-byte UTF-8 character it belongs to into + U+FFFD. Decoding the cumulative id list instead is correct, and the + next id completes a character that was still incomplete, so a + trailing replacement char is held back until it resolves (never emit + a partial multi-byte character).""" + safe = text.rstrip("\ufffd") + if safe.startswith(prev): + return safe[len(prev):] + # Defensive resync: a prefix-stable byte-level decoder never lands + # here, but if a decoder rewrote text below the emitted prefix, emit + # only past the longest shared prefix instead of re-emitting it. + i = 0 + for a, b in zip(safe, prev): + if a != b: + break + i += 1 + return safe[i:] + + def _chat_stream(server: QueueServer, payload: dict): req = parse_chat_request(payload) events = queue.Queue() @@ -79,8 +103,12 @@ def _chat_stream(server: QueueServer, payload: dict): request_id = f"chatcmpl-{int(time.time() * 1000)}" created = int(time.time()) - def on_token(tid: int): - text = decode_tokens(server.engine, [tid]) + # Keep every id seen so far, decode the cumulative list, emit only the + # newly completed suffix (see _incremental_suffix). + seen_ids: list[int] = [] + emitted = "" + + def _emit(text: str): events.put(sse_format({ "id": request_id, "object": "chat.completion.chunk", "created": created, "model": server.model_name, @@ -89,9 +117,26 @@ def on_token(tid: int): "finish_reason": None}], }).encode("utf-8")) + def on_token(tid: int): + nonlocal emitted + seen_ids.append(tid) + text = decode_tokens(server.engine, seen_ids) + delta = _incremental_suffix(emitted, text) + emitted = text.rstrip("\ufffd") + if delta: + _emit(delta) + def produce(): try: _, meta = server.chat(req, on_token=on_token) + # Flush whatever the incremental decode held back (an + # incomplete multi-byte sequence at the very end of the + # generation) so the concatenated deltas equal the + # non-streaming text exactly. + tail = _incremental_suffix( + emitted, decode_tokens(server.engine, seen_ids)) + if tail: + _emit(tail) events.put(sse_format({ "id": request_id, "object": "chat.completion.chunk", "created": created, "model": server.model_name, diff --git a/python/tests/test_server.py b/python/tests/test_server.py index bad235b..5e9351a 100644 --- a/python/tests/test_server.py +++ b/python/tests/test_server.py @@ -453,3 +453,106 @@ def test_flask_http_streams_before_generation_finishes(): eng.resume.set() response.close() assert b"data: [DONE]\n\n" in body + + +# ---- SSE incremental decode (multi-byte characters split across ids) ------ + + +def _byte_level_bpe(seed_texts): + """A real byte-level BPE (the family the shipped Qwen/Ling tokenizers + belong to): out-of-vocabulary characters fall back to per-byte tokens, + so a multi-byte character ends up split across several ids.""" + from tokenizers import (Tokenizer, decoders, models, pre_tokenizers, + trainers) + + tok = Tokenizer(models.BPE()) + tok.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=False) + tok.decoder = decoders.ByteLevel() + trainer = trainers.BpeTrainer( + vocab_size=300, special_tokens=["<|endoftext|>"], + initial_alphabet=pre_tokenizers.ByteLevel.alphabet()) + tok.train_from_iterator(seed_texts, trainer=trainer) + return tok + + +class _BpeTok: + """Engine-tokenizer facade over a byte-level BPE (what the server sees).""" + + def __init__(self, tok): + self._tok = tok + self.bos_token_id = 0 + + def apply_chat_template(self, messages, tokenize=False, + add_generation_prompt=True, + enable_thinking=False): + return "<|im_start|>user\nx<|im_end|>\n<|im_start|>assistant\n" + + def encode(self, text): + return list(self._tok.encode(text).ids) + + def decode(self, tokens): + return self._tok.decode(list(tokens)) + + +class SplitEngine: + """Engine that emits a fixed id sequence whose bytes split a character.""" + + name = "split" + + def __init__(self, tok, ids): + self._tok = tok + self.ids = list(ids) + self.pos = 0 + self.cfg = SimpleNamespace(gen=GenerationConfig()) + + def generate(self, ids, gen_config=None, on_token=None, **kw): + if on_token is not None: + for t in self.ids: + on_token(t) + return list(self.ids) + + def stats(self): + return {} + + def reset(self): + self.pos = 0 + + +def test_chat_stream_incremental_decode_of_split_characters(): + """The streamed deltas must equal the non-streaming text. + + A byte-level BPE token carries a byte fragment, not a character: + decoding each id in isolation turns any multi-byte UTF-8 character it + belongs to into U+FFFD, so the old per-id path corrupted every + non-ASCII reply while the non-streaming path (whole id list) stayed + correct. Assert on that exact trigger: an id sequence whose bytes + split a character, streamed through the handler. + """ + tok = _byte_level_bpe(["你好世界 hello world", "中文测试 沿着海滨", + "日本語テスト"]) + trigger = "你好,用一句话介绍海滨城市。" + ids = list(tok.encode(trigger).ids) + + # Precondition: the trigger really splits a character across ids, so + # per-id decode is corrupt by construction (this is the bug). + isolated = [tok.decode([i]) for i in ids] + assert any("\ufffd" in p for p in isolated), \ + "trigger does not split a multi-byte character across ids" + assert "\ufffd" not in tok.decode(ids) + + eng = SplitEngine(_BpeTok(tok), ids) + srv = QueueServer(eng, model_name="split") + events = _chat_stream(srv, { + "messages": [{"role": "user", "content": "hi"}], "stream": True, + }) + body = b"".join(events) + chunks = [json.loads(e[6:]) for e in + body.decode("utf-8").split("\n\n") + if e and e != "data: [DONE]"] + deltas = [c["choices"][0]["delta"]["content"] + for c in chunks if c["choices"][0]["delta"].get("content")] + + streamed = "".join(deltas) + non_streamed = decode_tokens(eng, ids) + assert streamed == non_streamed == trigger + assert "\ufffd" not in streamed