455 lines
13 KiB
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)
|
|
}
|