diff --git a/internal/core/key_auto_concurrency_test.go b/internal/core/key_auto_concurrency_test.go new file mode 100644 index 0000000..3d1a19f --- /dev/null +++ b/internal/core/key_auto_concurrency_test.go @@ -0,0 +1,189 @@ +package core + +import ( + "fmt" + "sync" + "sync/atomic" + "testing" + "time" + + "llmsproxy/internal/config" +) + +// AutoChainFor sits on the request hot path and reads keyAutoChains with no +// lock, while SaveKeyAuto / rebuildRegistry replace that map from the admin +// path. A "call it twice" test never interleaves those, so it would pass even +// if the map were mutated in place. These tests run both sides at once under +// -race and assert what readers actually observe. +// +// The property: a reader sees a chain that is either the key's own or the +// global one — never a torn or empty one. If rebuildKeyAutoChains built the +// map in place instead of swapping a fresh pointer, readers would sometimes +// see a partially-filled map and silently fall back to the global chain, i.e. +// a user quietly losing their configured chain under load. + +// concurrentChainConfig builds N keys, half with their own chain pointing at a +// distinct model so a fallback to the global chain is distinguishable. +func concurrentChainConfig(t *testing.T, n int) *config.Config { + t.Helper() + cfg := perKeyTestConfig(t) + cfg.Keys = nil + cfg.GatewayKeys = []string{"sk-gw-admin"} + cfg.Auto = []config.ModelScope{{Model: "gpt-4o", Source: "s1", Tier: 1}} + for i := 0; i < n; i++ { + k := config.GWKey{ + Key: fmt.Sprintf("sk-gw-k%03d", i), + Role: "user", + Name: fmt.Sprintf("user%d", i), + } + if i%2 == 0 { + k.Auto = []config.ModelScope{{Model: "gpt-4o-mini", Source: "s1", Tier: 1}} + } + cfg.Keys = append(cfg.Keys, k) + } + return cfg +} + +func TestAutoChainForConcurrentWithSaveKeyAuto(t *testing.T) { + cfg := concurrentChainConfig(t, 40) + c := newTestCore(t, cfg) + + var stop atomic.Bool + var wg sync.WaitGroup + var bad atomic.Int64 + + for r := 0; r < 8; r++ { + wg.Add(1) + go func() { + defer wg.Done() + for !stop.Load() { + for i := 0; i < 40; i++ { + key := fmt.Sprintf("sk-gw-k%03d", i) + ch, ok := c.AutoChainFor(key) + if !ok || ch == nil || len(ch.Tiers) == 0 { + bad.Add(1) + continue + } + // Legal values are this key's own chain model (either + // side of the writer's flip) or the global chain's. The + // assertion is about the legal set, not a fixed value: + // the writers legitimately change the chain underneath. + got := ch.Tiers[0].Slots[0].Model + if got != "gpt-4o" && got != "gpt-4o-mini" { + bad.Add(1) + } + } + } + }() + } + + for w := 0; w < 2; w++ { + wg.Add(1) + go func() { + defer wg.Done() + model := "gpt-4o-mini" + for i := 0; i < 25 && !stop.Load(); i++ { + for j := 0; j < 40; j += 2 { + if err := c.SaveKeyAuto(fmt.Sprintf("sk-gw-k%03d", j), []config.ModelScope{ + {Model: model, Source: "s1", Tier: 1}, + }); err != nil { + bad.Add(1) + return + } + } + if model == "gpt-4o-mini" { + model = "gpt-4o" + } else { + model = "gpt-4o-mini" + } + } + }() + } + + time.Sleep(400 * time.Millisecond) + stop.Store(true) + wg.Wait() + + if n := bad.Load(); n > 0 { + t.Fatalf("%d reads saw a chain outside {own, global} — the cache is "+ + "being swapped in place instead of atomically", n) + } +} + +// The registry rebuild path (AddSource / RemoveSource) also replaces the +// cache, so readers must keep a usable chain throughout. +// +// Churned sources use a "churn" prefix on purpose: deleting a source a chain +// references legitimately empties that chain (BuildChain drops unresolvable +// slots), and asserting otherwise would be testing a fiction rather than the +// code. +func TestAutoChainForConcurrentWithSourceChanges(t *testing.T) { + cfg := concurrentChainConfig(t, 20) + c := newTestCore(t, cfg) + + var stop atomic.Bool + var wg sync.WaitGroup + var bad atomic.Int64 + + for r := 0; r < 6; r++ { + wg.Add(1) + go func() { + defer wg.Done() + for !stop.Load() { + for i := 0; i < 20; i++ { + ch, ok := c.AutoChainFor(fmt.Sprintf("sk-gw-k%03d", i)) + if !ok || ch == nil || len(ch.Tiers) == 0 { + bad.Add(1) + continue + } + if got := ch.Tiers[0].Slots[0].Model; got != "gpt-4o" && got != "gpt-4o-mini" { + bad.Add(1) + } + } + } + }() + } + + wg.Add(1) + go func() { + defer wg.Done() + for n := 0; n < 30 && !stop.Load(); n++ { + name := fmt.Sprintf("churn%d", n) + _ = c.AddSource(config.Source{ + Name: name, BaseURL: "http://127.0.0.1:9/v1", Adapter: "openai", + Models: []config.Model{{ID: "m-" + name, Priority: 10}}}, + ) + _ = c.RemoveSource(name) + } + }() + + time.Sleep(500 * time.Millisecond) + stop.Store(true) + wg.Wait() + + if n := bad.Load(); n > 0 { + t.Fatalf("%d reads lost the chain while unrelated sources were churned", n) + } +} + +// Many distinct keys: the lookup must stay a map hit and must not degrade into +// returning the global chain as the key count grows. +func TestAutoChainForManyKeys(t *testing.T) { + cfg := concurrentChainConfig(t, 500) + c := newTestCore(t, cfg) + + for i := 0; i < 500; i++ { + key := fmt.Sprintf("sk-gw-k%03d", i) + ch, _ := c.AutoChainFor(key) + if ch == nil || len(ch.Tiers) == 0 { + t.Fatalf("key %s lost its chain with 500 keys configured", key) + } + want := "gpt-4o" + if i%2 == 0 { + want = "gpt-4o-mini" + } + if got := ch.Tiers[0].Slots[0].Model; got != want { + t.Fatalf("key %s got %q, want %q", key, got, want) + } + } +}