185 lines
6.6 KiB
Go
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)
|
|
}
|
|
}
|
|
}
|