vibe-proxy/backend/sdk/api/handlers/handlers_interceptors_test.go
2026-08-24 00:10:41 +02:00

1483 lines
57 KiB
Go

package handlers
import (
"context"
"fmt"
"net/http"
"net/http/httptest"
"net/url"
"sync"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/router-for-me/CLIProxyAPI/v7/internal/interfaces"
"github.com/router-for-me/CLIProxyAPI/v7/internal/logging"
"github.com/router-for-me/CLIProxyAPI/v7/internal/registry"
coreauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth"
coreexecutor "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/executor"
sdkconfig "github.com/router-for-me/CLIProxyAPI/v7/sdk/config"
"github.com/router-for-me/CLIProxyAPI/v7/sdk/pluginapi"
)
type handlerInterceptorTestHost struct {
interceptRequestBeforeAuth func(context.Context, pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse
interceptRequestAfterAuth func(context.Context, pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse
interceptResponse func(context.Context, pluginapi.ResponseInterceptRequest) pluginapi.ResponseInterceptResponse
interceptStreamChunk func(context.Context, pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse
completeRequest func(context.Context, pluginapi.RequestCompletion)
// includeStreamChunkRequestBodies simulates legacy schema_version < 3 plugins.
includeStreamChunkRequestBodies bool
}
type handlerInterceptorNoStreamTestHost struct {
*handlerInterceptorTestHost
}
type handlerInterceptorDisabledRequestTestHost struct {
*handlerInterceptorTestHost
}
func (h *handlerInterceptorNoStreamTestHost) HasStreamInterceptors() bool {
return false
}
func (h *handlerInterceptorDisabledRequestTestHost) HasRequestInterceptors() bool {
return false
}
func (h *handlerInterceptorTestHost) InterceptRequestBeforeAuth(ctx context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
if h != nil && h.interceptRequestBeforeAuth != nil {
return h.interceptRequestBeforeAuth(ctx, req)
}
return pluginapi.RequestInterceptResponse{
Headers: cloneHeader(req.Headers),
Body: cloneBytes(req.Body),
}
}
func (h *handlerInterceptorTestHost) InterceptRequestAfterAuth(ctx context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
if h != nil && h.interceptRequestAfterAuth != nil {
return h.interceptRequestAfterAuth(ctx, req)
}
return pluginapi.RequestInterceptResponse{
Headers: cloneHeader(req.Headers),
Body: cloneBytes(req.Body),
}
}
func (h *handlerInterceptorTestHost) InterceptResponse(ctx context.Context, req pluginapi.ResponseInterceptRequest) pluginapi.ResponseInterceptResponse {
if h != nil && h.interceptResponse != nil {
return h.interceptResponse(ctx, req)
}
return pluginapi.ResponseInterceptResponse{
Headers: cloneHeader(req.ResponseHeaders),
Body: cloneBytes(req.Body),
}
}
func (h *handlerInterceptorTestHost) InterceptStreamChunk(ctx context.Context, req pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse {
if h != nil && h.interceptStreamChunk != nil {
return h.interceptStreamChunk(ctx, req)
}
return pluginapi.StreamChunkInterceptResponse{
Headers: cloneHeader(req.ResponseHeaders),
Body: cloneBytes(req.Body),
}
}
func (h *handlerInterceptorTestHost) CompleteRequest(ctx context.Context, completion pluginapi.RequestCompletion) {
if h != nil && h.completeRequest != nil {
h.completeRequest(ctx, completion)
}
}
// StreamChunkPayloadIncludesRequestBody implements streamChunkRequestBodyPolicy.
// Default false simulates schema_version >= 3 (omit request bodies on payload chunks).
func (h *handlerInterceptorTestHost) StreamChunkPayloadIncludesRequestBody() bool {
if h == nil {
return false
}
return h.includeStreamChunkRequestBodies
}
type interceptorCaptureExecutor struct {
provider string
mu sync.Mutex
lastRequest coreexecutor.Request
lastOptions coreexecutor.Options
execute func(context.Context, *coreauth.Auth, coreexecutor.Request, coreexecutor.Options) (coreexecutor.Response, error)
executeCount func(context.Context, *coreauth.Auth, coreexecutor.Request, coreexecutor.Options) (coreexecutor.Response, error)
stream func(context.Context, *coreauth.Auth, coreexecutor.Request, coreexecutor.Options) (*coreexecutor.StreamResult, error)
}
func (e *interceptorCaptureExecutor) Identifier() string {
if e.provider != "" {
return e.provider
}
return "codex"
}
func (e *interceptorCaptureExecutor) Execute(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (coreexecutor.Response, error) {
e.capture(req, opts)
if e.execute != nil {
return e.execute(ctx, auth, req, opts)
}
return coreexecutor.Response{Payload: []byte("ok")}, nil
}
func (e *interceptorCaptureExecutor) ExecuteStream(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
e.capture(req, opts)
if e.stream != nil {
return e.stream(ctx, auth, req, opts)
}
chunks := make(chan coreexecutor.StreamChunk)
close(chunks)
return &coreexecutor.StreamResult{Chunks: chunks}, nil
}
func (e *interceptorCaptureExecutor) Refresh(ctx context.Context, auth *coreauth.Auth) (*coreauth.Auth, error) {
return auth, nil
}
func (e *interceptorCaptureExecutor) CountTokens(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (coreexecutor.Response, error) {
e.capture(req, opts)
if e.executeCount != nil {
return e.executeCount(ctx, auth, req, opts)
}
return coreexecutor.Response{Payload: []byte("0")}, nil
}
func (e *interceptorCaptureExecutor) HttpRequest(ctx context.Context, auth *coreauth.Auth, req *http.Request) (*http.Response, error) {
return nil, &coreauth.Error{Code: "not_implemented", Message: "HttpRequest not implemented", HTTPStatus: http.StatusNotImplemented}
}
func (e *interceptorCaptureExecutor) capture(req coreexecutor.Request, opts coreexecutor.Options) {
e.mu.Lock()
defer e.mu.Unlock()
e.lastRequest = coreexecutor.Request{
Model: req.Model,
Payload: cloneBytes(req.Payload),
Format: req.Format,
Metadata: req.Metadata,
}
e.lastOptions = coreexecutor.Options{
Stream: opts.Stream,
Alt: opts.Alt,
Headers: cloneHeader(opts.Headers),
Query: opts.Query,
OriginalRequest: cloneBytes(opts.OriginalRequest),
SourceFormat: opts.SourceFormat,
Metadata: opts.Metadata,
}
}
func (e *interceptorCaptureExecutor) captured() (coreexecutor.Request, coreexecutor.Options) {
e.mu.Lock()
defer e.mu.Unlock()
return e.lastRequest, e.lastOptions
}
func newInterceptorHandler(t *testing.T, model string, executor *interceptorCaptureExecutor, cfg *sdkconfig.SDKConfig) *BaseAPIHandler {
t.Helper()
manager := coreauth.NewManager(nil, nil, nil)
manager.RegisterExecutor(executor)
auth := &coreauth.Auth{
ID: "handler-interceptor-" + model,
Provider: executor.Identifier(),
Status: coreauth.StatusActive,
Metadata: map[string]any{"email": model + "@example.com"},
}
if _, errRegister := manager.Register(context.Background(), auth); errRegister != nil {
t.Fatalf("manager.Register(): %v", errRegister)
}
registry.GetGlobalRegistry().RegisterClient(auth.ID, auth.Provider, []*registry.ModelInfo{{ID: model}})
t.Cleanup(func() {
registry.GetGlobalRegistry().UnregisterClient(auth.ID)
})
return NewBaseAPIHandlers(cfg, manager)
}
func contextWithHeaders(headers http.Header) context.Context {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
for key, values := range headers {
for _, value := range values {
c.Request.Header.Add(key, value)
}
}
return context.WithValue(context.Background(), "gin", c)
}
// contextWithQuery builds a context whose embedded gin request carries the given
// query parameters, mirroring how plain HTTP requests expose inbound query to
// queryFromContext.
func contextWithQuery(query url.Values) context.Context {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
target := "/v1/chat/completions"
if encoded := query.Encode(); encoded != "" {
target = target + "?" + encoded
}
c.Request = httptest.NewRequest(http.MethodPost, target, nil)
return context.WithValue(context.Background(), "gin", c)
}
func TestRequestLifecycleTrackerUsesUniqueExecutionIDs(t *testing.T) {
handler := NewBaseAPIHandlers(nil, nil)
ctx := logging.WithRequestID(context.Background(), "trace-1")
first := handler.newRequestLifecycleTracker(ctx, "openai", "model", "model", false, nil, "")
second := handler.newRequestLifecycleTracker(ctx, "openai", "model", "model", false, nil, "")
if first.requestID() == "" || second.requestID() == "" || first.requestID() == second.requestID() {
t.Fatalf("lifecycle request IDs = %q and %q", first.requestID(), second.requestID())
}
if first.completion.TraceID != "trace-1" || second.completion.TraceID != "trace-1" {
t.Fatalf("trace IDs = %q and %q", first.completion.TraceID, second.completion.TraceID)
}
}
func TestHandlerRequestInterceptorTerminatesBeforeAuth(t *testing.T) {
model := "handler-interceptor-terminate-before-auth"
executor := &interceptorCaptureExecutor{}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{})
var requestID string
var completion pluginapi.RequestCompletion
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptRequestBeforeAuth: func(_ context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
requestID = req.RequestID
return pluginapi.RequestInterceptResponse{
Terminate: true,
StatusCode: http.StatusForbidden,
ResponseHeaders: http.Header{"Content-Type": {"application/json"}, "X-Policy": {"blocked"}},
ResponseBody: []byte(`{"error":"blocked"}`),
}
},
completeRequest: func(_ context.Context, got pluginapi.RequestCompletion) {
completion = got
},
})
body, headers, errMsg := handler.ExecuteWithAuthManager(context.Background(), "openai", model, []byte(`{"model":"`+model+`"}`), "")
if body != nil || headers != nil {
t.Fatalf("terminated response body = %q, headers = %#v", body, headers)
}
if errMsg == nil || !errMsg.DirectResponse || errMsg.StatusCode != http.StatusForbidden {
t.Fatalf("termination error = %#v", errMsg)
}
if string(errMsg.Body) != `{"error":"blocked"}` || errMsg.Headers.Get("X-Policy") != "blocked" {
t.Fatalf("termination response = body %q, headers %#v", errMsg.Body, errMsg.Headers)
}
if requestID == "" || completion.RequestID != requestID {
t.Fatalf("request IDs = start %q, completion %q", requestID, completion.RequestID)
}
if completion.Outcome != pluginapi.RequestCompletionRejected || completion.StatusCode != http.StatusForbidden {
t.Fatalf("completion = %#v", completion)
}
capturedReq, _ := executor.captured()
if capturedReq.Model != "" {
t.Fatalf("executor received terminated request: %#v", capturedReq)
}
}
func TestHandlerRequestInterceptorTerminatesAfterAuth(t *testing.T) {
model := "handler-interceptor-terminate-after-auth"
executor := &interceptorCaptureExecutor{}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{})
var beforeRequestID string
var afterRequestID string
var afterCalls int
var completion pluginapi.RequestCompletion
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptRequestBeforeAuth: func(_ context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
beforeRequestID = req.RequestID
return pluginapi.RequestInterceptResponse{Headers: req.Headers, Body: req.Body}
},
interceptRequestAfterAuth: func(_ context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
afterCalls++
afterRequestID = req.RequestID
return pluginapi.RequestInterceptResponse{
Terminate: true,
StatusCode: http.StatusTooManyRequests,
ResponseHeaders: http.Header{"Retry-After": {"3"}},
ResponseBody: []byte(`{"error":"busy"}`),
}
},
completeRequest: func(_ context.Context, got pluginapi.RequestCompletion) {
completion = got
},
})
_, _, errMsg := handler.ExecuteWithAuthManager(context.Background(), "openai", model, []byte(`{"model":"`+model+`"}`), "")
if errMsg == nil || !errMsg.DirectResponse || errMsg.StatusCode != http.StatusTooManyRequests {
t.Fatalf("termination error = %#v", errMsg)
}
if beforeRequestID == "" || afterRequestID != beforeRequestID || completion.RequestID != beforeRequestID {
t.Fatalf("request IDs = before %q, after %q, completion %q", beforeRequestID, afterRequestID, completion.RequestID)
}
if completion.Outcome != pluginapi.RequestCompletionRejected {
t.Fatalf("completion outcome = %q", completion.Outcome)
}
if afterCalls != 1 {
t.Fatalf("after-auth interceptor calls = %d, want 1", afterCalls)
}
capturedReq, _ := executor.captured()
if capturedReq.Model != "" {
t.Fatalf("executor received terminated request: %#v", capturedReq)
}
}
func TestHandlerAfterAuthTerminationSkipsCountAndStreamExecutors(t *testing.T) {
for _, operation := range []string{"count", "stream"} {
t.Run(operation, func(t *testing.T) {
model := "handler-interceptor-terminate-" + operation
executor := &interceptorCaptureExecutor{}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{})
afterCalls := 0
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptRequestAfterAuth: func(_ context.Context, _ pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
afterCalls++
return pluginapi.RequestInterceptResponse{
Terminate: true,
StatusCode: http.StatusForbidden,
ResponseBody: []byte(`{"error":"blocked"}`),
}
},
})
var errMsg *interfaces.ErrorMessage
if operation == "count" {
_, _, errMsg = handler.ExecuteCountWithAuthManager(context.Background(), "openai", model, []byte(`{"model":"`+model+`"}`), "")
} else {
dataChan, _, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(`{"model":"`+model+`","stream":true}`), "")
if dataChan != nil {
t.Fatal("terminated stream returned a data channel")
}
errMsg = <-errChan
}
if errMsg == nil || !errMsg.DirectResponse || errMsg.StatusCode != http.StatusForbidden {
t.Fatalf("termination error = %#v", errMsg)
}
if afterCalls != 1 {
t.Fatalf("after-auth interceptor calls = %d, want 1", afterCalls)
}
capturedReq, _ := executor.captured()
if capturedReq.Model != "" {
t.Fatalf("executor received terminated request: %#v", capturedReq)
}
})
}
}
func TestHandlerLifecycleCompletesSuccessfulRequestOnce(t *testing.T) {
model := "handler-interceptor-lifecycle-success"
executor := &interceptorCaptureExecutor{}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{})
var requestID string
var completionCount int
var completion pluginapi.RequestCompletion
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptRequestBeforeAuth: func(_ context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
requestID = req.RequestID
return pluginapi.RequestInterceptResponse{Headers: req.Headers, Body: req.Body}
},
interceptResponse: func(_ context.Context, req pluginapi.ResponseInterceptRequest) pluginapi.ResponseInterceptResponse {
if req.RequestID != requestID {
t.Fatalf("response request ID = %q, want %q", req.RequestID, requestID)
}
return pluginapi.ResponseInterceptResponse{Headers: req.ResponseHeaders, Body: req.Body}
},
completeRequest: func(_ context.Context, got pluginapi.RequestCompletion) {
completionCount++
completion = got
},
})
body, _, errMsg := handler.ExecuteWithAuthManager(context.Background(), "openai", model, []byte(`{"model":"`+model+`"}`), "")
if errMsg != nil || string(body) != "ok" {
t.Fatalf("ExecuteWithAuthManager() body = %q, error = %#v", body, errMsg)
}
if completionCount != 1 || completion.Outcome != pluginapi.RequestCompletionSucceeded || completion.RequestID != requestID {
t.Fatalf("completion count = %d, completion = %#v", completionCount, completion)
}
if completion.StartedAt.IsZero() || completion.CompletedAt.Before(completion.StartedAt) {
t.Fatalf("completion timestamps = %#v", completion)
}
}
func TestHandlerLifecycleCompletesFailedRequest(t *testing.T) {
model := "handler-interceptor-lifecycle-failed"
executor := &interceptorCaptureExecutor{
execute: func(context.Context, *coreauth.Auth, coreexecutor.Request, coreexecutor.Options) (coreexecutor.Response, error) {
return coreexecutor.Response{}, fmt.Errorf("upstream failed")
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{})
var completion pluginapi.RequestCompletion
handler.SetPluginHost(&handlerInterceptorTestHost{
completeRequest: func(_ context.Context, got pluginapi.RequestCompletion) {
completion = got
},
})
_, _, errMsg := handler.ExecuteWithAuthManager(context.Background(), "openai", model, []byte(`{"model":"`+model+`"}`), "")
if errMsg == nil {
t.Fatal("ExecuteWithAuthManager() error = nil")
}
if completion.Outcome != pluginapi.RequestCompletionFailed || completion.Error == "" {
t.Fatalf("completion = %#v", completion)
}
}
func TestHandlerLifecycleCompletesSuccessfulStreamOnce(t *testing.T) {
model := "handler-interceptor-lifecycle-stream"
executor := &interceptorCaptureExecutor{
stream: func(context.Context, *coreauth.Auth, coreexecutor.Request, coreexecutor.Options) (*coreexecutor.StreamResult, error) {
chunks := make(chan coreexecutor.StreamChunk, 1)
chunks <- coreexecutor.StreamChunk{Payload: []byte(`{"chunk":true}`)}
close(chunks)
return &coreexecutor.StreamResult{Chunks: chunks}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{})
completions := make(chan pluginapi.RequestCompletion, 2)
handler.SetPluginHost(&handlerInterceptorTestHost{
completeRequest: func(_ context.Context, completion pluginapi.RequestCompletion) {
completions <- completion
},
})
dataChan, _, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(`{"model":"`+model+`","stream":true}`), "")
for dataChan != nil || errChan != nil {
select {
case _, ok := <-dataChan:
if !ok {
dataChan = nil
}
case errMsg, ok := <-errChan:
if ok && errMsg != nil {
t.Fatalf("stream error = %#v", errMsg)
}
if !ok {
errChan = nil
}
}
}
select {
case completion := <-completions:
if completion.Outcome != pluginapi.RequestCompletionSucceeded || !completion.Stream || completion.RequestID == "" {
t.Fatalf("stream completion = %#v", completion)
}
case <-time.After(time.Second):
t.Fatal("missing stream completion")
}
select {
case duplicate := <-completions:
t.Fatalf("duplicate stream completion = %#v", duplicate)
default:
}
}
func TestHandlerLifecycleCompletesCanceledStream(t *testing.T) {
model := "handler-interceptor-lifecycle-canceled-stream"
chunks := make(chan coreexecutor.StreamChunk, 1)
chunks <- coreexecutor.StreamChunk{Payload: []byte(`{"chunk":true}`)}
executor := &interceptorCaptureExecutor{
stream: func(context.Context, *coreauth.Auth, coreexecutor.Request, coreexecutor.Options) (*coreexecutor.StreamResult, error) {
return &coreexecutor.StreamResult{Chunks: chunks}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{})
completions := make(chan pluginapi.RequestCompletion, 1)
handler.SetPluginHost(&handlerInterceptorTestHost{
completeRequest: func(_ context.Context, completion pluginapi.RequestCompletion) {
completions <- completion
},
})
ctx, cancel := context.WithCancel(context.Background())
dataChan, _, errChan := handler.ExecuteStreamWithAuthManager(ctx, "openai", model, []byte(`{"model":"`+model+`","stream":true}`), "")
cancel()
for dataChan != nil || errChan != nil {
select {
case _, ok := <-dataChan:
if !ok {
dataChan = nil
}
case _, ok := <-errChan:
if !ok {
errChan = nil
}
}
}
completion := <-completions
if completion.Outcome != pluginapi.RequestCompletionCanceled || completion.StatusCode != 0 {
t.Fatalf("completion = %#v", completion)
}
}
func TestHandlerRequestInterceptorRewritesExecutorRequest(t *testing.T) {
model := "handler-interceptor-request-model"
executor := &interceptorCaptureExecutor{}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{})
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptRequestBeforeAuth: func(ctx context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
if req.SourceFormat != "openai" || req.Model != model || req.RequestedModel != model {
t.Fatalf("unexpected request context: %#v", req)
}
if req.Headers.Get("X-Original") != "client" {
t.Fatalf("request headers = %#v, want client header", req.Headers)
}
if req.Metadata == nil {
t.Fatal("metadata = nil, want request metadata")
}
headers := cloneHeader(req.Headers)
headers.Set("X-Original", "plugin")
headers.Set("X-Plugin", "1")
headers.Del("X-Remove")
return pluginapi.RequestInterceptResponse{
Headers: headers,
Body: []byte(fmt.Sprintf(`{"model":%q,"plugin":true}`, model)),
}
},
})
ctx := contextWithHeaders(http.Header{
"X-Original": []string{"client"},
"X-Remove": []string{"yes"},
})
body, _, errMsg := handler.ExecuteWithAuthManager(ctx, "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
if errMsg != nil {
t.Fatalf("ExecuteWithAuthManager() error = %+v", errMsg)
}
if string(body) != "ok" {
t.Fatalf("body = %q, want ok", body)
}
gotReq, gotOpts := executor.captured()
wantPayload := fmt.Sprintf(`{"model":%q,"plugin":true}`, model)
if string(gotReq.Payload) != wantPayload {
t.Fatalf("executor payload = %q, want %q", gotReq.Payload, wantPayload)
}
if string(gotOpts.OriginalRequest) != wantPayload {
t.Fatalf("executor original request = %q, want %q", gotOpts.OriginalRequest, wantPayload)
}
if gotOpts.Headers.Get("X-Original") != "plugin" || gotOpts.Headers.Get("X-Plugin") != "1" {
t.Fatalf("executor headers = %#v, want plugin rewrite", gotOpts.Headers)
}
if gotOpts.Headers.Get("X-Remove") != "" {
t.Fatalf("executor headers kept cleared header: %#v", gotOpts.Headers)
}
if gotOpts.Metadata[coreexecutor.RequestedModelMetadataKey] != model {
t.Fatalf("metadata = %#v, want requested model", gotOpts.Metadata)
}
}
func TestHandlerSkipsDisabledRequestInterceptorsWithoutCopyingPayload(t *testing.T) {
payload := []byte(`{"model":"disabled-interceptor-model"}`)
called := false
handler := NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, nil)
handler.SetPluginHost(&handlerInterceptorDisabledRequestTestHost{
handlerInterceptorTestHost: &handlerInterceptorTestHost{
interceptRequestBeforeAuth: func(context.Context, pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
called = true
return pluginapi.RequestInterceptResponse{Body: []byte(`{"unexpected":true}`)}
},
},
})
req := coreexecutor.Request{Model: "disabled-interceptor-model", Payload: payload}
opts := coreexecutor.Options{OriginalRequest: payload}
gotReq, gotOpts, err := handler.applyRequestInterceptorsBeforeAuth(context.Background(), "openai", req.Model, "test-req", req, opts, "")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if called {
t.Fatal("disabled request interceptor was called")
}
if len(gotReq.Payload) != len(payload) || &gotReq.Payload[0] != &payload[0] {
t.Fatal("request payload was copied")
}
if len(gotOpts.OriginalRequest) != len(payload) || &gotOpts.OriginalRequest[0] != &payload[0] {
t.Fatal("original request was copied")
}
}
func BenchmarkHandlerRequestInterceptors(b *testing.B) {
sizes := []struct {
name string
bytes int
}{
{name: "1KiB", bytes: 1 << 10},
{name: "1MiB", bytes: 1 << 20},
{name: "8MiB", bytes: 8 << 20},
}
hosts := []struct {
name string
host PluginInterceptorHost
}{
{
name: "disabled",
host: &handlerInterceptorDisabledRequestTestHost{
handlerInterceptorTestHost: &handlerInterceptorTestHost{},
},
},
{name: "active", host: &handlerInterceptorTestHost{}},
}
for _, size := range sizes {
payload := make([]byte, size.bytes)
req := coreexecutor.Request{Model: "benchmark-model", Payload: payload}
opts := coreexecutor.Options{OriginalRequest: payload}
for _, host := range hosts {
b.Run(host.name+"/"+size.name, func(b *testing.B) {
handler := NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, nil)
handler.SetPluginHost(host.host)
b.ReportAllocs()
b.ResetTimer()
for range b.N {
gotReq, gotOpts, _ := handler.applyRequestInterceptorsBeforeAuth(context.Background(), "openai", req.Model, "benchmark-req", req, opts, "")
if len(gotReq.Payload) != size.bytes || len(gotOpts.OriginalRequest) != size.bytes {
b.Fatal("request payload length changed")
}
}
})
}
}
}
func TestHandlerRequestInterceptorEmptyBodyKeepsOriginalPayload(t *testing.T) {
model := "handler-interceptor-empty-body-model"
executor := &interceptorCaptureExecutor{}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{})
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptRequestBeforeAuth: func(ctx context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
return pluginapi.RequestInterceptResponse{
Headers: http.Header{"X-Plugin": []string{"empty-body"}},
Body: []byte{},
}
},
})
originalBody := []byte(fmt.Sprintf(`{"model":%q}`, model))
body, _, errMsg := handler.ExecuteWithAuthManager(context.Background(), "openai", model, originalBody, "")
if errMsg != nil {
t.Fatalf("ExecuteWithAuthManager() error = %+v", errMsg)
}
if string(body) != "ok" {
t.Fatalf("body = %q, want ok", body)
}
gotReq, gotOpts := executor.captured()
if string(gotReq.Payload) != string(originalBody) {
t.Fatalf("executor payload = %q, want original payload %q", gotReq.Payload, originalBody)
}
if gotOpts.Headers.Get("X-Plugin") != "empty-body" {
t.Fatalf("executor headers = %#v, want plugin header", gotOpts.Headers)
}
}
func TestHandlerRequestInterceptorAfterAuthRewritesExecutorRequest(t *testing.T) {
model := "handler-interceptor-after-auth-model"
executor := &interceptorCaptureExecutor{}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{})
var calls []string
var responseChecked bool
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptRequestBeforeAuth: func(ctx context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
calls = append(calls, "before")
headers := cloneHeader(req.Headers)
if headers == nil {
headers = http.Header{}
}
headers.Set("X-Stage", "before")
return pluginapi.RequestInterceptResponse{
Headers: headers,
Body: []byte(`{"stage":"before"}`),
}
},
interceptRequestAfterAuth: func(ctx context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
calls = append(calls, "after")
if req.SourceFormat != "openai" || req.ToFormat != "codex" {
t.Fatalf("request formats = %q -> %q, want openai -> codex", req.SourceFormat, req.ToFormat)
}
if req.Model != model || req.RequestedModel != model {
t.Fatalf("request models = %q/%q, want %q/%q", req.Model, req.RequestedModel, model, model)
}
if string(req.Body) != `{"stage":"before"}` {
t.Fatalf("after-auth body = %q, want before-auth rewrite", req.Body)
}
headers := cloneHeader(req.Headers)
if headers == nil {
headers = http.Header{}
}
headers.Set("X-Stage", "after")
return pluginapi.RequestInterceptResponse{
Headers: headers,
Body: []byte(`{"stage":"after"}`),
}
},
interceptResponse: func(ctx context.Context, req pluginapi.ResponseInterceptRequest) pluginapi.ResponseInterceptResponse {
responseChecked = true
if req.RequestHeaders.Get("X-Stage") != "after" {
t.Fatalf("response request headers = %#v, want after-auth header", req.RequestHeaders)
}
if string(req.OriginalRequest) != `{"stage":"after"}` {
t.Fatalf("response original request = %q, want after-auth body", req.OriginalRequest)
}
if string(req.RequestBody) != `{"stage":"after"}` {
t.Fatalf("response request body = %q, want after-auth body", req.RequestBody)
}
return pluginapi.ResponseInterceptResponse{
Headers: cloneHeader(req.ResponseHeaders),
Body: cloneBytes(req.Body),
}
},
})
body, _, errMsg := handler.ExecuteWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
if errMsg != nil {
t.Fatalf("ExecuteWithAuthManager() error = %+v", errMsg)
}
if string(body) != "ok" {
t.Fatalf("body = %q, want ok", body)
}
if fmt.Sprint(calls) != "[before after]" {
t.Fatalf("interceptor calls = %v, want [before after]", calls)
}
gotReq, gotOpts := executor.captured()
if string(gotReq.Payload) != `{"stage":"after"}` {
t.Fatalf("executor payload = %q, want after-auth body", gotReq.Payload)
}
if string(gotOpts.OriginalRequest) != `{"stage":"after"}` {
t.Fatalf("executor original request = %q, want after-auth body", gotOpts.OriginalRequest)
}
if gotOpts.Headers.Get("X-Stage") != "after" {
t.Fatalf("executor headers = %#v, want after-auth header", gotOpts.Headers)
}
if !responseChecked {
t.Fatal("response interceptor was not called")
}
}
func TestHandlerResponseInterceptorRewritesSuccessfulNonStreamResponse(t *testing.T) {
model := "handler-interceptor-response-model"
executor := &interceptorCaptureExecutor{
execute: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (coreexecutor.Response, error) {
return coreexecutor.Response{
Payload: []byte("upstream-body"),
Headers: http.Header{
"X-Upstream": []string{"1"},
"X-Clear": []string{"yes"},
},
}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: true})
var responseCalls int
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptResponse: func(ctx context.Context, req pluginapi.ResponseInterceptRequest) pluginapi.ResponseInterceptResponse {
responseCalls++
if req.StatusCode != http.StatusOK || req.Stream {
t.Fatalf("unexpected response context: %#v", req)
}
if req.ResponseHeaders.Get("X-Upstream") != "1" {
t.Fatalf("response headers = %#v, want upstream header", req.ResponseHeaders)
}
if string(req.Body) != "upstream-body" {
t.Fatalf("response body = %q, want upstream-body", req.Body)
}
headers := cloneHeader(req.ResponseHeaders)
headers.Set("X-Upstream", "2")
headers.Set("X-Plugin", "response")
headers.Del("X-Clear")
return pluginapi.ResponseInterceptResponse{
Headers: headers,
Body: []byte("plugin-body"),
}
},
})
body, headers, errMsg := handler.ExecuteWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
if errMsg != nil {
t.Fatalf("ExecuteWithAuthManager() error = %+v", errMsg)
}
if string(body) != "plugin-body" {
t.Fatalf("body = %q, want plugin-body", body)
}
if headers.Get("X-Upstream") != "2" || headers.Get("X-Plugin") != "response" {
t.Fatalf("headers = %#v, want plugin rewrite", headers)
}
if headers.Get("X-Clear") != "" {
t.Fatalf("headers kept cleared value: %#v", headers)
}
if responseCalls != 1 {
t.Fatalf("response interceptor calls = %d, want 1", responseCalls)
}
}
func TestHandlerExecutorErrorSkipsResponseInterceptor(t *testing.T) {
model := "handler-interceptor-error-model"
executor := &interceptorCaptureExecutor{
execute: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (coreexecutor.Response, error) {
return coreexecutor.Response{}, &coreauth.Error{
Code: "upstream_failed",
Message: "upstream failed",
HTTPStatus: http.StatusBadGateway,
}
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: true})
var responseCalls int
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptResponse: func(ctx context.Context, req pluginapi.ResponseInterceptRequest) pluginapi.ResponseInterceptResponse {
responseCalls++
return pluginapi.ResponseInterceptResponse{Body: []byte("should-not-run")}
},
})
body, headers, errMsg := handler.ExecuteWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
if errMsg == nil {
t.Fatal("ExecuteWithAuthManager() error = nil, want upstream error")
}
if body != nil || headers != nil {
t.Fatalf("body/header = %q/%#v, want nil on error", body, headers)
}
if responseCalls != 0 {
t.Fatalf("response interceptor calls = %d, want 0", responseCalls)
}
}
func TestHandlerStreamExecutorErrorSkipsResponseInterceptors(t *testing.T) {
model := "handler-interceptor-stream-error-model"
executor := &interceptorCaptureExecutor{
stream: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
return nil, &coreauth.Error{
Code: "stream_failed",
Message: "stream failed",
HTTPStatus: http.StatusBadGateway,
}
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: true})
var responseCalls int
var streamCalls int
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptResponse: func(ctx context.Context, req pluginapi.ResponseInterceptRequest) pluginapi.ResponseInterceptResponse {
responseCalls++
return pluginapi.ResponseInterceptResponse{Body: []byte("should-not-run")}
},
interceptStreamChunk: func(ctx context.Context, req pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse {
streamCalls++
return pluginapi.StreamChunkInterceptResponse{Body: []byte("should-not-run")}
},
})
dataChan, headers, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
if dataChan != nil || headers != nil {
t.Fatalf("stream data/header = %#v/%#v, want nil on execute error", dataChan, headers)
}
msg, ok := <-errChan
if !ok || msg == nil {
t.Fatal("stream error channel did not return error message")
}
if msg.StatusCode != http.StatusBadGateway {
t.Fatalf("stream error status = %d, want %d", msg.StatusCode, http.StatusBadGateway)
}
if responseCalls != 0 || streamCalls != 0 {
t.Fatalf("interceptor calls = response:%d stream:%d, want 0", responseCalls, streamCalls)
}
}
func TestHandlerStreamChunkErrorBeforePayloadSkipsResponseInterceptors(t *testing.T) {
model := "handler-interceptor-stream-chunk-error-model"
executor := &interceptorCaptureExecutor{
stream: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
chunks := make(chan coreexecutor.StreamChunk, 1)
chunks <- coreexecutor.StreamChunk{
Err: &coreauth.Error{
Code: "stream_failed",
Message: "stream failed before payload",
HTTPStatus: http.StatusBadGateway,
},
}
close(chunks)
return &coreexecutor.StreamResult{
Headers: http.Header{"X-Upstream": []string{"stream"}},
Chunks: chunks,
}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: true})
var responseCalls int
var streamCalls int
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptResponse: func(ctx context.Context, req pluginapi.ResponseInterceptRequest) pluginapi.ResponseInterceptResponse {
responseCalls++
return pluginapi.ResponseInterceptResponse{Body: []byte("should-not-run")}
},
interceptStreamChunk: func(ctx context.Context, req pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse {
streamCalls++
return pluginapi.StreamChunkInterceptResponse{Headers: cloneHeader(req.ResponseHeaders), Body: []byte("should-not-run")}
},
})
dataChan, headers, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
if dataChan == nil || errChan == nil {
t.Fatalf("stream data/error channels = %#v/%#v, want non-nil channels", dataChan, errChan)
}
for chunk := range dataChan {
t.Fatalf("unexpected stream payload before error: %q", chunk)
}
msg, ok := <-errChan
if !ok || msg == nil {
t.Fatal("stream error channel did not return error message")
}
if msg.StatusCode != http.StatusBadGateway {
t.Fatalf("stream error status = %d, want %d", msg.StatusCode, http.StatusBadGateway)
}
for msg := range errChan {
if msg != nil {
t.Fatalf("unexpected extra stream error: %+v", msg)
}
}
if headers.Get("X-Upstream") != "stream" {
t.Fatalf("headers = %#v, want original upstream headers", headers)
}
if responseCalls != 0 || streamCalls != 0 {
t.Fatalf("interceptor calls = response:%d stream:%d, want 0", responseCalls, streamCalls)
}
}
func TestHandlerStreamInterceptorRewritesAndDropsChunks(t *testing.T) {
model := "handler-interceptor-stream-model"
executor := &interceptorCaptureExecutor{
stream: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
chunks := make(chan coreexecutor.StreamChunk, 3)
chunks <- coreexecutor.StreamChunk{Payload: []byte("first")}
chunks <- coreexecutor.StreamChunk{Payload: []byte("drop")}
chunks <- coreexecutor.StreamChunk{Payload: []byte("second")}
close(chunks)
return &coreexecutor.StreamResult{
Headers: http.Header{"X-Upstream": []string{"stream"}},
Chunks: chunks,
}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: true})
var streamCalls int
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptRequestBeforeAuth: func(ctx context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
headers := cloneHeader(req.Headers)
if headers == nil {
headers = http.Header{}
}
headers.Set("X-Stage", "before")
return pluginapi.RequestInterceptResponse{
Headers: headers,
Body: []byte(`{"stage":"before-stream"}`),
}
},
interceptRequestAfterAuth: func(ctx context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
if string(req.Body) != `{"stage":"before-stream"}` {
t.Fatalf("after-auth stream body = %q, want before-auth rewrite", req.Body)
}
headers := cloneHeader(req.Headers)
if headers == nil {
headers = http.Header{}
}
headers.Set("X-Stage", "after")
return pluginapi.RequestInterceptResponse{
Headers: headers,
Body: []byte(`{"stage":"after-stream"}`),
}
},
interceptStreamChunk: func(ctx context.Context, req pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse {
streamCalls++
if req.RequestHeaders.Get("X-Stage") != "after" {
t.Fatalf("stream request headers = %#v, want after-auth header", req.RequestHeaders)
}
if req.ChunkIndex == pluginapi.StreamChunkHeaderInitIndex {
if string(req.OriginalRequest) != `{"stage":"after-stream"}` {
t.Fatalf("stream original request = %q, want after-auth body", req.OriginalRequest)
}
if string(req.RequestBody) != `{"stage":"after-stream"}` {
t.Fatalf("stream request body = %q, want after-auth body", req.RequestBody)
}
headers := cloneHeader(req.ResponseHeaders)
headers.Set("X-Stream", "plugin")
return pluginapi.StreamChunkInterceptResponse{Headers: headers}
}
if len(req.OriginalRequest) != 0 {
t.Fatalf("payload chunk OriginalRequest = %q, want omitted for schema v3+", req.OriginalRequest)
}
if len(req.RequestBody) != 0 {
t.Fatalf("payload chunk RequestBody = %q, want omitted for schema v3+", req.RequestBody)
}
if req.ResponseHeaders.Get("X-Upstream") != "stream" {
t.Fatalf("stream response headers = %#v, want upstream header", req.ResponseHeaders)
}
if string(req.Body) == "drop" {
return pluginapi.StreamChunkInterceptResponse{DropChunk: true}
}
if string(req.Body) == "second" {
if len(req.HistoryChunks) != 1 || string(req.HistoryChunks[0]) != "first|plugin" {
t.Fatalf("history = %#v, want first transformed chunk", req.HistoryChunks)
}
}
headers := cloneHeader(req.ResponseHeaders)
headers.Set("X-Stream", "plugin")
return pluginapi.StreamChunkInterceptResponse{
Headers: headers,
Body: append(req.Body, []byte("|plugin")...),
}
},
})
dataChan, upstreamHeaders, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
var got []byte
for chunk := range dataChan {
got = append(got, chunk...)
}
for msg := range errChan {
if msg != nil {
t.Fatalf("unexpected stream error: %+v", msg)
}
}
if string(got) != "first|pluginsecond|plugin" {
t.Fatalf("stream payload = %q, want transformed chunks without dropped chunk", got)
}
if upstreamHeaders.Get("X-Stream") != "plugin" {
t.Fatalf("upstream headers = %#v, want stream plugin header", upstreamHeaders)
}
if streamCalls != 4 {
t.Fatalf("stream interceptor calls = %d, want 4", streamCalls)
}
}
func TestHandlerStreamInterceptorLegacySchemaClonesRequestBodiesOnPayloadChunks(t *testing.T) {
model := "handler-interceptor-stream-legacy-clone-model"
executor := &interceptorCaptureExecutor{
stream: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
chunks := make(chan coreexecutor.StreamChunk, 2)
chunks <- coreexecutor.StreamChunk{Payload: []byte("first")}
chunks <- coreexecutor.StreamChunk{Payload: []byte("second")}
close(chunks)
return &coreexecutor.StreamResult{
Headers: http.Header{"X-Upstream": []string{"stream"}},
Chunks: chunks,
}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: true})
var payloadBodies [][]byte
handler.SetPluginHost(&handlerInterceptorTestHost{
includeStreamChunkRequestBodies: true,
interceptRequestAfterAuth: func(ctx context.Context, req pluginapi.RequestInterceptRequest) pluginapi.RequestInterceptResponse {
return pluginapi.RequestInterceptResponse{Body: []byte(`{"stage":"legacy-stream"}`)}
},
interceptStreamChunk: func(ctx context.Context, req pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse {
if req.ChunkIndex == pluginapi.StreamChunkHeaderInitIndex {
if string(req.OriginalRequest) != `{"stage":"legacy-stream"}` || string(req.RequestBody) != `{"stage":"legacy-stream"}` {
t.Fatalf("header-init bodies = original:%q body:%q", req.OriginalRequest, req.RequestBody)
}
// Mutate delivered slices; later chunks must not observe this mutation.
req.OriginalRequest[0] = 'X'
req.RequestBody[0] = 'Y'
return pluginapi.StreamChunkInterceptResponse{}
}
if string(req.OriginalRequest) != `{"stage":"legacy-stream"}` {
t.Fatalf("payload OriginalRequest = %q, want isolated clone of after-auth body", req.OriginalRequest)
}
if string(req.RequestBody) != `{"stage":"legacy-stream"}` {
t.Fatalf("payload RequestBody = %q, want isolated clone of after-auth body", req.RequestBody)
}
payloadBodies = append(payloadBodies, req.OriginalRequest)
req.OriginalRequest[0] = 'Z'
return pluginapi.StreamChunkInterceptResponse{Body: req.Body}
},
})
dataChan, _, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
for range dataChan {
}
for msg := range errChan {
if msg != nil {
t.Fatalf("unexpected stream error: %+v", msg)
}
}
if len(payloadBodies) != 2 {
t.Fatalf("payload body deliveries = %d, want 2", len(payloadBodies))
}
if &payloadBodies[0][0] == &payloadBodies[1][0] {
t.Fatal("payload OriginalRequest slices alias across chunks; want fresh clones")
}
}
func TestHandlerStreamInterceptorInitializesHeadersBeforeReturn(t *testing.T) {
model := "handler-interceptor-stream-header-before-return-model"
initStarted := make(chan struct{})
allowInit := make(chan struct{})
executor := &interceptorCaptureExecutor{
stream: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
chunks := make(chan coreexecutor.StreamChunk, 1)
chunks <- coreexecutor.StreamChunk{Payload: []byte("payload")}
close(chunks)
return &coreexecutor.StreamResult{
Headers: http.Header{"X-Upstream": []string{"stream"}},
Chunks: chunks,
}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: true})
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptStreamChunk: func(ctx context.Context, req pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse {
headers := cloneHeader(req.ResponseHeaders)
if req.ChunkIndex == pluginapi.StreamChunkHeaderInitIndex {
close(initStarted)
<-allowInit
headers.Set("X-Init", "plugin")
}
return pluginapi.StreamChunkInterceptResponse{
Headers: headers,
Body: cloneBytes(req.Body),
}
},
})
type streamResult struct {
dataChan <-chan []byte
upstreamHeaders http.Header
errChan <-chan *interfaces.ErrorMessage
}
resultChan := make(chan streamResult, 1)
go func() {
dataChan, upstreamHeaders, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
resultChan <- streamResult{dataChan: dataChan, upstreamHeaders: upstreamHeaders, errChan: errChan}
}()
select {
case result := <-resultChan:
t.Fatalf("ExecuteStreamWithAuthManager returned before stream header init: %#v", result.upstreamHeaders)
case <-initStarted:
}
select {
case result := <-resultChan:
t.Fatalf("ExecuteStreamWithAuthManager returned while stream header init was blocked: %#v", result.upstreamHeaders)
default:
}
close(allowInit)
result := <-resultChan
dataChan := result.dataChan
upstreamHeaders := result.upstreamHeaders
errChan := result.errChan
if upstreamHeaders.Get("X-Init") != "plugin" {
t.Fatalf("upstream headers before first payload = %#v, want initialized plugin header", upstreamHeaders)
}
for range dataChan {
}
for msg := range errChan {
if msg != nil {
t.Fatalf("unexpected stream error: %+v", msg)
}
}
}
func TestHandlerStreamSkipsInterceptorsWhenHostReportsNoStreamInterceptors(t *testing.T) {
model := "handler-interceptor-no-stream-capability-model"
executor := &interceptorCaptureExecutor{
stream: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
chunks := make(chan coreexecutor.StreamChunk, 1)
chunks <- coreexecutor.StreamChunk{Payload: []byte("payload")}
close(chunks)
return &coreexecutor.StreamResult{
Headers: http.Header{"X-Upstream": []string{"stream"}},
Chunks: chunks,
}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: false})
var streamCalls int
handler.SetPluginHost(&handlerInterceptorNoStreamTestHost{
handlerInterceptorTestHost: &handlerInterceptorTestHost{
interceptStreamChunk: func(ctx context.Context, req pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse {
streamCalls++
return pluginapi.StreamChunkInterceptResponse{Headers: cloneHeader(req.ResponseHeaders), Body: cloneBytes(req.Body)}
},
},
})
dataChan, upstreamHeaders, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
var got []byte
for chunk := range dataChan {
got = append(got, chunk...)
}
for msg := range errChan {
if msg != nil {
t.Fatalf("unexpected stream error: %+v", msg)
}
}
if string(got) != "payload" {
t.Fatalf("stream payload = %q, want payload", got)
}
if upstreamHeaders != nil {
t.Fatalf("upstream headers = %#v, want nil without passthrough or stream interceptors", upstreamHeaders)
}
if streamCalls != 0 {
t.Fatalf("stream interceptor calls = %d, want 0", streamCalls)
}
}
func TestAppendStreamInterceptorHistoryBoundsRetainedChunks(t *testing.T) {
var history [][]byte
for i := 0; i < maxStreamInterceptorHistoryChunks+10; i++ {
history = appendStreamInterceptorHistory(history, []byte{byte(i)})
}
if len(history) != maxStreamInterceptorHistoryChunks {
t.Fatalf("history chunks = %d, want %d", len(history), maxStreamInterceptorHistoryChunks)
}
if got := history[0][0]; got != 10 {
t.Fatalf("first retained history chunk = %d, want 10", got)
}
history = nil
largeChunk := make([]byte, maxStreamInterceptorHistoryBytes/2+1)
for i := 0; i < 3; i++ {
history = appendStreamInterceptorHistory(history, largeChunk)
}
if gotBytes := byteSlicesSize(history); gotBytes > maxStreamInterceptorHistoryBytes {
t.Fatalf("history bytes = %d, want <= %d", gotBytes, maxStreamInterceptorHistoryBytes)
}
}
func TestHandlerStreamInterceptorKeepsReturnedHeadersStableAfterFirstPayload(t *testing.T) {
model := "handler-interceptor-stream-stable-headers-model"
releaseSecond := make(chan struct{})
executor := &interceptorCaptureExecutor{
stream: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
chunks := make(chan coreexecutor.StreamChunk)
go func() {
defer close(chunks)
chunks <- coreexecutor.StreamChunk{Payload: []byte("first")}
<-releaseSecond
chunks <- coreexecutor.StreamChunk{Payload: []byte("second")}
}()
return &coreexecutor.StreamResult{
Headers: http.Header{"X-Upstream": []string{"stream"}},
Chunks: chunks,
}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: true})
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptStreamChunk: func(ctx context.Context, req pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse {
headers := cloneHeader(req.ResponseHeaders)
switch req.ChunkIndex {
case pluginapi.StreamChunkHeaderInitIndex:
headers.Set("X-Stage", "init")
case 0:
headers.Set("X-Chunk", "first")
case 1:
headers.Set("X-Chunk", "second")
}
return pluginapi.StreamChunkInterceptResponse{
Headers: headers,
Body: cloneBytes(req.Body),
}
},
})
dataChan, upstreamHeaders, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
firstChunk, ok := <-dataChan
if !ok {
t.Fatal("data channel closed before first chunk")
}
if string(firstChunk) != "first" {
t.Fatalf("first chunk = %q, want first", firstChunk)
}
if upstreamHeaders.Get("X-Chunk") != "first" || upstreamHeaders.Get("X-Stage") != "init" {
t.Fatalf("upstream headers after first chunk = %#v, want first transformed chunk headers", upstreamHeaders)
}
close(releaseSecond)
got := append([]byte(nil), firstChunk...)
for chunk := range dataChan {
got = append(got, chunk...)
}
for msg := range errChan {
if msg != nil {
t.Fatalf("unexpected stream error: %+v", msg)
}
}
if string(got) != "firstsecond" {
t.Fatalf("stream payload = %q, want firstsecond", got)
}
if upstreamHeaders.Get("X-Chunk") != "first" {
t.Fatalf("upstream headers changed after return: %#v", upstreamHeaders)
}
}
func TestHandlerStreamInterceptorReturnedHeadersImmutableAfterReturn(t *testing.T) {
model := "handler-interceptor-stream-immutable-headers-model"
releaseSecond := make(chan struct{})
bodyStarted := make(chan struct{})
releaseBody := make(chan struct{})
executor := &interceptorCaptureExecutor{
stream: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
chunks := make(chan coreexecutor.StreamChunk)
go func() {
defer close(chunks)
chunks <- coreexecutor.StreamChunk{Payload: []byte("first")}
<-releaseSecond
chunks <- coreexecutor.StreamChunk{Payload: []byte("second")}
}()
return &coreexecutor.StreamResult{
Headers: http.Header{"X-Upstream": []string{"stream"}},
Chunks: chunks,
}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: true})
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptStreamChunk: func(ctx context.Context, req pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse {
headers := cloneHeader(req.ResponseHeaders)
switch req.ChunkIndex {
case pluginapi.StreamChunkHeaderInitIndex:
headers.Set("X-Init", "plugin")
case 1:
close(bodyStarted)
<-releaseBody
headers.Set("X-Body", "plugin")
}
return pluginapi.StreamChunkInterceptResponse{Headers: headers, Body: cloneBytes(req.Body)}
},
})
dataChan, upstreamHeaders, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
dataDone := make(chan struct{})
go func() {
defer close(dataDone)
for range dataChan {
}
}()
stopReading := make(chan struct{})
readerDone := make(chan struct{})
go func() {
defer close(readerDone)
for {
select {
case <-stopReading:
return
default:
_ = upstreamHeaders.Get("X-Init")
}
}
}()
close(releaseSecond)
<-bodyStarted
close(releaseBody)
<-dataDone
for msg := range errChan {
if msg != nil {
t.Fatalf("unexpected stream error: %+v", msg)
}
}
close(stopReading)
<-readerDone
if upstreamHeaders.Get("X-Init") != "plugin" || upstreamHeaders.Get("X-Body") != "" {
t.Fatalf("returned headers mutated after return: %#v", upstreamHeaders)
}
}
func TestHandlerStreamInterceptorInitializesHeadersWithoutPayload(t *testing.T) {
model := "handler-interceptor-stream-header-only-model"
executor := &interceptorCaptureExecutor{
stream: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (*coreexecutor.StreamResult, error) {
chunks := make(chan coreexecutor.StreamChunk, 1)
chunks <- coreexecutor.StreamChunk{Payload: []byte("payload")}
close(chunks)
return &coreexecutor.StreamResult{
Headers: http.Header{"X-Upstream": []string{"stream"}},
Chunks: chunks,
}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: true})
var initCalls int
var payloadCalls int
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptStreamChunk: func(ctx context.Context, req pluginapi.StreamChunkInterceptRequest) pluginapi.StreamChunkInterceptResponse {
if req.ChunkIndex != pluginapi.StreamChunkHeaderInitIndex {
payloadCalls++
if string(req.Body) != "payload" || req.ResponseHeaders.Get("X-Init") != "plugin" {
t.Fatalf("payload stream request = %#v, want initialized headers and payload", req)
}
return pluginapi.StreamChunkInterceptResponse{Headers: cloneHeader(req.ResponseHeaders), Body: cloneBytes(req.Body)}
}
initCalls++
headers := cloneHeader(req.ResponseHeaders)
headers.Set("X-Init", "plugin")
return pluginapi.StreamChunkInterceptResponse{Headers: headers}
},
})
dataChan, upstreamHeaders, errChan := handler.ExecuteStreamWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
for chunk := range dataChan {
if string(chunk) != "payload" {
t.Fatalf("stream chunk = %q, want payload", chunk)
}
}
for msg := range errChan {
if msg != nil {
t.Fatalf("unexpected stream error: %+v", msg)
}
}
if initCalls != 1 {
t.Fatalf("initial stream calls = %d, want 1", initCalls)
}
if payloadCalls != 1 {
t.Fatalf("payload stream calls = %d, want 1", payloadCalls)
}
if upstreamHeaders.Get("X-Init") != "plugin" {
t.Fatalf("upstream headers = %#v, want initial plugin header", upstreamHeaders)
}
}
func TestHandlerResponseInterceptorSeesRawHeadersWhenPassthroughDisabled(t *testing.T) {
model := "handler-interceptor-raw-headers-model"
executor := &interceptorCaptureExecutor{
execute: func(ctx context.Context, auth *coreauth.Auth, req coreexecutor.Request, opts coreexecutor.Options) (coreexecutor.Response, error) {
return coreexecutor.Response{
Payload: []byte("upstream-body"),
Headers: http.Header{
"X-Upstream": []string{"raw"},
},
}, nil
},
}
handler := newInterceptorHandler(t, model, executor, &sdkconfig.SDKConfig{PassthroughHeaders: false})
handler.SetPluginHost(&handlerInterceptorTestHost{
interceptResponse: func(ctx context.Context, req pluginapi.ResponseInterceptRequest) pluginapi.ResponseInterceptResponse {
if req.ResponseHeaders.Get("X-Upstream") != "raw" {
t.Fatalf("response headers = %#v, want raw upstream header", req.ResponseHeaders)
}
headers := cloneHeader(req.ResponseHeaders)
headers.Set("X-Plugin", "response")
return pluginapi.ResponseInterceptResponse{Headers: headers}
},
})
_, headers, errMsg := handler.ExecuteWithAuthManager(context.Background(), "openai", model, []byte(fmt.Sprintf(`{"model":%q}`, model)), "")
if errMsg != nil {
t.Fatalf("ExecuteWithAuthManager() error = %+v", errMsg)
}
if headers.Get("X-Plugin") != "response" {
t.Fatalf("headers = %#v, want plugin header", headers)
}
if headers.Get("X-Upstream") != "" {
t.Fatalf("headers leaked raw upstream header with passthrough disabled: %#v", headers)
}
}