Skip to content
Merged
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
41 changes: 41 additions & 0 deletions internal/inference/llama.go
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,8 @@ type llamaEngine struct {
logFile *os.File
hasMultimodal bool
nativeToolStreaming bool
// slots pins conversations to llama-server slots; nil with one slot.
slots *llamaSlotScheduler
}

type inferenceHTTPError struct {
Expand Down Expand Up @@ -533,6 +535,9 @@ func newLlamaEngineWithMode(modelPath, modelName string, verbose bool, progress
effectiveNumCtx = CapNumCtxToEmbeddingModelMax(filepath.Dir(modelPath), effectiveNumCtx)
}
effectiveNumParallel := ResolveNumParallel(numParallel)
if !embedding {
engine.slots = newLlamaSlotScheduler(effectiveNumParallel)
}
effectiveNGPULayers := ResolveNGPULayers(nGPULayers)
normalizedCacheTypeK, err := NormalizeCacheType(cacheTypeK)
if err != nil {
Expand Down Expand Up @@ -723,6 +728,42 @@ func (e *llamaEngine) baseURL() string {
}

func (e *llamaEngine) ChatCompletion(ctx context.Context, reqBody map[string]interface{}) (*http.Response, error) {
release := func() {}
if _, explicit := reqBody["id_slot"]; !explicit {
var slot int
slot, release = e.slots.acquire(SlotAffinity(ctx))
if slot >= 0 {
pinned := make(map[string]interface{}, len(reqBody)+1)
for k, v := range reqBody {
pinned[k] = v
}
pinned["id_slot"] = slot
reqBody = pinned
}
}
resp, err := e.postChatCompletion(ctx, reqBody)
if err != nil {
release()
return nil, err
}
resp.Body = &releasingBody{ReadCloser: resp.Body, release: release}
return resp, nil
}

// releasingBody frees the request's slot once the caller has finished with
// the response, streamed or not.
type releasingBody struct {
io.ReadCloser
release func()
}

func (b *releasingBody) Close() error {
err := b.ReadCloser.Close()
b.release()
return err
}

func (e *llamaEngine) postChatCompletion(ctx context.Context, reqBody map[string]interface{}) (*http.Response, error) {
body, err := json.Marshal(reqBody)
if err != nil {
return nil, fmt.Errorf("marshaling request: %w", err)
Expand Down
135 changes: 135 additions & 0 deletions internal/inference/llama_slots.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,135 @@
package inference

import (
"context"
"sync"
)

// llama-server keeps one KV cache per slot, and picks a slot for each request
// by the longest prefix it shares with what a slot already holds. Two
// conversations from the same agent share their system prompt and tool
// definitions, so the slot holding one of them always outscores an empty
// slot: every request lands in that one slot, the others stay unused, and
// each switch between the conversations throws away everything after the
// shared prefix. Measured on Qwen3.5-2B with two slots, two interleaved agent
// conversations kept 41% of their prompt tokens cached; pinning each to its
// own slot with id_slot kept 81% and cut the wall time from 92 s to 30 s.
//
// llamaSlotScheduler does that pinning. It only acts when the server has two
// or more slots, so a default single-slot load behaves exactly as before.

type slotAffinityContextKey struct{}

// WithSlotAffinity tags a request with the conversation it belongs to, so a
// local llama-server keeps that conversation in the same slot.
func WithSlotAffinity(ctx context.Context, key string) context.Context {
if key == "" {
return ctx
}
return context.WithValue(ctx, slotAffinityContextKey{}, key)
}

// SlotAffinity returns the conversation key WithSlotAffinity stored.
func SlotAffinity(ctx context.Context) string {
key, _ := ctx.Value(slotAffinityContextKey{}).(string)
return key
}

type llamaSlotScheduler struct {
mu sync.Mutex
slots []llamaSlot
clock uint64
}

type llamaSlot struct {
owner string // conversation whose context the slot holds
busy int // requests currently running on the slot
lastUsed uint64
}

func newLlamaSlotScheduler(numParallel int) *llamaSlotScheduler {
if numParallel < 2 {
return nil
}
return &llamaSlotScheduler{slots: make([]llamaSlot, numParallel)}
}

// acquire picks the slot for a request of conversation key, marks it busy and
// returns it. It returns -1 only when scheduling is off. release must be
// called once the response is done.
//
// Every request is pinned, and whichever request runs on a slot becomes its
// owner: its prompt replaces the KV cache there, so the owner is always the
// conversation whose context the slot really holds. A request without a key
// leaves the slot unowned.
//
// A conversation always goes back to its own slot, even when that slot is
// busy: llama-server then queues the request for it, which keeps the cache,
// where running on another slot would both miss the cache and overwrite
// someone else's. Any other request takes an idle slot nobody owns, then the
// idle slot used longest ago, and only when every slot is busy queues on the
// busy slot used longest ago. Leaving the choice to llama-server there would
// bring back the longest-prefix collisions this scheduler exists to avoid.
func (s *llamaSlotScheduler) acquire(key string) (int, func()) {
if s == nil {
return -1, func() {}
}
s.mu.Lock()
defer s.mu.Unlock()
s.clock++

slot := -1
if key != "" {
slot = s.ownedSlot(key)
}
if slot < 0 {
slot = s.freeSlot()
}
s.slots[slot].owner = key
s.slots[slot].busy++
s.slots[slot].lastUsed = s.clock

var once sync.Once
return slot, func() {
once.Do(func() {
s.mu.Lock()
s.slots[slot].busy--
s.mu.Unlock()
})
}
}

func (s *llamaSlotScheduler) ownedSlot(key string) int {
for i := range s.slots {
if s.slots[i].owner == key {
return i
}
}
return -1
}

// freeSlot picks the slot for a request that has none of its own: an idle
// slot nobody owns, then the idle slot used longest ago, then the busy slot
// used longest ago.
func (s *llamaSlotScheduler) freeSlot() int {
idle, busy := -1, -1
for i := range s.slots {
slot := s.slots[i]
if slot.busy == 0 {
if slot.owner == "" {
return i
}
if idle < 0 || slot.lastUsed < s.slots[idle].lastUsed {
idle = i
}
continue
}
if busy < 0 || slot.lastUsed < s.slots[busy].lastUsed {
busy = i
}
}
if idle >= 0 {
return idle
}
return busy
}
199 changes: 199 additions & 0 deletions internal/inference/llama_slots_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,199 @@
package inference

import (
"context"
"encoding/json"
"io"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strconv"
"sync"
"testing"
)

func TestSlotSchedulerOffWithOneSlot(t *testing.T) {
if s := newLlamaSlotScheduler(1); s != nil {
t.Fatal("scheduler created for a single slot")
}
var s *llamaSlotScheduler
if slot, release := s.acquire("a"); slot != -1 {
t.Fatalf("nil scheduler picked slot %d", slot)
} else {
release()
}
}

func TestSlotSchedulerKeepsConversationsInTheirSlots(t *testing.T) {
s := newLlamaSlotScheduler(2)
use := func(key string) int {
slot, release := s.acquire(key)
release()
return slot
}
a, b := use("A"), use("B")
if a == b || a < 0 || b < 0 {
t.Fatalf("A=%d B=%d, want two different slots", a, b)
}
for i := 0; i < 3; i++ {
if got := use("A"); got != a {
t.Fatalf("A moved to slot %d", got)
}
if got := use("B"); got != b {
t.Fatalf("B moved to slot %d", got)
}
}
}

func TestSlotSchedulerQueuesOnOwnBusySlot(t *testing.T) {
s := newLlamaSlotScheduler(2)
own, releaseFirst := s.acquire("A")
// A second A request waits for A's slot rather than overwriting another.
if again, release := s.acquire("A"); again != own {
t.Fatalf("second A request got slot %d, want its own busy slot %d", again, own)
} else {
defer release()
}
b, releaseB := s.acquire("B")
if b == own || b < 0 {
t.Fatalf("B got slot %d while slot %d is A's", b, own)
}
releaseFirst()
releaseFirst() // releasing twice must not free a slot twice
if s.slots[own].busy != 1 {
t.Fatalf("busy count = %d after double release, want 1", s.slots[own].busy)
}
releaseB()
}

func TestSlotSchedulerPinsEvenWhenEverySlotIsBusy(t *testing.T) {
s := newLlamaSlotScheduler(2)
a, _ := s.acquire("A")
b, _ := s.acquire("B")
c, releaseC := s.acquire("C")
defer releaseC()
if c != a {
t.Fatalf("C got slot %d, want the least recently used busy slot %d", c, a)
}
if s.slots[c].owner != "C" {
t.Fatalf("slot %d owner = %q, want C, whose prompt will replace A's", c, s.slots[c].owner)
}
// A lost its slot, so it no longer finds one of its own there.
if s.ownedSlot("A") != -1 || s.ownedSlot("B") != b {
t.Fatalf("owners after C: A=%d B=%d", s.ownedSlot("A"), s.ownedSlot("B"))
}
}

func TestSlotSchedulerEvictsLeastRecentlyUsed(t *testing.T) {
s := newLlamaSlotScheduler(2)
use := func(key string) int {
slot, release := s.acquire(key)
release()
return slot
}
a, b := use("A"), use("B")
use("B")
if c := use("C"); c != a {
t.Fatalf("C took slot %d, want A's slot %d (least recently used)", c, a)
}
if got := use("B"); got != b {
t.Fatalf("B lost its slot to C: got %d", got)
}
// A was evicted, so it now takes the least recently used slot, C's.
if got := use("A"); got != a {
t.Fatalf("A got %d, want %d", got, a)
}
}

func TestSlotSchedulerKeylessRequestsTakeOverWhatTheyOverwrite(t *testing.T) {
s := newLlamaSlotScheduler(2)
use := func(key string) int {
slot, release := s.acquire(key)
release()
return slot
}
a := use("A")
if anon := use(""); anon == a {
t.Fatalf("keyless request used A's slot %d while unowned slots were idle", anon)
}
b := use("B") // takes the slot the keyless request left unowned
use("B")
// Both slots are owned now, and A's is the least recently used one.
if anon := use(""); anon != a {
t.Fatalf("keyless request took slot %d, want A's least recently used slot %d", anon, a)
}
if s.ownedSlot("A") != -1 {
t.Fatal("A still owns the slot a keyless request overwrote")
}
if got := use("B"); got != b {
t.Fatalf("B moved to slot %d", got)
}
}

// fakeLlamaServer records the id_slot of every chat completion it receives.
func fakeLlamaServer(t *testing.T) (*llamaEngine, func() []any) {
t.Helper()
var mu sync.Mutex
var slots []any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var body map[string]any
_ = json.NewDecoder(r.Body).Decode(&body)
mu.Lock()
slots = append(slots, body["id_slot"])
mu.Unlock()
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"choices":[]}`))
}))
t.Cleanup(server.Close)
parsed, _ := url.Parse(server.URL)
_, portText, _ := net.SplitHostPort(parsed.Host)
port, _ := strconv.Atoi(portText)
return &llamaEngine{port: port, client: server.Client()}, func() []any {
mu.Lock()
defer mu.Unlock()
return append([]any(nil), slots...)
}
}

func TestLlamaChatCompletionPinsConversationSlot(t *testing.T) {
engine, sent := fakeLlamaServer(t)
engine.slots = newLlamaSlotScheduler(2)
call := func(key string, body map[string]interface{}) {
resp, err := engine.ChatCompletion(WithSlotAffinity(context.Background(), key), body)
if err != nil {
t.Fatal(err)
}
_, _ = io.Copy(io.Discard, resp.Body)
_ = resp.Body.Close()
}
shared := map[string]interface{}{"messages": []any{}}
for i := 0; i < 2; i++ {
call("A", shared)
call("B", shared)
}
call("A", map[string]interface{}{"id_slot": 1})

got := sent()
if got[0] != got[2] || got[1] != got[3] || got[0] == got[1] || got[0] == nil {
t.Fatalf("id_slot sent = %v, want A and B each in their own slot", got)
}
if got[4] != float64(1) {
t.Fatalf("explicit id_slot overridden: %v", got[4])
}
if _, ok := shared["id_slot"]; ok {
t.Fatal("caller's request body was modified")
}
}

func TestLlamaChatCompletionSingleSlotSendsNoSlot(t *testing.T) {
engine, sent := fakeLlamaServer(t)
resp, err := engine.ChatCompletion(WithSlotAffinity(context.Background(), "A"), map[string]interface{}{"messages": []any{}})
if err != nil {
t.Fatal(err)
}
_ = resp.Body.Close()
if got := sent(); got[0] != nil {
t.Fatalf("single-slot engine sent id_slot %v", got[0])
}
}
Loading
Loading