248 lines
6.6 KiB
Go
248 lines
6.6 KiB
Go
package vllm
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"net/url"
|
|
"strings"
|
|
"time"
|
|
|
|
"iop/apps/node/internal/runtime"
|
|
)
|
|
|
|
func messagesFromInput(input map[string]any) []vllmMessage {
|
|
if raw, ok := input["messages"].([]any); ok {
|
|
out := make([]vllmMessage, 0, len(raw))
|
|
for _, item := range raw {
|
|
m, ok := item.(map[string]any)
|
|
if !ok {
|
|
continue
|
|
}
|
|
role, _ := m["role"].(string)
|
|
content, _ := m["content"].(string)
|
|
toolCallID, _ := m["tool_call_id"].(string)
|
|
toolCalls := anySlice(m["tool_calls"])
|
|
role = strings.TrimSpace(role)
|
|
content = strings.TrimSpace(content)
|
|
toolCallID = strings.TrimSpace(toolCallID)
|
|
if role == "" || (content == "" && len(toolCalls) == 0) {
|
|
continue
|
|
}
|
|
out = append(out, vllmMessage{
|
|
Role: role,
|
|
Content: content,
|
|
ToolCalls: toolCalls,
|
|
ToolCallID: toolCallID,
|
|
})
|
|
}
|
|
if len(out) > 0 {
|
|
return out
|
|
}
|
|
}
|
|
if prompt := strings.TrimSpace(stringInput(input, "prompt")); prompt != "" {
|
|
return []vllmMessage{{Role: "user", Content: prompt}}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func emitError(ctx context.Context, sink runtime.EventSink, runID, msg string) error {
|
|
return sink.Emit(ctx, runtime.RuntimeEvent{
|
|
RunID: runID,
|
|
Type: runtime.EventTypeError,
|
|
Error: msg,
|
|
Timestamp: time.Now(),
|
|
})
|
|
}
|
|
|
|
func (v *Vllm) doChatCompletion(ctx context.Context, body []byte) (*http.Response, error) {
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, joinURL(v.endpoint, "/v1/chat/completions"), bytes.NewReader(body))
|
|
if err != nil {
|
|
return nil, fmt.Errorf("build request: %w", err)
|
|
}
|
|
req.Header.Set("Content-Type", "application/json")
|
|
return v.client.Do(req)
|
|
}
|
|
|
|
func isAutoToolChoiceUnsupportedError(errorBody string) bool {
|
|
errorBody = strings.ToLower(errorBody)
|
|
return strings.Contains(errorBody, "auto") &&
|
|
strings.Contains(errorBody, "tool choice") &&
|
|
isToolCallingUnsupportedError(errorBody)
|
|
}
|
|
|
|
func isToolCallingUnsupportedError(errorBody string) bool {
|
|
errorBody = strings.ToLower(errorBody)
|
|
return strings.Contains(errorBody, "enable-auto-tool-choice") ||
|
|
strings.Contains(errorBody, "tool-call-parser")
|
|
}
|
|
|
|
func forcedToolChoiceForSingleTool(tools any) (map[string]any, bool) {
|
|
items, ok := tools.([]any)
|
|
if !ok || len(items) != 1 {
|
|
return nil, false
|
|
}
|
|
tool, ok := items[0].(map[string]any)
|
|
if !ok {
|
|
return nil, false
|
|
}
|
|
name := ""
|
|
if fn, ok := tool["function"].(map[string]any); ok {
|
|
name, _ = fn["name"].(string)
|
|
}
|
|
if name == "" {
|
|
name, _ = tool["name"].(string)
|
|
}
|
|
name = strings.TrimSpace(name)
|
|
if name == "" {
|
|
return nil, false
|
|
}
|
|
return map[string]any{
|
|
"type": "function",
|
|
"function": map[string]any{
|
|
"name": name,
|
|
},
|
|
}, true
|
|
}
|
|
|
|
func textToolFallbackChatRequest(req vllmChatRequest, errorBody string) (vllmChatRequest, bool) {
|
|
if !isToolCallingUnsupportedError(errorBody) {
|
|
return vllmChatRequest{}, false
|
|
}
|
|
instruction, ok := textToolFallbackInstruction(req.Tools)
|
|
if !ok {
|
|
return vllmChatRequest{}, false
|
|
}
|
|
next := req
|
|
next.Tools = nil
|
|
next.ToolChoice = nil
|
|
next.Messages = prependTextToolFallbackInstruction(req.Messages, instruction)
|
|
return next, true
|
|
}
|
|
|
|
func prependTextToolFallbackInstruction(messages []vllmMessage, instruction string) []vllmMessage {
|
|
system := vllmMessage{Role: "system", Content: instruction}
|
|
out := make([]vllmMessage, 0, len(messages)+1)
|
|
for _, msg := range messages {
|
|
if strings.EqualFold(strings.TrimSpace(msg.Role), "system") {
|
|
system.Content = joinTextToolFallbackSystemContent(system.Content, msg.Content)
|
|
continue
|
|
}
|
|
out = append(out, msg)
|
|
}
|
|
return append([]vllmMessage{system}, out...)
|
|
}
|
|
|
|
func joinTextToolFallbackSystemContent(first, next string) string {
|
|
first = strings.TrimSpace(first)
|
|
next = strings.TrimSpace(next)
|
|
switch {
|
|
case first == "":
|
|
return next
|
|
case next == "":
|
|
return first
|
|
default:
|
|
return first + "\n\n" + next
|
|
}
|
|
}
|
|
|
|
func textToolFallbackInstruction(tools any) (string, bool) {
|
|
items, ok := anyItems(tools)
|
|
if !ok || len(items) == 0 {
|
|
return "", false
|
|
}
|
|
encoded, err := json.Marshal(items)
|
|
if err != nil {
|
|
return "", false
|
|
}
|
|
return "Tool calls must be emitted as plain text because this backend does not support native OpenAI tool calling. When a tool is needed, respond with exactly one tool call and no markdown:\n" +
|
|
"<tool_call>\n<function=TOOL_NAME>\n<parameter=PARAMETER_NAME>JSON_VALUE</parameter>\n</function>\n</tool_call>\n" +
|
|
"run_commands executes from the client workspace root. Do not prepend cd to an absolute workspace path unless the user explicitly asks to operate in a different directory; prefer current-workspace commands such as git status.\n" +
|
|
"Use valid JSON for each parameter value and follow the supplied parameter schema. Available tools JSON: " + string(encoded), true
|
|
}
|
|
|
|
func anyItems(value any) ([]any, bool) {
|
|
switch v := value.(type) {
|
|
case []any:
|
|
return v, true
|
|
default:
|
|
encoded, err := json.Marshal(v)
|
|
if err != nil {
|
|
return nil, false
|
|
}
|
|
var out []any
|
|
if err := json.Unmarshal(encoded, &out); err != nil {
|
|
return nil, false
|
|
}
|
|
return out, true
|
|
}
|
|
}
|
|
|
|
func completeEvent(runID, finishReason string, usage *runtime.UsageStats, outputTokens int, toolCalls []any, textToolFallback bool) runtime.RuntimeEvent {
|
|
if usage == nil {
|
|
usage = &runtime.UsageStats{OutputTokens: outputTokens}
|
|
}
|
|
var metadata map[string]string
|
|
if finishReason != "" {
|
|
metadata = map[string]string{"finish_reason": finishReason}
|
|
}
|
|
if textToolFallback {
|
|
if metadata == nil {
|
|
metadata = make(map[string]string, 2)
|
|
}
|
|
metadata[runtimeMetadataOpenAITextToolFallback] = "true"
|
|
}
|
|
if len(toolCalls) > 0 {
|
|
if metadata == nil {
|
|
metadata = make(map[string]string, 2)
|
|
}
|
|
if finishReason == "" {
|
|
metadata["finish_reason"] = "tool_calls"
|
|
}
|
|
if encoded, err := json.Marshal(toolCalls); err == nil {
|
|
metadata[runtimeMetadataOpenAIToolCalls] = string(encoded)
|
|
}
|
|
}
|
|
return runtime.RuntimeEvent{
|
|
RunID: runID,
|
|
Type: runtime.EventTypeComplete,
|
|
Message: "vllm chat complete",
|
|
Usage: usage,
|
|
Metadata: metadata,
|
|
Timestamp: time.Now(),
|
|
}
|
|
}
|
|
|
|
func anySlice(v any) []any {
|
|
if items, ok := v.([]any); ok {
|
|
return items
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func stringInput(input map[string]any, key string) string {
|
|
if input == nil {
|
|
return ""
|
|
}
|
|
if v, ok := input[key].(string); ok {
|
|
return v
|
|
}
|
|
return ""
|
|
}
|
|
|
|
func joinURL(baseURL, path string) string {
|
|
u, err := url.Parse(baseURL)
|
|
if err != nil {
|
|
return strings.TrimRight(baseURL, "/") + path
|
|
}
|
|
u.Path = strings.TrimRight(u.Path, "/") + path
|
|
return u.String()
|
|
}
|
|
|
|
func readLimited(r io.Reader, limit int64) string {
|
|
b, _ := io.ReadAll(io.LimitReader(r, limit))
|
|
return strings.TrimSpace(string(b))
|
|
}
|