package openai_compat import ( "bytes" "context" "net/http" "net/http/httptest" "strings" "testing" "time" "go.uber.org/zap" "iop/apps/node/internal/runtime" "iop/packages/go/config" ) func TestOpenAICompatTunnelProvider(t *testing.T) { expectedBody := "data: {\"choices\":[{\"delta\":{\"reasoning_content\":\"think Qwen\"}}]}\n\ndata: [DONE]\n\n" server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { t.Errorf("expected POST method, got %s", r.Method) } if r.URL.Path != "/v1/chat/completions" { t.Errorf("expected path /v1/chat/completions, got %s", r.URL.Path) } if r.Header.Get("Authorization") != "Bearer test-key" { t.Errorf("expected Authorization header, got %s", r.Header.Get("Authorization")) } w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("X-Custom-Header", "custom-value") w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte(expectedBody)) })) defer server.Close() adapter := New(config.OpenAICompatConf{ Endpoint: server.URL, }, zap.NewNop()) sink := &fakeTunnelSink{} req := runtime.ProviderTunnelRequest{ RunID: "run-1", TunnelID: "tunnel-1", Method: "POST", Path: "/v1/chat/completions", Headers: map[string]string{"Authorization": "Bearer test-key"}, Body: []byte(`{"model":"qwen3.6:35b","prompt":"hi"}`), } err := adapter.TunnelProvider(context.Background(), req, sink) if err != nil { t.Fatalf("TunnelProvider failed: %v", err) } frames := sink.all() if len(frames) < 3 { t.Fatalf("expected at least 3 frames, got %d", len(frames)) } startFrame := frames[0] if startFrame.Kind != runtime.ProviderTunnelFrameKindResponseStart { t.Errorf("expected RESPONSE_START, got %s", startFrame.Kind) } if startFrame.StatusCode != http.StatusOK { t.Errorf("expected 200 OK, got %d", startFrame.StatusCode) } if startFrame.Headers["X-Custom-Header"] != "custom-value" { t.Errorf("expected Custom Header, got %v", startFrame.Headers) } var bodyBuffer bytes.Buffer var lastSeq int64 = 0 for _, f := range frames[1 : len(frames)-1] { if f.Kind != runtime.ProviderTunnelFrameKindBody { t.Errorf("expected BODY kind, got %s", f.Kind) } if f.Sequence != lastSeq+1 { t.Errorf("expected seq %d, got %d", lastSeq+1, f.Sequence) } bodyBuffer.Write(f.Body) lastSeq = f.Sequence } if bodyBuffer.String() != expectedBody { t.Errorf("body mismatch: got %q, want %q", bodyBuffer.String(), expectedBody) } endFrame := frames[len(frames)-1] if endFrame.Kind != runtime.ProviderTunnelFrameKindEnd { t.Errorf("expected END, got %s", endFrame.Kind) } if !endFrame.End { t.Errorf("expected End true, got %v", endFrame.End) } if endFrame.Sequence != lastSeq+1 { t.Errorf("expected seq %d, got %d", lastSeq+1, endFrame.Sequence) } assertOrderedTunnelFrames(t, frames) } func TestOpenAICompatTunnelProvider_RelaysProviderHTTPError(t *testing.T) { expectedBody := `{"error":{"message":"unsupported field","type":"invalid_request_error","param":"custom_provider_options"}}` server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { t.Errorf("expected POST method, got %s", r.Method) } if r.URL.Path != "/v1/chat/completions" { t.Errorf("expected path /v1/chat/completions, got %s", r.URL.Path) } w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusUnprocessableEntity) _, _ = w.Write([]byte(expectedBody)) })) defer server.Close() adapter := New(config.OpenAICompatConf{ Endpoint: server.URL, }, zap.NewNop()) sink := &fakeTunnelSink{} req := runtime.ProviderTunnelRequest{ RunID: "run-1", TunnelID: "tunnel-1", Method: "POST", Path: "/v1/chat/completions", Body: []byte(`{"model":"qwen3.6:35b","custom_provider_options":{"reject":true}}`), } err := adapter.TunnelProvider(context.Background(), req, sink) if err != nil { t.Fatalf("TunnelProvider must relay provider HTTP errors as response frames, got error: %v", err) } frames := sink.all() if len(frames) < 3 { t.Fatalf("expected at least 3 frames, got %d", len(frames)) } if frames[0].Kind != runtime.ProviderTunnelFrameKindResponseStart { t.Fatalf("expected RESPONSE_START, got %s", frames[0].Kind) } if frames[0].StatusCode != http.StatusUnprocessableEntity { t.Fatalf("expected provider status 422, got %d", frames[0].StatusCode) } if got := frames[0].Headers["Content-Type"]; got != "application/json" { t.Fatalf("expected provider Content-Type header, got %q", got) } var bodyBuffer bytes.Buffer for _, f := range frames[1 : len(frames)-1] { if f.Kind != runtime.ProviderTunnelFrameKindBody { t.Errorf("expected BODY kind, got %s", f.Kind) } bodyBuffer.Write(f.Body) } if bodyBuffer.String() != expectedBody { t.Fatalf("provider error body mismatch:\n got: %q\nwant: %q", bodyBuffer.String(), expectedBody) } if frames[len(frames)-1].Kind != runtime.ProviderTunnelFrameKindEnd { t.Fatalf("expected END, got %s", frames[len(frames)-1].Kind) } assertOrderedTunnelFrames(t, frames) } func TestOpenAICompatTunnelProvider_ResponsesPath(t *testing.T) { requestBody := `{"model":"served-model","input":"hi","max_output_tokens":123,"store":false}` expectedBody := "data: {\"type\":\"response.output_text.delta\",\"delta\":\"hi\"}\n\ndata: [DONE]\n\n" server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodPost { t.Errorf("expected POST method, got %s", r.Method) } if r.URL.Path != "/v1/responses" { t.Errorf("expected path /v1/responses, got %s", r.URL.Path) } if r.Header.Get("Authorization") != "Bearer test-key" { t.Errorf("expected Authorization header, got %s", r.Header.Get("Authorization")) } var got bytes.Buffer if _, err := got.ReadFrom(r.Body); err != nil { t.Fatalf("read request body: %v", err) } if got.String() != requestBody { t.Errorf("request body mismatch: got %q want %q", got.String(), requestBody) } w.Header().Set("Content-Type", "text/event-stream") w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte(expectedBody)) })) defer server.Close() adapter := New(config.OpenAICompatConf{ Endpoint: server.URL, }, zap.NewNop()) sink := &fakeTunnelSink{} req := runtime.ProviderTunnelRequest{ RunID: "run-1", TunnelID: "tunnel-1", Method: "POST", Path: "/v1/responses", Headers: map[string]string{"Authorization": "Bearer test-key"}, Body: []byte(requestBody), Stream: true, } err := adapter.TunnelProvider(context.Background(), req, sink) if err != nil { t.Fatalf("TunnelProvider failed: %v", err) } frames := sink.all() if len(frames) < 3 { t.Fatalf("expected at least 3 frames, got %d", len(frames)) } if frames[0].Kind != runtime.ProviderTunnelFrameKindResponseStart { t.Errorf("expected RESPONSE_START, got %s", frames[0].Kind) } if frames[0].StatusCode != http.StatusOK { t.Errorf("expected 200 OK, got %d", frames[0].StatusCode) } var bodyBuffer bytes.Buffer for _, f := range frames[1 : len(frames)-1] { if f.Kind != runtime.ProviderTunnelFrameKindBody { t.Errorf("expected BODY kind, got %s", f.Kind) } bodyBuffer.Write(f.Body) } if bodyBuffer.String() != expectedBody { t.Errorf("body mismatch: got %q, want %q", bodyBuffer.String(), expectedBody) } if frames[len(frames)-1].Kind != runtime.ProviderTunnelFrameKindEnd { t.Errorf("expected END, got %s", frames[len(frames)-1].Kind) } assertOrderedTunnelFrames(t, frames) } func TestOpenAICompatTunnelProvider_Cancel(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) handlerObservedCancel := make(chan struct{}) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") w.WriteHeader(http.StatusOK) flusher, ok := w.(http.Flusher) for { _, err := w.Write([]byte("data: chunk\n\n")) if err != nil { close(handlerObservedCancel) return } if ok { flusher.Flush() } select { case <-r.Context().Done(): close(handlerObservedCancel) return case <-time.After(10 * time.Millisecond): } } })) defer server.Close() adapter := New(config.OpenAICompatConf{ Endpoint: server.URL, }, zap.NewNop()) sink := &fakeTunnelSink{} req := runtime.ProviderTunnelRequest{ RunID: "run-1", TunnelID: "tunnel-1", Method: "POST", Path: "/v1/chat/completions", Body: []byte(`{"prompt":"hi"}`), } errCh := make(chan error, 1) go func() { errCh <- adapter.TunnelProvider(ctx, req, sink) }() time.Sleep(100 * time.Millisecond) cancel() select { case err := <-errCh: if err == nil { t.Error("expected error from TunnelProvider on cancellation, got nil") } case <-time.After(2 * time.Second): t.Fatal("timeout waiting for TunnelProvider to return") } select { case <-handlerObservedCancel: case <-time.After(2 * time.Second): t.Fatal("handler did not observe context cancel") } frames := sink.all() if len(frames) < 2 { t.Fatalf("expected at least 2 frames, got %d: %+v", len(frames), frames) } assertOrderedTunnelFrames(t, frames) if frames[0].Kind != runtime.ProviderTunnelFrameKindResponseStart { t.Errorf("expected first frame RESPONSE_START, got %s", frames[0].Kind) } for i, f := range frames[1 : len(frames)-1] { if f.Kind != runtime.ProviderTunnelFrameKindBody { t.Errorf("mid frame %d: expected BODY before the terminal ERROR, got %s", i+1, f.Kind) } } lastFrame := frames[len(frames)-1] if lastFrame.Kind != runtime.ProviderTunnelFrameKindError { t.Errorf("expected last frame to be ERROR, got %s", lastFrame.Kind) } if lastFrame.Error == "" { t.Error("expected non-empty error message in frame") } if !strings.Contains(lastFrame.Error, "context canceled") { t.Errorf("expected cancellation error in terminal frame, got %q", lastFrame.Error) } } func TestOpenAICompatTunnelProvider_BodyReadError(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { hj, ok := w.(http.Hijacker) if !ok { t.Fatal("webserver doesn't support hijacking") } conn, _, err := hj.Hijack() if err != nil { t.Fatalf("hijack failed: %v", err) } defer conn.Close() _, _ = conn.Write([]byte("HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: 100\r\n\r\n")) _, _ = conn.Write([]byte("data: chunk1\n\n")) })) defer server.Close() adapter := New(config.OpenAICompatConf{ Endpoint: server.URL, }, zap.NewNop()) sink := &fakeTunnelSink{} req := runtime.ProviderTunnelRequest{ RunID: "run-1", TunnelID: "tunnel-1", Method: "POST", Path: "/v1/chat/completions", Body: []byte(`{"prompt":"hi"}`), } err := adapter.TunnelProvider(context.Background(), req, sink) if err == nil { t.Fatal("expected error from body read failure, got nil") } frames := sink.all() if len(frames) < 2 { t.Fatalf("expected at least 2 frames, got %d: %+v", len(frames), frames) } assertOrderedTunnelFrames(t, frames) lastFrame := frames[len(frames)-1] if lastFrame.Kind != runtime.ProviderTunnelFrameKindError { t.Errorf("expected last frame to be ERROR, got %s", lastFrame.Kind) } if !strings.Contains(lastFrame.Error, "read response body") && !strings.Contains(lastFrame.Error, "EOF") { t.Errorf("expected body read error msg, got %q", lastFrame.Error) } } // TestOpenAICompatTunnelProvider_SinkWriteFailureStopsUpstreamRelay covers the // Edge-write-failure leg of S09: when EmitTunnelFrame fails mid-stream the // adapter stops relaying immediately, returns the emit error without appending // further frames, and the upstream provider request is torn down. func TestOpenAICompatTunnelProvider_SinkWriteFailureStopsUpstreamRelay(t *testing.T) { handlerObservedTeardown := make(chan struct{}) server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") w.WriteHeader(http.StatusOK) flusher, ok := w.(http.Flusher) for { _, err := w.Write([]byte("data: chunk\n\n")) if err != nil { close(handlerObservedTeardown) return } if ok { flusher.Flush() } select { case <-r.Context().Done(): close(handlerObservedTeardown) return case <-time.After(10 * time.Millisecond): } } })) defer server.Close() adapter := New(config.OpenAICompatConf{ Endpoint: server.URL, }, zap.NewNop()) // RESPONSE_START and the first BODY frame succeed; the next emit fails as // if the Edge socket write broke after the caller went away. sink := &failingTunnelSink{failAfter: 2} req := runtime.ProviderTunnelRequest{ RunID: "run-1", TunnelID: "tunnel-1", Method: "POST", Path: "/v1/chat/completions", Body: []byte(`{"prompt":"hi"}`), } errCh := make(chan error, 1) go func() { errCh <- adapter.TunnelProvider(context.Background(), req, sink) }() var tunnelErr error select { case tunnelErr = <-errCh: case <-time.After(2 * time.Second): t.Fatal("timeout waiting for TunnelProvider to return after sink write failure") } if tunnelErr == nil || !strings.Contains(tunnelErr.Error(), "edge socket write failed") { t.Fatalf("expected the sink emit error to propagate, got %v", tunnelErr) } select { case <-handlerObservedTeardown: case <-time.After(2 * time.Second): t.Fatal("provider handler did not observe upstream request teardown") } frames := sink.all() if len(frames) != 2 { t.Fatalf("expected relay to stop at 2 frames after emit failure, got %d: %+v", len(frames), frames) } if frames[0].Kind != runtime.ProviderTunnelFrameKindResponseStart { t.Errorf("expected first frame RESPONSE_START, got %s", frames[0].Kind) } if frames[1].Kind != runtime.ProviderTunnelFrameKindBody { t.Errorf("expected second frame BODY, got %s", frames[1].Kind) } for i, f := range frames { if f.Sequence != int64(i) { t.Errorf("frame %d: sequence %d, want %d", i, f.Sequence, i) } } }