190 lines
5.5 KiB
Go
190 lines
5.5 KiB
Go
package vllm
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
|
|
"go.uber.org/zap"
|
|
|
|
"iop/apps/node/internal/runtime"
|
|
"iop/packages/go/config"
|
|
)
|
|
|
|
type fakeSink struct {
|
|
mu sync.Mutex
|
|
events []runtime.RuntimeEvent
|
|
}
|
|
|
|
func (s *fakeSink) Emit(_ context.Context, event runtime.RuntimeEvent) error {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
s.events = append(s.events, event)
|
|
return nil
|
|
}
|
|
|
|
func (s *fakeSink) all() []runtime.RuntimeEvent {
|
|
s.mu.Lock()
|
|
defer s.mu.Unlock()
|
|
return append([]runtime.RuntimeEvent(nil), s.events...)
|
|
}
|
|
|
|
func TestVllmCapabilitiesQueryModels(t *testing.T) {
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
if r.URL.Path != "/v1/models" {
|
|
t.Fatalf("unexpected path %s", r.URL.Path)
|
|
}
|
|
_, _ = w.Write([]byte(`{"object":"list","data":[{"id":"model-a"},{"id":"model-b"}]}`))
|
|
}))
|
|
defer server.Close()
|
|
|
|
adapter := New(config.VllmConf{Endpoint: server.URL}, zap.NewNop())
|
|
caps, err := adapter.Capabilities(context.Background())
|
|
if err != nil {
|
|
t.Fatalf("Capabilities failed: %v", err)
|
|
}
|
|
if got := strings.Join(caps.Targets, ","); got != "model-a,model-b" {
|
|
t.Fatalf("targets: got %q", got)
|
|
}
|
|
}
|
|
|
|
func TestVllmExecuteStreamsDeltas(t *testing.T) {
|
|
var gotModel string
|
|
var gotMessages int
|
|
var gotStream bool
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
if r.URL.Path != "/v1/chat/completions" {
|
|
t.Fatalf("unexpected path %s", r.URL.Path)
|
|
}
|
|
var req struct {
|
|
Model string `json:"model"`
|
|
Messages []any `json:"messages"`
|
|
Stream bool `json:"stream"`
|
|
}
|
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
|
t.Fatalf("decode request: %v", err)
|
|
}
|
|
gotModel = req.Model
|
|
gotMessages = len(req.Messages)
|
|
gotStream = req.Stream
|
|
w.Header().Set("Content-Type", "text/event-stream")
|
|
_, _ = fmt.Fprintf(w, "data: %s\n\n", `{"choices":[{"delta":{"content":"hello "}}]}`)
|
|
_, _ = fmt.Fprintf(w, "data: %s\n\n", `{"choices":[{"delta":{"content":"world"}}]}`)
|
|
_, _ = fmt.Fprintf(w, "data: [DONE]\n\n")
|
|
}))
|
|
defer server.Close()
|
|
|
|
adapter := New(config.VllmConf{Endpoint: server.URL}, zap.NewNop())
|
|
sink := &fakeSink{}
|
|
err := adapter.Execute(context.Background(), runtime.ExecutionSpec{
|
|
RunID: "run-1",
|
|
Target: "llama-3",
|
|
Input: map[string]any{"prompt": "say hello"},
|
|
}, sink)
|
|
if err != nil {
|
|
t.Fatalf("Execute failed: %v", err)
|
|
}
|
|
if gotModel != "llama-3" {
|
|
t.Fatalf("model: got %q", gotModel)
|
|
}
|
|
if gotMessages != 1 {
|
|
t.Fatalf("messages: got %d", gotMessages)
|
|
}
|
|
if !gotStream {
|
|
t.Fatal("expected stream=true")
|
|
}
|
|
|
|
events := sink.all()
|
|
if len(events) != 4 {
|
|
t.Fatalf("expected 4 events (start+delta+delta+complete), got %d: %+v", len(events), events)
|
|
}
|
|
if events[0].Type != runtime.EventTypeStart {
|
|
t.Fatalf("expected start event, got %+v", events[0])
|
|
}
|
|
if events[1].Delta+events[2].Delta != "hello world" {
|
|
t.Fatalf("unexpected deltas: %q + %q", events[1].Delta, events[2].Delta)
|
|
}
|
|
if events[3].Type != runtime.EventTypeComplete {
|
|
t.Fatalf("expected complete event, got %+v", events[3])
|
|
}
|
|
}
|
|
|
|
func TestVllmExecuteUsesMessagesInput(t *testing.T) {
|
|
var gotMessages []map[string]any
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
var req struct {
|
|
Messages []map[string]any `json:"messages"`
|
|
}
|
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
|
t.Fatalf("decode request: %v", err)
|
|
}
|
|
gotMessages = req.Messages
|
|
w.Header().Set("Content-Type", "text/event-stream")
|
|
_, _ = fmt.Fprintf(w, "data: %s\n\n", `{"choices":[{"delta":{"content":"ok"}}]}`)
|
|
_, _ = fmt.Fprintf(w, "data: [DONE]\n\n")
|
|
}))
|
|
defer server.Close()
|
|
|
|
adapter := New(config.VllmConf{Endpoint: server.URL}, zap.NewNop())
|
|
sink := &fakeSink{}
|
|
err := adapter.Execute(context.Background(), runtime.ExecutionSpec{
|
|
RunID: "run-2",
|
|
Target: "llama-3",
|
|
Input: map[string]any{
|
|
"messages": []any{
|
|
map[string]any{"role": "system", "content": "You are helpful."},
|
|
map[string]any{"role": "user", "content": "hi"},
|
|
},
|
|
},
|
|
}, sink)
|
|
if err != nil {
|
|
t.Fatalf("Execute failed: %v", err)
|
|
}
|
|
if len(gotMessages) != 2 {
|
|
t.Fatalf("messages: got %d", len(gotMessages))
|
|
}
|
|
if gotMessages[0]["role"] != "system" || gotMessages[1]["role"] != "user" {
|
|
t.Fatalf("unexpected messages: %+v", gotMessages)
|
|
}
|
|
}
|
|
|
|
func TestVllmExecuteEmitsErrorForHTTPFailure(t *testing.T) {
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
http.Error(w, "service unavailable", http.StatusServiceUnavailable)
|
|
}))
|
|
defer server.Close()
|
|
|
|
adapter := New(config.VllmConf{Endpoint: server.URL}, zap.NewNop())
|
|
sink := &fakeSink{}
|
|
err := adapter.Execute(context.Background(), runtime.ExecutionSpec{
|
|
RunID: "run-err",
|
|
Target: "llama-3",
|
|
Input: map[string]any{"prompt": "hi"},
|
|
}, sink)
|
|
if err == nil {
|
|
t.Fatal("expected error")
|
|
}
|
|
events := sink.all()
|
|
if len(events) < 2 || events[1].Type != runtime.EventTypeError {
|
|
t.Fatalf("expected error event after start, got %+v", events)
|
|
}
|
|
if !strings.Contains(events[1].Error, "503") {
|
|
t.Fatalf("expected status in error, got %q", events[1].Error)
|
|
}
|
|
}
|
|
|
|
func TestVllmInstanceKey(t *testing.T) {
|
|
adapter := New(config.VllmConf{Endpoint: "http://localhost:8000"}, zap.NewNop(), "vllm-gpu")
|
|
caps, _ := adapter.Capabilities(context.Background())
|
|
if caps.InstanceKey != "vllm-gpu" {
|
|
t.Fatalf("InstanceKey: got %q, want vllm-gpu", caps.InstanceKey)
|
|
}
|
|
if caps.AdapterName != Name {
|
|
t.Fatalf("AdapterName: got %q, want %q", caps.AdapterName, Name)
|
|
}
|
|
}
|