[Integration] ollama: three-stage reuse orchestrator

Ties the pieces together into Lookup -> Prepare -> Commit:

  Lookup   build block-prefix keys, find the longest prefix present in the store
           (one batched, multi-node-correct existence call) and consult the
           index + arbiter.
  Prepare  if the arbiter approves, GetFile the matched KV and restore it into a
           llama.cpp slot so only the tail is prefilled; feed the full restore
           wall back to the arbiter so its bandwidth estimate self-calibrates.
  Commit   save the slot KV and PutFile it under the prefix's block key, with
           skip-if-exists dedup so concurrent agents store a shared prefix once.

Also owns per-model KV bytes/token learning and the Prometheus accounting.

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
zbtrs2 2026-06-23 19:10:16 +08:00
parent 113076bc96
commit cef91c4dcd
1 changed files with 494 additions and 0 deletions

View File

@ -0,0 +1,494 @@
// Package orchestrator implements the three-stage KV reuse flow:
//
// Lookup compute block-prefix keys, find the longest prefix present in the
// store (authoritative, batched), and consult the radix index.
// Prepare Lookup + (if the cost arbiter approves) GetFile the matched KV and
// restore it into a llama.cpp slot, so only the tail is prefilled.
// Commit Save the slot KV and PutFile it to the store under its block key,
// with single-writer dedup so concurrent agents store a shared prefix
// exactly once.
//
// It owns the cross-cutting policy: cache-key construction, longest-prefix
// matching, the restore-vs-recompute arbiter, the radix index, per-model
// bytes/token learning, and Prometheus accounting.
package orchestrator
import (
"context"
"fmt"
"os"
"path/filepath"
"sync"
"sync/atomic"
"time"
"github.com/mooncake-ai/ollama-mooncake-bridge/internal/arbiter"
"github.com/mooncake-ai/ollama-mooncake-bridge/internal/cachekey"
"github.com/mooncake-ai/ollama-mooncake-bridge/internal/llamabridge"
"github.com/mooncake-ai/ollama-mooncake-bridge/internal/metrics"
"github.com/mooncake-ai/ollama-mooncake-bridge/internal/prefixindex"
"github.com/mooncake-ai/ollama-mooncake-bridge/internal/store"
)
type Config struct {
SlotSavePath string // == llama.cpp --slot-save-path (shared filesystem)
DefaultBlockSize int
DefaultReplicaNum uint32
MinPrefixBlocks int
CleanupFiles bool // delete the local save file after store put / after restore
}
func (c *Config) defaults() {
if c.DefaultBlockSize <= 0 {
c.DefaultBlockSize = 256
}
if c.DefaultReplicaNum == 0 {
c.DefaultReplicaNum = 1
}
if c.MinPrefixBlocks <= 0 {
c.MinPrefixBlocks = 1
}
}
type Orchestrator struct {
store store.Backend
llama *llamabridge.Client
index *prefixindex.Index
arb *arbiter.Arbiter
mx *metrics.Metrics
cfg Config
fpMu sync.Mutex
fpCache map[string]cachekey.ModelFingerprint // gguf path -> completed fingerprint
bptMu sync.Mutex
bytesPerTok map[string]float64 // fingerprint prefix -> learned KV bytes/token
nonce atomic.Uint64 // makes per-request slot filenames unique
}
func New(b store.Backend, arb *arbiter.Arbiter, mx *metrics.Metrics, cfg Config) *Orchestrator {
cfg.defaults()
return &Orchestrator{
store: b, llama: llamabridge.New(), index: prefixindex.New(),
arb: arb, mx: mx, cfg: cfg,
fpCache: map[string]cachekey.ModelFingerprint{},
bytesPerTok: map[string]float64{},
}
}
func (o *Orchestrator) Index() *prefixindex.Index { return o.index }
func (o *Orchestrator) Arbiter() *arbiter.Arbiter { return o.arb }
// ---- request/response types (decoupled from gRPC wire types) ----
type Policy struct {
Enable bool
Namespace string
Read bool
Write bool
BlockSize int
ReplicaNum uint32
SoftPin bool
MinPrefixBlocks int
}
type Target struct {
BaseURL string
Slot int
}
type LookupResult struct {
Hit bool
MatchedBlocks int
MatchedTokens int
TotalBlocks int
TotalTokens int
Decision string // restore | recompute | miss
MatchedKey string
Reason string
}
type PrepareResult struct {
LookupResult
Restored bool
RestoredTokens int
RestoreMs float64
StoreGetMs float64
Bytes uint64
}
type CommitResult struct {
OK bool
Stored bool
StoredBlocks int
StoredTokens int
Bytes uint64
SaveMs float64
StorePutMs float64
Key string
Reason string
}
// completeFingerprint fills empty fields by parsing the GGUF at ModelPath
// (cached). The completed fingerprint is what every key derives from, enforcing
// the "different model/tokenizer/rope => different key space" safety rule.
func (o *Orchestrator) completeFingerprint(fp cachekey.ModelFingerprint, modelPath string, swa bool) cachekey.ModelFingerprint {
if fp.BlockSize <= 0 {
fp.BlockSize = o.cfg.DefaultBlockSize
}
needEnrich := fp.Arch == "" || fp.TokenizerHash == "" || fp.RopeHash == ""
if modelPath == "" || !needEnrich {
if fp.KVType == "" {
fp.KVType = "f16"
}
if fp.ModelDigest == "" {
fp.ModelDigest = "unknown"
}
return fp
}
o.fpMu.Lock()
defer o.fpMu.Unlock()
if cached, ok := o.fpCache[modelPath]; ok {
// keep caller-provided block size / kv type
cached.BlockSize = fp.BlockSize
if fp.KVType != "" {
cached.KVType = fp.KVType
}
return cached
}
meta, err := cachekey.ReadGGUFMeta(modelPath)
if err != nil {
if fp.KVType == "" {
fp.KVType = "f16"
}
if fp.ModelDigest == "" {
fp.ModelDigest = "ggufpath-" + filepath.Base(modelPath)
}
return fp
}
kv := fp.KVType
if kv == "" {
kv = "f16"
}
full := meta.Fingerprint(fp.ModelDigest, kv, swa, fp.BlockSize)
o.fpCache[modelPath] = full
return full
}
type planned struct {
fp cachekey.ModelFingerprint
rootKey string
modelKey string // per-model arbiter key (model digest + kv type), namespace-independent
chain cachekey.Chain
keys []string // keys[i] = boundary i+1 (i.e. prefix of i+1 blocks)
}
func (o *Orchestrator) plan(fp cachekey.ModelFingerprint, pol Policy, tokens []int32, modelPath string) planned {
full := o.completeFingerprint(fp, modelPath, pol.BlockSize < 0)
bs := full.BlockSize
if pol.BlockSize > 0 {
bs = pol.BlockSize
full.BlockSize = bs
}
rootKey := full.Prefix(pol.Namespace)
modelKey := full.ModelDigest + "|" + full.KVType + "|" + full.Arch
chain := cachekey.ChainBlockHashes(rootKey, tokens, bs)
keys := make([]string, chain.FullBlocks)
for i := 0; i < chain.FullBlocks; i++ {
keys[i] = full.BlockKey(pol.Namespace, i+1, chain.Hex[i])
}
return planned{fp: full, rootKey: rootKey, modelKey: modelKey, chain: chain, keys: keys}
}
// longestPresent returns the largest boundary M (in blocks) whose key exists in
// the store, querying all boundaries in one batched call (multi-node correct).
func (o *Orchestrator) longestPresent(ctx context.Context, p planned) (int, string, error) {
if len(p.keys) == 0 {
return 0, "", nil
}
present, err := o.store.Exists(ctx, p.keys)
if err != nil {
return 0, "", err
}
for i := len(present) - 1; i >= 0; i-- {
if i < len(present) && present[i] == 1 {
return i + 1, p.keys[i], nil
}
}
return 0, "", nil
}
func (o *Orchestrator) learnedBytesPerTok(rootKey string, fallbackTokens int, fallbackBytes uint64) float64 {
o.bptMu.Lock()
defer o.bptMu.Unlock()
if v, ok := o.bytesPerTok[rootKey]; ok && v > 0 {
return v
}
if fallbackTokens > 0 && fallbackBytes > 0 {
return float64(fallbackBytes) / float64(fallbackTokens)
}
return 0
}
func (o *Orchestrator) updateBytesPerTok(rootKey string, tokens int, bytes uint64) {
if tokens <= 0 || bytes == 0 {
return
}
o.bptMu.Lock()
defer o.bptMu.Unlock()
bpt := float64(bytes) / float64(tokens)
if old, ok := o.bytesPerTok[rootKey]; ok {
o.bytesPerTok[rootKey] = 0.5*old + 0.5*bpt
} else {
o.bytesPerTok[rootKey] = bpt
}
}
// Lookup is read-only: longest-prefix match + arbiter decision, no llama I/O.
func (o *Orchestrator) Lookup(ctx context.Context, fp cachekey.ModelFingerprint, pol Policy, tokens []int32, modelPath string) (LookupResult, error) {
o.mx.LookupTotal.Inc()
p := o.plan(fp, pol, tokens, modelPath)
res := LookupResult{TotalBlocks: p.chain.FullBlocks, TotalTokens: p.chain.FullBlocks * p.fp.BlockSize}
if p.chain.FullBlocks == 0 {
res.Decision = "miss"
res.Reason = "prompt shorter than one block"
return res, nil
}
m, key, err := o.longestPresent(ctx, p)
if err != nil {
return res, err
}
if m == 0 {
res.Decision = "miss"
res.Reason = "no cached prefix"
return res, nil
}
matchedTokens := m * p.fp.BlockSize
// size estimate for the arbiter
im := o.index.LongestMatch(p.rootKey, p.chain.Hex, false)
var estBytes uint64
if im.Found && im.Blocks == m {
estBytes = im.Bytes
} else {
bpt := o.learnedBytesPerTok(p.rootKey, 0, 0)
estBytes = uint64(float64(matchedTokens) * bpt)
}
dec := o.arb.Decide(p.modelKey, matchedTokens, estBytes)
res.Hit = true
res.MatchedBlocks = m
res.MatchedTokens = matchedTokens
res.MatchedKey = key
res.Reason = dec.Reason
if dec.Restore {
res.Decision = "restore"
} else {
res.Decision = "recompute"
}
return res, nil
}
// Prepare runs Lookup and, if the arbiter approves and a target is given,
// restores the matched KV into target.Slot.
func (o *Orchestrator) Prepare(ctx context.Context, fp cachekey.ModelFingerprint, pol Policy, tokens []int32, modelPath string, tgt *Target) (PrepareResult, error) {
o.mx.PrepareTotal.Inc()
prepStart := time.Now() // full restore cost incl lookup + transfer + GPU load + orchestration
p := o.plan(fp, pol, tokens, modelPath)
out := PrepareResult{}
out.TotalBlocks = p.chain.FullBlocks
out.TotalTokens = p.chain.FullBlocks * p.fp.BlockSize
if !pol.Read || p.chain.FullBlocks == 0 {
out.Decision = "miss"
out.Reason = "read disabled or prompt < 1 block"
o.mx.Misses.Inc()
o.mx.MissBlocks.Add(float64(p.chain.FullBlocks))
return out, nil
}
m, key, err := o.longestPresent(ctx, p)
if err != nil {
return out, err
}
if m == 0 {
out.Decision = "miss"
out.Reason = "no cached prefix"
o.mx.Misses.Inc()
o.mx.MissBlocks.Add(float64(p.chain.FullBlocks))
return out, nil
}
out.Hit = true
out.MatchedBlocks = m
out.MatchedTokens = m * p.fp.BlockSize
out.MatchedKey = key
im := o.index.LongestMatch(p.rootKey, p.chain.Hex, true)
var estBytes uint64
if im.Found && im.Blocks == m {
estBytes = im.Bytes
} else {
estBytes = uint64(float64(out.MatchedTokens) * o.learnedBytesPerTok(p.rootKey, 0, 0))
}
dec := o.arb.Decide(p.modelKey, out.MatchedTokens, estBytes)
out.Reason = dec.Reason
if !dec.Restore {
out.Decision = "recompute"
o.mx.RecomputeChosen.Inc()
o.mx.MissBlocks.Add(float64(p.chain.FullBlocks)) // will be recomputed
return out, nil
}
out.Decision = "restore"
if tgt == nil || tgt.BaseURL == "" {
// plan-only (no target): report the decision without doing I/O.
return out, nil
}
// ---- Load stage ----
fname := fmt.Sprintf("omb-r-%x-%d.bin", hashKey(key), o.nonce.Add(1))
abspath := filepath.Join(o.cfg.SlotSavePath, fname)
gr, err := o.store.GetFile(ctx, key, abspath)
if err != nil {
return out, fmt.Errorf("store get: %w", err)
}
if !gr.Found {
// raced with eviction; degrade to miss
out.Decision = "miss"
out.Hit = false
out.Reason = "matched key vanished (evicted); recompute"
o.mx.Misses.Inc()
o.mx.MissBlocks.Add(float64(p.chain.FullBlocks))
return out, nil
}
out.StoreGetMs = gr.ElapsedMs
out.Bytes = gr.Bytes
o.mx.StoreGetLatency.Observe(gr.ElapsedMs)
o.mx.AddBytesGet(gr.Bytes)
tRestore := time.Now()
rr, err := o.llama.RestoreSlot(ctx, tgt.BaseURL, tgt.Slot, fname)
restoreMs := float64(time.Since(tRestore).Microseconds()) / 1000.0
if o.cfg.CleanupFiles {
os.Remove(abspath)
}
if err != nil {
return out, fmt.Errorf("llama restore: %w", err)
}
out.Restored = true
out.RestoredTokens = rr.NRestored
out.RestoreMs = restoreMs
// Feed the arbiter the FULL prepare wall (lookup + store fetch + GPU load +
// orchestration). For small blobs this is dominated by fixed overhead, so
// the learned "restore bandwidth" is low and the arbiter declines next time;
// for large blobs it is transfer-dominated and restore wins. This is what
// makes the policy adaptive and loss-free across model/hardware regimes.
fullMs := float64(time.Since(prepStart).Microseconds()) / 1000.0
o.arb.ObserveGet(p.modelKey, gr.Bytes, fullMs)
o.mx.RestoreLatency.Observe(fullMs)
o.mx.RestoreCount.Inc()
o.mx.Hits.Inc()
o.mx.HitBlocks.Add(float64(m))
o.mx.MissBlocks.Add(float64(p.chain.FullBlocks - m))
o.mx.AddSavedTokens(out.MatchedTokens)
if im.Found {
// learned size correction
o.updateBytesPerTok(p.rootKey, out.MatchedTokens, gr.Bytes)
}
return out, nil
}
// Commit saves the KV currently in target.Slot and stores it under the block
// key for the largest block-aligned prefix of `tokens`. skip-if-exists gives
// single-writer dedup across concurrent agents sharing a prefix.
func (o *Orchestrator) Commit(ctx context.Context, fp cachekey.ModelFingerprint, pol Policy, tokens []int32, modelPath string, tgt *Target, prefillN int, prefillMs float64) (CommitResult, error) {
o.mx.CommitTotal.Inc()
out := CommitResult{}
if !pol.Write {
out.Reason = "write disabled"
return out, nil
}
p := o.plan(fp, pol, tokens, modelPath)
// Learn the live prefill rate (per model) so the arbiter can compare.
o.arb.ObservePrefill(p.modelKey, prefillN, prefillMs)
minB := pol.MinPrefixBlocks
if minB <= 0 {
minB = o.cfg.MinPrefixBlocks
}
if p.chain.FullBlocks < minB {
out.Reason = fmt.Sprintf("prefix %d blocks < min %d; not cached", p.chain.FullBlocks, minB)
return out, nil
}
if tgt == nil || tgt.BaseURL == "" {
out.Reason = "no target to save from"
return out, nil
}
blocks := p.chain.FullBlocks
key := p.keys[blocks-1]
out.Key = key
out.StoredBlocks = blocks
out.StoredTokens = blocks * p.fp.BlockSize
// Fast path: someone already stored this exact prefix.
if present, err := o.store.Exists(ctx, []string{key}); err == nil && len(present) == 1 && present[0] == 1 {
out.OK = true
out.Stored = false
out.Reason = "already present (dedup)"
return out, nil
}
fname := fmt.Sprintf("omb-s-%x-%d.bin", hashKey(key), o.nonce.Add(1))
abspath := filepath.Join(o.cfg.SlotSavePath, fname)
tSave := time.Now()
sr, err := o.llama.SaveSlot(ctx, tgt.BaseURL, tgt.Slot, fname)
if err != nil {
return out, fmt.Errorf("llama save: %w", err)
}
out.SaveMs = float64(time.Since(tSave).Microseconds()) / 1000.0
o.mx.SaveLatency.Observe(out.SaveMs)
repl := pol.ReplicaNum
if repl == 0 {
repl = o.cfg.DefaultReplicaNum
}
pr, err := o.store.PutFile(ctx, key, abspath, repl, pol.SoftPin, true)
if o.cfg.CleanupFiles {
os.Remove(abspath)
}
if err != nil {
return out, fmt.Errorf("store put: %w", err)
}
out.StorePutMs = pr.ElapsedMs
out.Bytes = pr.Bytes
out.OK = true
out.Stored = !pr.Existed
o.mx.StorePutLatency.Observe(pr.ElapsedMs)
if out.Stored {
o.mx.AddBytesPut(pr.Bytes)
o.index.Insert(p.rootKey, p.chain.Hex, blocks, out.StoredTokens, key, pr.Bytes)
}
// learn bytes/token from the real save (n_written / n_saved)
if sr.NSaved > 0 {
o.updateBytesPerTok(p.rootKey, sr.NSaved, sr.NWritten)
}
if out.Stored {
out.Reason = "stored"
} else {
out.Reason = "already present (dedup)"
}
return out, nil
}
// RefreshGauges pushes learned arbiter rates + index stats into Prometheus.
func (o *Orchestrator) RefreshGauges() {
s := o.arb.Snapshot()
o.mx.GetGBps.Set(s.GetGBps)
o.mx.PrefillToksS.Set(s.PrefillToksS)
is := o.index.Stats()
o.mx.IndexSnapshots.Set(float64(is.Snapshots))
o.mx.IndexBytes.Set(float64(is.Bytes))
}
func hashKey(s string) uint64 {
// FNV-1a 64
var h uint64 = 1469598103934665603
for i := 0; i < len(s); i++ {
h ^= uint64(s[i])
h *= 1099511628211
}
return h
}