summaryrefslogtreecommitdiff
path: root/internal/providerhealth/prober_test.go
blob: 69b6e807718ad39b3bc16c7f4a3df9e5af919eef (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
package providerhealth

import (
	"context"
	"io"
	"log/slog"
	"net/http"
	"net/http/httptest"
	"sync/atomic"
	"testing"
	"time"

	"aigw/internal/catalog"
	"aigw/internal/domain"
)

type probeMetricCounter struct {
	total  atomic.Int64
	failed atomic.Int64
}

func (m *probeMetricCounter) ProviderProbe(success bool) {
	m.total.Add(1)
	if !success {
		m.failed.Add(1)
	}
}

func TestProberAuthenticatesDeduplicatesAndOpensRouteCircuits(t *testing.T) {
	var calls atomic.Int64
	server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		calls.Add(1)
		if r.URL.Path != "/v1/models" || r.Header.Get("Authorization") != "Bearer secret" {
			t.Errorf("unexpected probe path=%s auth=%q", r.URL.Path, r.Header.Get("Authorization"))
		}
		http.Error(w, "unavailable", http.StatusServiceUnavailable)
	}))
	defer server.Close()

	provider := domain.Provider{ID: "provider", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: server.URL + "/v1", APIKey: "secret"}
	model := domain.Model{ID: "public/model", Routes: []domain.Route{{Provider: provider, UpstreamModel: "one"}, {Provider: provider, UpstreamModel: "two"}}}
	tracker := New(Options{FailureThreshold: 1, OpenDuration: time.Minute})
	metrics := &probeMetricCounter{}
	prober := NewProber(ProbeOptions{Enabled: true, Catalog: catalog.NewModels([]domain.Model{model}), Tracker: tracker,
		Metrics: metrics, Timeout: time.Second, Logger: slog.New(slog.NewTextHandler(io.Discard, nil))})
	prober.ProbeOnce(context.Background())

	key := RouteKey{ModelID: model.ID, ProviderID: provider.ID, WireAPI: provider.WireAPI}
	if calls.Load() != 1 || metrics.total.Load() != 1 || metrics.failed.Load() != 1 || !tracker.CircuitOpen(key) {
		t.Fatalf("calls=%d total=%d failed=%d open=%v", calls.Load(), metrics.total.Load(), metrics.failed.Load(), tracker.CircuitOpen(key))
	}
	status := tracker.Snapshot()[0]
	if status.ActiveProbes != 1 || status.LastProbeAt == nil || status.LastStatusCode != http.StatusServiceUnavailable {
		t.Fatalf("unexpected probe status: %+v", status)
	}
}

func TestProberRejectsHTMLSuccessResponse(t *testing.T) {
	server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
		w.Header().Set("Content-Type", "text/html")
		_, _ = io.WriteString(w, "<html>provider console</html>")
	}))
	defer server.Close()
	provider := domain.Provider{ID: "provider", Protocol: domain.ProtocolOpenAI, BaseURL: server.URL, APIKey: "secret"}
	model := domain.Model{ID: "model", Routes: []domain.Route{{Provider: provider, UpstreamModel: "model"}}}
	tracker := New(Options{FailureThreshold: 1})
	NewProber(ProbeOptions{Enabled: true, Catalog: catalog.NewModels([]domain.Model{model}), Tracker: tracker, Timeout: time.Second,
		Logger: slog.New(slog.NewTextHandler(io.Discard, nil))}).ProbeOnce(context.Background())
	if !tracker.CircuitOpen(RouteKey{ModelID: model.ID, ProviderID: provider.ID, WireAPI: "chat_completions"}) {
		t.Fatal("HTML success response must not be treated as a healthy API probe")
	}
}

func TestProberUsesAnthropicAuthentication(t *testing.T) {
	server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		if r.Header.Get("x-api-key") != "anthropic-secret" || r.Header.Get("anthropic-version") == "" {
			t.Errorf("unexpected Anthropic headers")
		}
		w.Header().Set("Content-Type", "application/json")
		_, _ = io.WriteString(w, `{"data":[]}`)
	}))
	defer server.Close()
	provider := domain.Provider{ID: "anthropic", Protocol: domain.ProtocolAnthropic, WireAPI: "messages", BaseURL: server.URL + "/v1", APIKey: "anthropic-secret"}
	model := domain.Model{ID: "anthropic/model", Routes: []domain.Route{{Provider: provider, UpstreamModel: "model"}}}
	tracker := New(Options{})
	NewProber(ProbeOptions{Enabled: true, Catalog: catalog.NewModels([]domain.Model{model}), Tracker: tracker, Timeout: time.Second}).ProbeOnce(context.Background())
	status := tracker.Snapshot()[0]
	if status.State != "healthy" || status.ActiveProbes != 1 {
		t.Fatalf("unexpected status: %+v", status)
	}
}