iop/apps/edge/internal/openai/anthropic_native_test.go

343 lines
16 KiB
Go

package openai
import (
"bytes"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"reflect"
"strings"
"testing"
"time"
"iop/apps/edge/internal/authprojection"
edgeservice "iop/apps/edge/internal/service"
"iop/packages/go/config"
iop "iop/proto/gen/iop"
)
func TestAnthropicNativeProviderFixturesPreserveBytesAndHeaders(t *testing.T) {
response := mustReadAnthropicFixture(t, "native_message.json")
requestBody := []byte("{\n \"model\" : \"claude-route\",\n \"max_tokens\":64,\n \"messages\":[{\"role\":\"user\",\"content\":[{\"type\":\"text\",\"text\":\"keep formatting\"}]}],\n \"vendor_extension\": {\"enabled\": true}\n}")
for _, tc := range []struct {
profile string
authHeader string
authValue string
}{
{profile: "anthropic", authHeader: "X-Api-Key", authValue: "provider-secret"},
{profile: "minimax_messages", authHeader: "Authorization", authValue: "Bearer provider-secret"},
{profile: "mimo_messages", authHeader: "Api-Key", authValue: "provider-secret"},
{profile: "seulgi_messages", authHeader: "X-Api-Key", authValue: "provider-secret"},
} {
t.Run(tc.profile, func(t *testing.T) {
fake := &providerFakeRunService{
poolDispatchPath: string(edgeservice.ProviderPoolPathTunnel),
poolSelectedCandidate: anthropicTestCandidate(t, tc.profile),
tunnelServedTarget: "upstream-claude",
tunnelFrames: anthropicTunnelFrames(http.StatusOK, "application/json", response[:19], response[19:73], response[73:]),
}
cfg := config.EdgeOpenAIConf{
BearerToken: "iop-secret",
ProviderAuth: config.EdgeOpenAIProviderAuthConf{
Enabled: true, FromHeader: "X-Provider-Token", TargetHeader: "Authorization", Scheme: "Bearer", Required: true,
},
}
srv := NewServer(cfg, fake, nil)
srv.SetModelCatalog([]config.ModelCatalogEntry{{ID: "claude-route", Providers: map[string]string{"provider": "upstream-claude"}}})
req := httptest.NewRequest(http.MethodPost, "/v1/messages", bytes.NewReader(requestBody))
req.Header.Set("Authorization", "Bearer iop-secret")
req.Header.Set("X-Provider-Token", "provider-secret")
req.Header.Set(anthropicVersionHeader, anthropicSupportedVersion)
req.Header.Set(anthropicBetaHeader, "prompt-caching-2024-07-31")
req.Header.Set("User-Agent", "must-not-forward")
w := httptest.NewRecorder()
srv.routes().ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", w.Code, w.Body.String())
}
if !bytes.Equal(w.Body.Bytes(), response) {
t.Fatalf("response bytes changed:\n got: %q\nwant: %q", w.Body.Bytes(), response)
}
requests := fake.tunnelReqsSnapshot()
bodies := fake.tunnelBodiesSnapshot()
if len(requests) != 1 || len(bodies) != 1 {
t.Fatalf("wire request evidence missing: requests=%d bodies=%d", len(requests), len(bodies))
}
if requests[0].Operation != string(config.OperationMessages) {
t.Fatalf("operation=%q", requests[0].Operation)
}
if got := requests[0].Headers[tc.authHeader]; got != tc.authValue {
t.Fatalf("profile auth header %s=%q, want %q; all=%v", tc.authHeader, got, tc.authValue, requests[0].Headers)
}
if requests[0].Headers["User-Agent"] != "" || requests[0].Headers["X-Provider-Token"] != "" {
t.Fatalf("non-allowlisted request header forwarded: %v", requests[0].Headers)
}
if requests[0].Headers[anthropicVersionHeader] != anthropicSupportedVersion || requests[0].Headers[anthropicBetaHeader] == "" {
t.Fatalf("Anthropic compatibility headers missing: %v", requests[0].Headers)
}
expectedBody := bytes.Replace(requestBody, []byte(`"claude-route"`), []byte(`"upstream-claude"`), 1)
if !bytes.Equal(bodies[0], expectedBody) {
t.Fatalf("native request bytes changed outside model value:\n got=%q\nwant=%q", bodies[0], expectedBody)
}
var gotBody map[string]any
if err := json.Unmarshal(bodies[0], &gotBody); err != nil {
t.Fatal(err)
}
if gotBody["model"] != "upstream-claude" {
t.Fatalf("model rewrite=%v", gotBody["model"])
}
var original map[string]any
if err := json.Unmarshal(requestBody, &original); err != nil {
t.Fatal(err)
}
original["model"] = "upstream-claude"
if !reflect.DeepEqual(gotBody, original) {
t.Fatalf("native body semantics changed:\n got=%v\nwant=%v", gotBody, original)
}
})
}
}
func TestAnthropicNativeStreamPreservesFragmentOrderAndSingleTerminal(t *testing.T) {
streamBody := mustReadAnthropicFixture(t, "native_stream.sse")
frames := make(chan *iop.ProviderTunnelFrame, 8)
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: http.StatusOK, Headers: map[string]string{"Content-Type": "text/event-stream", "Connection": "must-not-copy"}}
for index, fragment := range splitAnthropicFixture(streamBody, 17, 81, 143, len(streamBody)-9) {
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Sequence: int64(index + 1), Body: fragment}
}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte("must-not-appear")}
close(frames)
fake := &providerFakeRunService{
poolDispatchPath: string(edgeservice.ProviderPoolPathTunnel),
poolSelectedCandidate: anthropicTestCandidate(t, "anthropic"),
tunnelServedTarget: "upstream-claude",
tunnelFrames: frames,
}
srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
srv.SetModelCatalog([]config.ModelCatalogEntry{{ID: "claude-route", Providers: map[string]string{"provider": "upstream-claude"}}})
req := httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(`{"model":"claude-route","max_tokens":32,"stream":true,"messages":[{"role":"user","content":"hello"}]}`))
req.Header.Set(anthropicVersionHeader, anthropicSupportedVersion)
w := httptest.NewRecorder()
srv.routes().ServeHTTP(w, req)
if w.Code != http.StatusOK || !bytes.Equal(w.Body.Bytes(), streamBody) {
t.Fatalf("native stream changed: status=%d\n got=%q\nwant=%q", w.Code, w.Body.Bytes(), streamBody)
}
if strings.Count(w.Body.String(), "event: message_stop") != 1 || strings.Contains(w.Body.String(), "must-not-appear") {
t.Fatalf("native terminal was not exactly once: %s", w.Body.String())
}
if got := w.Header().Get("Connection"); got != "" {
t.Fatalf("hop-by-hop response header leaked: %q", got)
}
}
func TestAnthropicNativeProviderErrorPreservesStatusAndBody(t *testing.T) {
errorBody := mustReadAnthropicFixture(t, "native_error.json")
frames := make(chan *iop.ProviderTunnelFrame, 5)
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: http.StatusTooManyRequests, Headers: map[string]string{"Content-Type": "application/json"}}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Sequence: 1, Body: errorBody[:11]}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Sequence: 2, Body: errorBody[11:]}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Sequence: 3, Body: []byte("must-not-appear")}
close(frames)
fake := &providerFakeRunService{
poolDispatchPath: string(edgeservice.ProviderPoolPathTunnel),
poolSelectedCandidate: anthropicTestCandidate(t, "anthropic"),
tunnelFrames: frames,
}
srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
srv.SetModelCatalog([]config.ModelCatalogEntry{{ID: "claude-route", Providers: map[string]string{"provider": "served"}}})
w := serveAnthropicRequest(srv, "/v1/messages", `{"model":"claude-route","max_tokens":16,"messages":[{"role":"user","content":"hello"}]}`)
if w.Code != http.StatusTooManyRequests || !bytes.Equal(w.Body.Bytes(), errorBody) {
t.Fatalf("native error changed: status=%d body=%q want=%q", w.Code, w.Body.Bytes(), errorBody)
}
}
func TestAnthropicNativeVirtualPresetPreservesPublicModelIdentity(t *testing.T) {
const (
virtualModelID = "virtual-public-model"
canonicalModel = "canonical-selector-model"
projectedRoute = "projected-selector-route"
credentialSlot = "selector-slot"
providerID = "provider-resource"
servedModel = "served-selector-model"
)
now := time.Date(2026, 8, 2, 12, 0, 0, 0, time.UTC)
preset := config.ExecutionPreset{
ID: "preset-native-public-identity",
Selector: config.ExecutionModelBinding{Model: canonicalModel},
AllowedModes: []string{config.ModeDirect},
Routes: map[string]config.ExecutionRoute{config.ModeDirect: {}},
}
newServer := func(t *testing.T, frames chan *iop.ProviderTunnelFrame) (*Server, *providerFakeRunService) {
t.Helper()
candidate := anthropicTestCandidate(t, "anthropic")
candidate.ProviderID = providerID
candidate.ActualModel = servedModel
route := authprojection.Route{
RouteID: projectedRoute, PrincipalRef: "principal-1", CredentialSlotRef: credentialSlot,
ProfileID: candidate.ProfileID, UpstreamModel: servedModel, ResourceSelector: providerID,
}
cache := authprojection.NewCache(authprojection.DefaultLimits(), func() time.Time { return now })
projection := makeTestProjection(1, now, time.Hour, map[string]string{"managed-token": "principal-1"}, map[string]authprojection.Route{"selector": route})
if err := cache.Apply(projection); err != nil {
t.Fatal(err)
}
fake := &providerFakeRunService{
poolDispatchPath: string(edgeservice.ProviderPoolPathTunnel),
poolSelectedCandidate: candidate,
tunnelFrames: frames,
}
srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
srv.SetEdgeID("edge-native-public-identity")
setManagedPrincipalProjection(srv, cache)
srv.SetExecutionPresets([]config.ExecutionPreset{preset})
srv.SetModelCatalog([]config.ModelCatalogEntry{
{ID: virtualModelID, ExecutionPreset: preset.ID},
{ID: canonicalModel, Providers: map[string]string{providerID: servedModel}},
})
return srv, fake
}
assertSelectorBinding := func(t *testing.T, fake *providerFakeRunService) {
t.Helper()
runs := fake.tunnelReqsSnapshot()
if len(runs) != 1 {
t.Fatalf("tunnel requests=%d, want 1", len(runs))
}
binding := runs[0].CredentialBinding
if binding == nil || binding.RouteID != projectedRoute || binding.CredentialSlotRef != credentialSlot {
t.Fatalf("credential binding=%+v, want projected selector route %q", binding, projectedRoute)
}
}
serve := func(t *testing.T, srv *Server, stream bool) *httptest.ResponseRecorder {
t.Helper()
body := fmt.Sprintf(`{"model":"virtual-public-model","max_tokens":8,"messages":[{"role":"user","content":"hi"}],"stream":%t}`, stream)
req := httptest.NewRequest(http.MethodPost, "/v1/messages", strings.NewReader(body))
req.Header.Set("Authorization", "Bearer managed-token")
req.Header.Set(anthropicVersionHeader, anthropicSupportedVersion)
w := httptest.NewRecorder()
srv.routes().ServeHTTP(w, req)
return w
}
t.Run("non-stream JSON", func(t *testing.T) {
body := []byte(`{"id":"msg-public","type":"message","role":"assistant","model":"served-selector-model","content":[{"type":"text","text":"ok"}],"stop_reason":"end_turn"}`)
frames := make(chan *iop.ProviderTunnelFrame, 5)
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: http.StatusOK, Headers: map[string]string{"Content-Type": "application/json", "Content-Length": "999"}}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: body[:23]}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: body[23:71]}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: body[71:]}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true}
close(frames)
srv, fake := newServer(t, frames)
w := serve(t, srv, false)
if w.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", w.Code, w.Body.String())
}
var response anthropicMessageResponse
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatal(err)
}
if response.ID != "msg-public" {
t.Fatalf("response id=%q, want exact provider ID %q", response.ID, "msg-public")
}
if response.Model != virtualModelID {
t.Fatalf("response model=%q, want %q", response.Model, virtualModelID)
}
if got := w.Header().Get("Content-Length"); got != "" {
t.Fatalf("content length=%q, want removed after rewrite", got)
}
assertSelectorBinding(t, fake)
assertHotPathTerminal(t, srv)
})
t.Run("fragmented SSE", func(t *testing.T) {
stream := []byte("event: message_start\r\ndata: {\"type\":\"message_start\",\"message\":{\"id\":\"msg-public\",\"type\":\"message\",\"role\":\"assistant\",\"model\":\"served-selector-model\",\"content\":[]}}\r\n\r\nevent: content_block_delta\r\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"ok\"}}\r\n\r\nevent: message_stop\r\ndata: {\"type\":\"message_stop\"}\r\n\r\n")
modelAt := bytes.Index(stream, []byte(servedModel))
if modelAt < 0 {
t.Fatal("served model missing from fixture")
}
fragments := splitAnthropicFixture(stream, 31, modelAt+7, modelAt+len(servedModel)-4, len(stream)-18)
frames := anthropicTunnelFrames(http.StatusOK, "text/event-stream", fragments...)
srv, fake := newServer(t, frames)
w := serve(t, srv, true)
if w.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", w.Code, w.Body.String())
}
if !strings.Contains(w.Body.String(), "event: message_start") ||
!strings.Contains(w.Body.String(), `"id":"msg-public"`) ||
!strings.Contains(w.Body.String(), `"model":"virtual-public-model"`) ||
strings.Contains(w.Body.String(), servedModel) {
t.Fatalf("direct stream did not preserve public identity: %s", w.Body.String())
}
if strings.Count(w.Body.String(), "event: message_stop") != 1 {
t.Fatalf("message stop count=%d, want 1", strings.Count(w.Body.String(), "event: message_stop"))
}
assertSelectorBinding(t, fake)
assertNoReservedPath(t, w.Body.String())
assertHotPathTerminal(t, srv)
})
for _, tc := range []struct {
name string
frames []*iop.ProviderTunnelFrame
wantStatus int
}{
{
name: "END before response start fails closed",
frames: []*iop.ProviderTunnelFrame{
{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true},
},
wantStatus: http.StatusBadGateway,
},
{
name: "BODY before response start fails closed",
frames: []*iop.ProviderTunnelFrame{
{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte(`{"model":"served-selector-model"}`)},
{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true},
},
wantStatus: http.StatusBadGateway,
},
} {
t.Run(tc.name, func(t *testing.T) {
frames := make(chan *iop.ProviderTunnelFrame, len(tc.frames))
for _, frame := range tc.frames {
frames <- frame
}
close(frames)
srv, fake := newServer(t, frames)
w := serve(t, srv, false)
if w.Code != tc.wantStatus || !strings.Contains(w.Body.String(), `"type":"api_error"`) ||
strings.Contains(w.Body.String(), servedModel) || strings.Contains(w.Body.String(), "run-") {
t.Fatalf("status=%d body=%q want sanitized status=%d api_error", w.Code, w.Body.Bytes(), tc.wantStatus)
}
assertSelectorBinding(t, fake)
assertHotPathTerminal(t, srv)
})
}
}
func splitAnthropicFixture(body []byte, offsets ...int) [][]byte {
parts := make([][]byte, 0, len(offsets)+1)
start := 0
for _, offset := range offsets {
if offset <= start || offset >= len(body) {
continue
}
parts = append(parts, append([]byte(nil), body[start:offset]...))
start = offset
}
parts = append(parts, append([]byte(nil), body[start:]...))
return parts
}