iop/apps/node/internal/adapters/vllm/vllm_tunnel_test.go

455 lines
13 KiB
Go

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)
}