mlxrunner: record committed MTP drafts before streaming them

The batched MTP accept paths advance the cache by the whole accepted run
before streaming it to the client. If the stream was cancelled partway
(e.g. the caller disconnects), the loop returned before recording the
remaining accepted tokens, leaving the cache offset ahead of
session.outputs. close() then indexed the token log past its end and
panicked with a slice-bounds error.

Record the whole run to session.outputs before streaming any of it, so a
cancelled stream can no longer desync the cache from the token log.

The same bug is present on main, with identical mechanics: the accept
paths there commit the cache to before+accepted and then stream in a loop
that returns on cancellation before recording the rest.
This commit is contained in:
Jesse Gross
2026-06-01 11:00:46 -07:00
parent ded2db7d86
commit 1abd56b6e6

View File

@@ -206,14 +206,11 @@ func (r *Runner) runGreedyMTPDecode(ctx context.Context, request Request, sessio
now = time.Now() now = time.Now()
} }
done, err := r.emitMTPToken(ctx, request, session, &dec, current, &final) done, err := r.emitTokens(ctx, request, session, &dec, []sampler.Result{current}, &final, &generated)
if err != nil { if err != nil {
return err return err
} }
if !done { if done {
generated++
}
if done || generated >= request.Options.NumPredict {
break break
} }
@@ -349,14 +346,11 @@ func (r *Runner) runSampleMTPDecode(ctx context.Context, request Request, sessio
now = time.Now() now = time.Now()
} }
done, err := r.emitMTPToken(ctx, request, session, &dec, current, &final) done, err := r.emitTokens(ctx, request, session, &dec, []sampler.Result{current}, &final, &generated)
if err != nil { if err != nil {
return err return err
} }
if !done { if done {
generated++
}
if done || generated >= request.Options.NumPredict {
break break
} }
@@ -634,26 +628,11 @@ func (r *Runner) acceptMTPDraftsBatched(ctx context.Context, request Request, se
commitSpeculation(caches, accepted, draftCount, before) commitSpeculation(caches, accepted, draftCount, before)
*position = before + accepted *position = before + accepted
for _, id := range draftIDs[:accepted] { emitted, err := r.emitTokens(ctx, request, session, dec, draftResults(draftIDs[:accepted]), final, generated)
if *generated >= request.Options.NumPredict { if err != nil {
done = true return sampler.Result{}, accepted, emitted || done, err
break
}
res := sampler.Result{Token: mlx.FromValues([]int32{int32(id)}, 1)}
var err error
done, err = r.emitMTPToken(ctx, request, session, dec, res, final)
if err != nil {
return sampler.Result{}, accepted, done, err
}
if !done {
(*generated)++
}
if done {
break
}
} }
if emitted || done {
if done || *generated >= request.Options.NumPredict {
return sampler.Result{}, accepted, true, nil return sampler.Result{}, accepted, true, nil
} }
if next.Token == nil { if next.Token == nil {
@@ -711,26 +690,11 @@ func (r *Runner) acceptSampleMTPDrafts(ctx context.Context, request Request, ses
commitSpeculation(caches, accepted, draftCount, before) commitSpeculation(caches, accepted, draftCount, before)
*position = before + accepted *position = before + accepted
for _, id := range draftIDs[:accepted] { emitted, err := r.emitTokens(ctx, request, session, dec, draftResults(draftIDs[:accepted]), final, generated)
if *generated >= request.Options.NumPredict { if err != nil {
done = true return sampler.Result{}, accepted, emitted || done, err
break
}
res := sampler.Result{Token: mlx.FromValues([]int32{int32(id)}, 1)}
var err error
done, err = r.emitMTPToken(ctx, request, session, dec, res, final)
if err != nil {
return sampler.Result{}, accepted, done, err
}
if !done {
(*generated)++
}
if done {
break
}
} }
if emitted || done {
if done || *generated >= request.Options.NumPredict {
r.Sampler.Commit(pipelineSlot, commitIDs) r.Sampler.Commit(pipelineSlot, commitIDs)
return sampler.Result{}, accepted, true, nil return sampler.Result{}, accepted, true, nil
} }
@@ -811,14 +775,11 @@ func (r *Runner) acceptMTPDraftsSerial(ctx context.Context, request Request, ses
accepted++ accepted++
res := sampler.Result{Token: mlx.FromValues([]int32{int32(id)}, 1)} res := sampler.Result{Token: mlx.FromValues([]int32{int32(id)}, 1)}
done, err := r.emitMTPToken(ctx, request, session, dec, res, final) done, err := r.emitTokens(ctx, request, session, dec, []sampler.Result{res}, final, generated)
if err != nil { if err != nil {
return sampler.Result{}, accepted, done, err return sampler.Result{}, accepted, done, err
} }
if !done { if done {
(*generated)++
}
if done || *generated >= request.Options.NumPredict {
return sampler.Result{}, accepted, true, nil return sampler.Result{}, accepted, true, nil
} }
@@ -828,23 +789,51 @@ func (r *Runner) acceptMTPDraftsSerial(ctx context.Context, request Request, ses
return sampler.Result{Token: greedyTokenFromLogits(logits)}, accepted, false, nil return sampler.Result{Token: greedyTokenFromLogits(logits)}, accepted, false, nil
} }
func (r *Runner) emitMTPToken(ctx context.Context, request Request, session *cacheSession, dec *decoder, res sampler.Result, final *CompletionResponse) (bool, error) { // emitTokens records a run of generated tokens to session.outputs, then streams
output := int32(tokenID(res.Token)) // them. A trailing EOS stops generation and is recorded but not streamed.
session.outputs = append(session.outputs, output) // Returns whether to stop and any cancellation error.
func (r *Runner) emitTokens(ctx context.Context, request Request, session *cacheSession, dec *decoder, results []sampler.Result, final *CompletionResponse, generated *int) (done bool, err error) {
if r.Tokenizer.IsEOS(output) { stream := len(results)
final.DoneReason = 0 for i, res := range results {
return true, nil id := int32(tokenID(res.Token))
session.outputs = append(session.outputs, id)
if r.Tokenizer.IsEOS(id) {
final.DoneReason = 0
done = true
stream = i
break
}
(*generated)++
}
if *generated >= request.Options.NumPredict {
done = true
} }
if resp, ok := dec.decode(res); ok { // Record the whole run before streaming any of it: streaming returns early on
// a cancelled context, and a partial stream must not leave the cache ahead of
// session.outputs.
for _, res := range results[:stream] {
resp, ok := dec.decode(res)
if !ok {
continue
}
select { select {
case <-ctx.Done(): case <-ctx.Done():
return false, ctx.Err() return done, ctx.Err()
case request.Responses <- resp: case request.Responses <- resp:
} }
} }
return false, nil return done, nil
}
// draftResults wraps accepted draft token ids as sampler results for emitTokens.
// Accepted drafts carry no logprobs, so only the token id is set.
func draftResults(ids []int) []sampler.Result {
results := make([]sampler.Result, len(ids))
for i, id := range ids {
results[i] = sampler.Result{Token: mlx.FromValues([]int32{int32(id)}, 1)}
}
return results
} }
func (r *Runner) lastLogits(hidden *mlx.Array) *mlx.Array { func (r *Runner) lastLogits(hidden *mlx.Array) *mlx.Array {