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

611 lines
21 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,
Capacity: 10,
MaxQueue: 20,
QueueTimeoutMS: 5000,
RequestTimeoutMS: 10000,
}, 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)
}
if caps.MaxConcurrency != 10 {
t.Fatalf("expected MaxConcurrency 10, got %d", caps.MaxConcurrency)
}
if caps.MaxQueue != 20 {
t.Fatalf("expected MaxQueue 20, got %d", caps.MaxQueue)
}
if caps.QueueTimeoutMS != 5000 {
t.Fatalf("expected QueueTimeoutMS 5000, got %d", caps.QueueTimeoutMS)
}
if caps.RequestTimeoutMS != 10000 {
t.Fatalf("expected RequestTimeoutMS 10000, got %d", caps.RequestTimeoutMS)
}
// Verify default capacity fallback
adapterDefault := New(config.VllmConf{Endpoint: server.URL}, zap.NewNop())
capsDefault, err := adapterDefault.Capabilities(context.Background())
if err != nil {
t.Fatalf("Capabilities failed: %v", err)
}
if capsDefault.MaxConcurrency != 8 {
t.Fatalf("expected default MaxConcurrency 8, got %d", capsDefault.MaxConcurrency)
}
}
func TestVllmExecuteParsesReasoningAndCachedInputTokens(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
_, _ = fmt.Fprintf(w, "data: %s\n\n", `{"choices":[{"delta":{"content":"hi"},"finish_reason":"stop"}],"usage":{"prompt_tokens":9,"completion_tokens":6,"prompt_tokens_details":{"cached_tokens":2},"completion_tokens_details":{"reasoning_tokens":5}}}`)
_, _ = fmt.Fprintf(w, "data: [DONE]\n\n")
}))
defer server.Close()
adapter := New(config.VllmConf{Endpoint: server.URL}, zap.NewNop())
sink := &fakeSink{}
if err := adapter.Execute(context.Background(), runtime.ExecutionSpec{
RunID: "run-1",
Target: "llama-3",
Input: map[string]any{"prompt": "hi"},
}, sink); err != nil {
t.Fatalf("Execute failed: %v", err)
}
events := sink.all()
complete := events[len(events)-1]
if complete.Type != runtime.EventTypeComplete || complete.Usage == nil {
t.Fatalf("expected complete event with usage, got %+v", complete)
}
u := complete.Usage
if u.InputTokens != 9 || u.OutputTokens != 6 || u.CachedInputTokens != 2 || u.ReasoningTokens != 5 {
t.Fatalf("usage: got in=%d out=%d cached=%d reasoning=%d, want 9/6/2/5",
u.InputTokens, u.OutputTokens, u.CachedInputTokens, u.ReasoningTokens)
}
}
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 TestVllmExecutePassesToolsAndPreservesNativeToolCalls(t *testing.T) {
var gotToolChoice any
var gotTools []any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var req map[string]any
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
t.Fatalf("decode request: %v", err)
}
gotToolChoice = req["tool_choice"]
gotTools, _ = req["tools"].([]any)
w.Header().Set("Content-Type", "text/event-stream")
_, _ = fmt.Fprintf(w, "data: %s\n\n", `{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"run_commands","arguments":"{\"commands\":"}}]}}]}`)
_, _ = fmt.Fprintf(w, "data: %s\n\n", `{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"[\"git status\"]}"}}]},"finish_reason":"tool_calls"}]}`)
_, _ = 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-tools",
Target: "llama-3",
Input: map[string]any{
"prompt": "status",
"tools": []any{map[string]any{"type": "function"}},
"tool_choice": "auto",
},
}, sink)
if err != nil {
t.Fatalf("Execute failed: %v", err)
}
if gotToolChoice != "auto" || len(gotTools) != 1 {
t.Fatalf("tools/tool_choice not passed: tools=%+v choice=%+v", gotTools, gotToolChoice)
}
events := sink.all()
complete := events[len(events)-1]
if complete.Metadata["finish_reason"] != "tool_calls" {
t.Fatalf("finish_reason: %+v", complete.Metadata)
}
var toolCalls []map[string]any
if err := json.Unmarshal([]byte(complete.Metadata[runtimeMetadataOpenAIToolCalls]), &toolCalls); err != nil {
t.Fatalf("tool_calls metadata JSON: %v", err)
}
fn := toolCalls[0]["function"].(map[string]any)
if fn["arguments"] != `{"commands":["git status"]}` {
t.Fatalf("arguments: %+v", fn["arguments"])
}
}
func TestVllmExecuteRetriesSingleToolWhenAutoUnsupported(t *testing.T) {
attempts := 0
var retryBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
attempts++
if attempts == 1 {
w.WriteHeader(http.StatusBadRequest)
_, _ = w.Write([]byte(`{"error":{"message":"\"auto\" tool choice requires --enable-auto-tool-choice and --tool-call-parser to be set"}}`))
return
}
if err := json.NewDecoder(r.Body).Decode(&retryBody); err != nil {
t.Fatalf("decode retry: %v", err)
}
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-retry",
Target: "llama-3",
Input: map[string]any{
"prompt": "status",
"tools": []any{map[string]any{
"type": "function",
"function": map[string]any{
"name": "run_commands",
},
}},
},
}, sink)
if err != nil {
t.Fatalf("Execute failed: %v", err)
}
if attempts != 2 {
t.Fatalf("attempts: got %d, want 2", attempts)
}
choice, ok := retryBody["tool_choice"].(map[string]any)
if !ok {
t.Fatalf("retry tool_choice missing: %+v", retryBody)
}
fn := choice["function"].(map[string]any)
if choice["type"] != "function" || fn["name"] != "run_commands" {
t.Fatalf("retry tool_choice: %+v", choice)
}
}
func TestVllmExecuteFallsBackToTextToolsWhenNativeToolsUnsupported(t *testing.T) {
attempts := 0
var fallbackBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
attempts++
switch attempts {
case 1:
w.WriteHeader(http.StatusBadRequest)
_, _ = w.Write([]byte(`{"error":{"message":"\"auto\" tool choice requires --enable-auto-tool-choice and --tool-call-parser to be set"}}`))
return
case 2:
w.WriteHeader(http.StatusBadRequest)
_, _ = w.Write([]byte(`{"error":{"message":"tool_choice=\"function=ChatCompletionNamedFunction(name='run_commands') type='function'\" requires --tool-call-parser to be set"}}`))
return
}
if err := json.NewDecoder(r.Body).Decode(&fallbackBody); err != nil {
t.Fatalf("decode fallback: %v", err)
}
w.Header().Set("Content-Type", "text/event-stream")
_, _ = fmt.Fprintf(w, "data: %s\n\n", `{"choices":[{"delta":{"content":"<tool_call>\n<function=run_commands>\n<parameter=commands>[\"git status\"]</parameter>\n</function>\n</tool_call>"}}]}`)
_, _ = 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-text-fallback",
Target: "llama-3",
Input: map[string]any{
"messages": []any{
map[string]any{"role": "system", "content": "Existing Cline system prompt."},
map[string]any{"role": "user", "content": "status"},
},
"tools": []any{map[string]any{
"type": "function",
"function": map[string]any{
"name": "run_commands",
"parameters": map[string]any{
"type": "object",
"properties": map[string]any{
"commands": map[string]any{"type": "array", "items": map[string]any{"type": "string"}},
},
},
},
}},
},
}, sink)
if err != nil {
t.Fatalf("Execute failed: %v", err)
}
if attempts != 3 {
t.Fatalf("attempts: got %d, want 3", attempts)
}
if _, ok := fallbackBody["tools"]; ok {
t.Fatalf("fallback must omit tools: %+v", fallbackBody["tools"])
}
if _, ok := fallbackBody["tool_choice"]; ok {
t.Fatalf("fallback must omit tool_choice: %+v", fallbackBody["tool_choice"])
}
messages := fallbackBody["messages"].([]any)
first := messages[0].(map[string]any)
if first["role"] != "system" || !strings.Contains(first["content"].(string), "<tool_call>") || !strings.Contains(first["content"].(string), "run_commands") {
t.Fatalf("fallback system instruction missing tool format/name: %+v", first)
}
if !strings.Contains(first["content"].(string), "client workspace root") || !strings.Contains(first["content"].(string), "Do not prepend cd") {
t.Fatalf("fallback system instruction missing workspace-root guidance: %+v", first)
}
if !strings.Contains(first["content"].(string), "Existing Cline system prompt.") {
t.Fatalf("fallback system instruction did not preserve existing system content: %+v", first)
}
if len(messages) != 2 || messages[1].(map[string]any)["role"] != "user" {
t.Fatalf("fallback should keep a single leading system message: %+v", messages)
}
events := sink.all()
complete := events[len(events)-1]
if complete.Metadata[runtimeMetadataOpenAITextToolFallback] != "true" {
t.Fatalf("fallback metadata missing: %+v", complete.Metadata)
}
}
func TestVllmExecutePreservesTypedTools(t *testing.T) {
var body map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Fatalf("decode: %v", err)
}
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-typed-tools",
Target: "llama-3",
Input: map[string]any{
"prompt": "status",
"tools": []map[string]any{
{
"type": "function",
"function": map[string]any{
"name": "run_commands",
"description": "Execute commands",
},
},
},
},
}, sink)
if err != nil {
t.Fatalf("Execute failed: %v", err)
}
tools, ok := body["tools"].([]any)
if !ok {
t.Fatalf("tools field missing or not a slice: %+v", body["tools"])
}
if len(tools) != 1 {
t.Fatalf("expected 1 tool, got %d", len(tools))
}
tool := tools[0].(map[string]any)
if tool["type"] != "function" {
t.Fatalf("expected function tool, got %+v", tool)
}
fn := tool["function"].(map[string]any)
if fn["name"] != "run_commands" {
t.Fatalf("expected name run_commands, got %v", fn["name"])
}
}
func TestVllmExecuteStreamsReasoningDeltas(t *testing.T) {
type expectation struct {
Type runtime.EventType
Delta string
}
tests := []struct {
name string
chunks []string
expect []expectation
}{
{
name: "precedence",
chunks: []string{`{"choices":[{"delta":{"reasoning_content":"think_content","reasoning":"think_alias"}}]}`, `{"choices":[{"delta":{"content":"answer"}}]}`},
expect: []expectation{{runtime.EventTypeReasoningDelta, "think_content"}, {runtime.EventTypeDelta, "answer"}},
},
{
name: "fallback",
chunks: []string{`{"choices":[{"delta":{"reasoning":"think_alias"}}]}`, `{"choices":[{"delta":{"content":"answer"}}]}`},
expect: []expectation{{runtime.EventTypeReasoningDelta, "think_alias"}, {runtime.EventTypeDelta, "answer"}},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
for _, chunk := range tc.chunks {
_, _ = fmt.Fprintf(w, "data: %s\n\n", chunk)
}
_, _ = fmt.Fprintf(w, "data: [DONE]\n\n")
}))
defer server.Close()
adapter := New(config.VllmConf{Endpoint: server.URL}, zap.NewNop())
sink := &fakeSink{}
if err := adapter.Execute(context.Background(), runtime.ExecutionSpec{RunID: "run-reasoning", Target: "llama-3", Input: map[string]any{"prompt": "hi"}}, sink); err != nil {
t.Fatalf("Execute failed: %v", err)
}
events := sink.all()
if len(events) != 4 {
t.Fatalf("expected 4 events (start+reasoning_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].Type != tc.expect[0].Type || events[1].Delta != tc.expect[0].Delta {
t.Fatalf("expected reasoning_delta %s=%q, got type=%s delta=%q", tc.expect[0].Type, tc.expect[0].Delta, events[1].Type, events[1].Delta)
}
if events[2].Type != tc.expect[1].Type || events[2].Delta != tc.expect[1].Delta {
t.Fatalf("expected delta %s=%q, got type=%s delta=%q", tc.expect[1].Type, tc.expect[1].Delta, events[2].Type, 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)
}
}
func TestVllmProbeProviderAvailability(t *testing.T) {
t.Run("200_ok_and_target_hit", func(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())
res, err := adapter.ProbeProvider(context.Background(), "model-a")
if err != nil {
t.Fatalf("ProbeProvider failed: %v", err)
}
if res.Status != runtime.ProviderStatusAvailable {
t.Errorf("expected Status available, got %s", res.Status)
}
if len(res.Targets) != 2 || res.Targets[0] != "model-a" {
t.Errorf("unexpected Targets: %+v", res.Targets)
}
})
t.Run("200_ok_and_target_miss", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte(`{"object":"list","data":[{"id":"model-a"}]}`))
}))
defer server.Close()
adapter := New(config.VllmConf{Endpoint: server.URL}, zap.NewNop())
res, err := adapter.ProbeProvider(context.Background(), "model-b")
if err != nil {
t.Fatalf("ProbeProvider failed: %v", err)
}
if res.Status != runtime.ProviderStatusUnavailable {
t.Errorf("expected Status unavailable, got %s", res.Status)
}
if !strings.Contains(res.Detail, "not found") {
t.Errorf("expected 'not found' in detail, got %s", res.Detail)
}
})
t.Run("500_internal_error", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
}))
defer server.Close()
adapter := New(config.VllmConf{Endpoint: server.URL}, zap.NewNop())
res, err := adapter.ProbeProvider(context.Background(), "model-a")
if err != nil {
t.Fatalf("ProbeProvider failed: %v", err)
}
if res.Status != runtime.ProviderStatusUnavailable {
t.Errorf("expected Status unavailable, got %s", res.Status)
}
})
t.Run("empty_target", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte(`{"object":"list","data":[{"id":"model-a"}]}`))
}))
defer server.Close()
adapter := New(config.VllmConf{Endpoint: server.URL}, zap.NewNop())
res, err := adapter.ProbeProvider(context.Background(), "")
if err != nil {
t.Fatalf("ProbeProvider failed: %v", err)
}
if res.Status != runtime.ProviderStatusAvailable {
t.Errorf("expected Status available, got %s", res.Status)
}
})
}