package vllm import ( "bytes" "context" "errors" "net/http" "net/http/httptest" "strings" "sync" "testing" "time" "go.uber.org/zap" "iop/apps/node/internal/runtime" "iop/packages/go/config" ) type fakeTunnelSink struct { mu sync.Mutex frames []runtime.ProviderTunnelFrame } func (s *fakeTunnelSink) EmitTunnelFrame(_ context.Context, frame runtime.ProviderTunnelFrame) error { s.mu.Lock() defer s.mu.Unlock() s.frames = append(s.frames, frame) return nil } func (s *fakeTunnelSink) all() []runtime.ProviderTunnelFrame { s.mu.Lock() defer s.mu.Unlock() return append([]runtime.ProviderTunnelFrame(nil), s.frames...) } // failingTunnelSink emulates the Edge socket write failing after failAfter // successful frames, as when the Edge relay aborts an abandoned tunnel. type failingTunnelSink struct { mu sync.Mutex frames []runtime.ProviderTunnelFrame failAfter int } func (s *failingTunnelSink) EmitTunnelFrame(_ context.Context, frame runtime.ProviderTunnelFrame) error { s.mu.Lock() defer s.mu.Unlock() if len(s.frames) >= s.failAfter { return errors.New("edge socket write failed") } s.frames = append(s.frames, frame) return nil } func (s *failingTunnelSink) all() []runtime.ProviderTunnelFrame { s.mu.Lock() defer s.mu.Unlock() return append([]runtime.ProviderTunnelFrame(nil), s.frames...) } // assertOrderedTunnelFrames verifies S09 frame semantics: contiguous sequences // starting at 0 (no drop/reorder) and exactly one terminal END/ERROR frame, // which must be the last frame emitted. func assertOrderedTunnelFrames(t *testing.T, frames []runtime.ProviderTunnelFrame) { t.Helper() if len(frames) == 0 { t.Fatal("no tunnel frames emitted") } terminals := 0 for i, f := range frames { if f.Sequence != int64(i) { t.Errorf("frame %d: sequence %d, want %d (ordered, lossless)", i, f.Sequence, i) } switch f.Kind { case runtime.ProviderTunnelFrameKindEnd, runtime.ProviderTunnelFrameKindError: terminals++ if i != len(frames)-1 { t.Errorf("terminal %s frame at index %d must be the last of %d frames", f.Kind, i, len(frames)) } } } if terminals != 1 { t.Errorf("expected exactly one terminal END/ERROR frame, got %d", terminals) } } func TestVllmTunnelProvider(t *testing.T) { expectedBody := "data: {\"choices\":[{\"delta\":{\"content\":\"hello\"}}]}\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) } 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.VllmConf{ 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","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 TestVllmTunnelProvider_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 { if _, err := w.Write([]byte("data: chunk\n\n")); 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.VllmConf{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 cancel") } case <-time.After(2 * time.Second): t.Fatal("timeout waiting for TunnelProvider") } select { case <-handlerObservedCancel: case <-time.After(2 * time.Second): t.Fatal("handler did not observe 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 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, got %s", i+1, f.Kind) } } lastFrame := frames[len(frames)-1] if lastFrame.Kind != runtime.ProviderTunnelFrameKindError || lastFrame.Error == "" || !strings.Contains(lastFrame.Error, "context canceled") { t.Errorf("expected terminal ERROR with context canceled, got %+v", lastFrame) } } func TestVllmTunnelProvider_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.VllmConf{ 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) } } // TestVllmTunnelProvider_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 TestVllmTunnelProvider_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.VllmConf{ 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) } } } func TestVllmTunnelProvider_GemmaContentToolCallRawBytes(t *testing.T) { expectedBody := "data: {\"choices\":[{\"delta\":{\"content\":\"Gemma response \"}}]}\n\n" + "data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"index\":0,\"id\":\"call_gemma_1\",\"type\":\"function\",\"function\":{\"name\":\"run_commands\",\"arguments\":\"{\\\"commands\\\": \"}}]}}]}\n\n" + "data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"index\":0,\"function\":{\"arguments\":\"[\\\"git status\\\"]}\"}}]}}]}\n\n" + "data: [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) } w.Header().Set("Content-Type", "text/event-stream") w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte(expectedBody)) })) defer server.Close() adapter := New(config.VllmConf{ Endpoint: server.URL, }, zap.NewNop()) sink := &fakeTunnelSink{} req := runtime.ProviderTunnelRequest{ RunID: "run-gemma-1", TunnelID: "tunnel-gemma-1", Method: "POST", Path: "/v1/chat/completions", Body: []byte(`{"model":"gemma-2b","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) } 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:\ngot: %q\nwant: %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) }