diff --git a/python/src/edge0/server/app.py b/python/src/edge0/server/app.py index fcbda96..a0f836c 100644 --- a/python/src/edge0/server/app.py +++ b/python/src/edge0/server/app.py @@ -89,6 +89,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() @@ -106,10 +130,12 @@ def _chat_stream(server: QueueServer, payload: dict): buffering = bool(req.tools) and req.tool_choice != "none" and callable( getattr(server.engine, "parse_tool_calls", None)) - def on_token(tid: int): - if buffering: - return - 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, @@ -118,6 +144,17 @@ def on_token(tid: int): "finish_reason": None}], }).encode("utf-8")) + def on_token(tid: int): + if buffering: + return + 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: tokens, meta = server.chat(req, on_token=on_token) @@ -138,6 +175,15 @@ def produce(): finish_reason = "tool_calls" else: delta = {"content": content} + else: + # 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 c153da2..3eba4d2 100644 --- a/python/tests/test_server.py +++ b/python/tests/test_server.py @@ -119,13 +119,20 @@ def reset(self): class ToolCallTok(FakeTok): """Decodes generated tokens to a Ling-style ```` block, the - same shape the real edge0-8b template asks the model to emit.""" + same shape the real edge0-8b template asks the model to emit. + Like a real tokenizer, decode is prefix-stable: the cumulative + decode of the id list equals the whole block.""" + + # FakeEngine.generate() emits ids 11, 12, 13; each id maps to + # one fixed piece so decode([11, 12, 13]) is the full block. + _PIECES = { + 11: "calculator\n", + 12: "expression\n2 + 2\n", + 13: "\n", + } def decode(self, tokens): - return ("calculator\n" - "expression\n" - "2 + 2\n" - "") + return "".join(self._PIECES.get(t, "") for t in tokens) class ToolCallEngine(FakeEngine): @@ -587,8 +594,8 @@ def test_chat_stream_without_tools_keeps_immediate_per_token_deltas(): chunks = [json.loads(e[6:]) for e in frames[:-1]] deltas = [c["choices"][0]["delta"].get("content") for c in chunks if c["choices"][0]["delta"]] - # unchanged from the no-tools path: whatever decode_tokens([tid]) - # returns per call, not the whole buffered response + # unchanged from the no-tools path: one delta per token, each + # emitted as the token arrives -- not the whole buffered response assert len(deltas) == 3 assert chunks[-1]["choices"][0]["finish_reason"] == "stop" @@ -794,3 +801,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