Files
chorus/internal/core/router/routing_test.go
T

185 lines
6.6 KiB
Go

package router
import (
"encoding/json"
"errors"
"reflect"
"slices"
"strings"
"testing"
"git.ilapage.cn/OPC/chorus/internal/core/model"
"git.ilapage.cn/OPC/chorus/internal/core/provider"
)
type sequenceRandom struct {
values []uint64
index int
}
func (r *sequenceRandom) Uint64n(limit uint64) uint64 {
if r.index >= len(r.values) {
return 0
}
value := r.values[r.index]
r.index++
return value
}
func validSnapshot() RouteSnapshot {
return RouteSnapshot{
Capability: model.CapabilityText,
RoutePoolID: 11,
RoutePoolVersion: 3,
PromptTemplateID: 22,
PromptTemplateKey: "text-default",
PromptTemplateVersion: 7,
MaxFailover: 2,
Members: []MemberSnapshot{
{RoutePoolMemberID: 1, ProviderModelID: 101, Weight: 1, FailureThreshold: 3, OpenSeconds: 60, HalfOpenMax: 1},
{RoutePoolMemberID: 2, ProviderModelID: 102, Weight: 3, FailureThreshold: 3, OpenSeconds: 60, HalfOpenMax: 2},
{RoutePoolMemberID: 3, ProviderModelID: 103, Weight: 6, FailureThreshold: 3, OpenSeconds: 60, HalfOpenMax: 1},
},
}
}
func availableState(memberID uint64) MemberState {
return MemberState{RoutePoolMemberID: memberID, Enabled: true, ProviderEnabled: true, ModelEnabled: true, SupportsCapability: true, CircuitState: CircuitClosed}
}
func TestRouteSnapshotEncodesOnlyRoutingMetadata(t *testing.T) {
snapshot := validSnapshot()
encoded, err := snapshot.Encode()
if err != nil {
t.Fatal(err)
}
for _, forbidden := range []string{"api_key", "credential", "base_url", "nonce", "ciphertext"} {
if strings.Contains(string(encoded), forbidden) {
t.Fatalf("snapshot contains forbidden %q: %s", forbidden, encoded)
}
}
var decoded RouteSnapshot
if err := json.Unmarshal(encoded, &decoded); err != nil || !reflect.DeepEqual(decoded, snapshot) {
t.Fatalf("snapshot round trip = %#v, %v", decoded, err)
}
generation := &model.Generation{}
if err := ApplySnapshot(generation, snapshot); err != nil {
t.Fatal(err)
}
if generation.RoutePoolID == nil || *generation.RoutePoolID != snapshot.RoutePoolID || generation.RoutePoolVersion == nil || *generation.RoutePoolVersion != snapshot.RoutePoolVersion || generation.PromptTemplateID == nil || *generation.PromptTemplateID != snapshot.PromptTemplateID || !slices.Equal(generation.RouteSnapshot, encoded) {
t.Fatalf("generation route fields = %#v", generation)
}
}
func TestRouteSnapshotRejectsDuplicateModelsAndInvalidMembers(t *testing.T) {
snapshot := validSnapshot()
snapshot.Members[1].ProviderModelID = snapshot.Members[0].ProviderModelID
if err := snapshot.Validate(); !errors.Is(err, ErrInvalidSnapshot) {
t.Fatalf("duplicate model error = %v", err)
}
snapshot = validSnapshot()
snapshot.Members[0].HalfOpenMax = 0
if err := snapshot.Validate(); !errors.Is(err, ErrInvalidSnapshot) {
t.Fatalf("invalid half-open limit error = %v", err)
}
}
func TestCapabilityForSubmission(t *testing.T) {
tests := []struct {
kind model.GenerationKind
hasInputs bool
want model.Capability
}{
{model.KindText, false, model.CapabilityText},
{model.KindImage, false, model.CapabilityImageGenerate},
{model.KindImage, true, model.CapabilityImageEdit},
}
for _, test := range tests {
got, err := CapabilityForSubmission(test.kind, test.hasInputs)
if err != nil || got != test.want {
t.Fatalf("CapabilityForSubmission(%q, %t) = %q, %v; want %q", test.kind, test.hasInputs, got, err, test.want)
}
}
if _, err := CapabilityForSubmission("invalid", false); !errors.Is(err, ErrRouteNotConfigured) {
t.Fatalf("invalid capability error = %v", err)
}
}
func TestSelectCandidatesUsesWeightedSamplingWithoutReplacement(t *testing.T) {
snapshot := validSnapshot()
states := []MemberState{availableState(1), availableState(2), availableState(3)}
selected, err := SelectCandidates(snapshot, states, &sequenceRandom{values: []uint64{9, 0, 0}})
if err != nil {
t.Fatal(err)
}
got := []uint64{selected[0].RoutePoolMemberID, selected[1].RoutePoolMemberID, selected[2].RoutePoolMemberID}
if want := []uint64{3, 1, 2}; !slices.Equal(got, want) {
t.Fatalf("weighted candidates = %v, want %v", got, want)
}
snapshot.MaxFailover = 0
selected, err = SelectCandidates(snapshot, states, &sequenceRandom{values: []uint64{1}})
if err != nil || len(selected) != 1 {
t.Fatalf("max_failover=0 candidates = %#v, %v", selected, err)
}
}
func TestSelectCandidatesSkipsDisabledOpenAndCapabilityMismatchMembers(t *testing.T) {
snapshot := validSnapshot()
states := []MemberState{
availableState(1),
{RoutePoolMemberID: 2, Enabled: true, ProviderEnabled: true, ModelEnabled: true, SupportsCapability: true, CircuitState: CircuitOpen},
{RoutePoolMemberID: 3, Enabled: true, ProviderEnabled: true, ModelEnabled: true, SupportsCapability: false, CircuitState: CircuitClosed},
}
selected, err := SelectCandidates(snapshot, states, &sequenceRandom{values: []uint64{0}})
if err != nil || len(selected) != 1 || selected[0].RoutePoolMemberID != 1 {
t.Fatalf("eligible candidates = %#v, %v", selected, err)
}
if _, err := SelectCandidates(snapshot, states[1:], &sequenceRandom{values: []uint64{0}}); !errors.Is(err, ErrRouteUnavailable) {
t.Fatalf("unavailable route error = %v", err)
}
}
func TestCandidateForAttemptReturnsStableExhaustionCode(t *testing.T) {
candidates := validSnapshot().Members[:1]
member, err := CandidateForAttempt(candidates, 0)
if err != nil || member.RoutePoolMemberID != candidates[0].RoutePoolMemberID {
t.Fatalf("first candidate = %#v, %v", member, err)
}
if _, err := CandidateForAttempt(candidates, 1); !errors.Is(err, ErrFailoverExhausted) || err.Error() != string(CodeFailoverExhausted) {
t.Fatalf("exhausted candidate error = %v", err)
}
}
func TestHalfOpenEligibilityRespectsProbeLimit(t *testing.T) {
state := availableState(1)
state.CircuitState = CircuitHalfOpen
state.HalfOpenMax = 2
for inFlight, want := range map[uint16]bool{0: true, 1: true, 2: false, 3: false} {
state.HalfOpenInFlight = inFlight
if got := state.Eligible(); got != want {
t.Errorf("half-open probes %d eligible = %t, want %t", inFlight, got, want)
}
}
}
func TestFailureActionPreservesRetryableRedLines(t *testing.T) {
tests := []struct {
class provider.FailureClass
want FailureAction
}{
{provider.FailureRateLimited, FailureTryNext},
{provider.FailureServer, FailureTryNext},
{provider.FailureTimeout, FailureTryNext},
{provider.FailureConnection, FailureTryNext},
{provider.FailureBadRequest, FailureStop},
{provider.FailureUnauthorized, FailureStop},
{provider.FailurePolicyRejected, FailureStop},
}
for _, test := range tests {
if got := ActionForFailure(test.class); got != test.want {
t.Errorf("ActionForFailure(%s) = %s, want %s", test.class, got, test.want)
}
}
}