Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 50 additions & 4 deletions python/src/edge0/server/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand All @@ -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,
Expand All @@ -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)
Expand All @@ -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,
Expand Down
124 changes: 117 additions & 7 deletions python/tests/test_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -119,13 +119,20 @@ def reset(self):

class ToolCallTok(FakeTok):
"""Decodes generated tokens to a Ling-style ``<tool_call>`` 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: "<tool_call>calculator\n",
12: "<arg_key>expression</arg_key>\n<arg_value>2 + 2</arg_value>\n",
13: "</tool_call>\n",
}

def decode(self, tokens):
return ("<tool_call>calculator\n"
"<arg_key>expression</arg_key>\n"
"<arg_value>2 + 2</arg_value>\n"
"</tool_call>")
return "".join(self._PIECES.get(t, "") for t in tokens)


class ToolCallEngine(FakeEngine):
Expand Down Expand Up @@ -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"

Expand Down Expand Up @@ -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
Loading