// gr2gw: Universal Gradio to OpenAI LLM proxy gateway in Go
// Created by Luxferre in 2026, released into the public domain
package main
import (
"bufio"
"bytes"
"context"
"crypto/rand"
"encoding/json"
"flag"
"fmt"
"io"
"log"
"net"
"net/http"
"net/url"
"os"
"regexp"
"strconv"
"strings"
"sync"
"time"
)
var (
DefaultSpaceURL = "https://lucasmarchettidelima-digital-twin.hf.space"
DefaultUserAgent = "Mozilla/5.0 (X11; Linux x86_64; rv:153.0) Gecko/20100101 Firefox/153.0"
ConfiguredUserAgent string
)
// ---------------------------------------------------------------------------
// OpenAI API data structures
// ---------------------------------------------------------------------------
type ModelItem struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
OwnedBy string `json:"owned_by"`
}
type ModelsResponse struct {
Object string `json:"object"`
Data []ModelItem `json:"data"`
}
type ToolCallFunction struct {
Name string `json:"name,omitempty"`
Arguments string `json:"arguments"`
}
type ToolCall struct {
Index *int `json:"index,omitempty"`
ID string `json:"id,omitempty"`
Type string `json:"type,omitempty"`
Function ToolCallFunction `json:"function"`
}
type Tool struct {
Type string `json:"type"`
Function interface{} `json:"function"`
}
type ChatMessage struct {
Role string `json:"role"`
Content interface{} `json:"content"`
ReasoningContent string `json:"reasoning_content,omitempty"`
ToolCalls []ToolCall `json:"tool_calls,omitempty"`
ToolCallID string `json:"tool_call_id,omitempty"`
Name string `json:"name,omitempty"`
}
func (m *ChatMessage) GetContentString() string {
if m.Content == nil {
return ""
}
if str, ok := m.Content.(string); ok {
return str
}
if parts, ok := m.Content.([]interface{}); ok {
var sb strings.Builder
for _, p := range parts {
if str, ok := p.(string); ok {
sb.WriteString(str)
} else if itemMap, ok := p.(map[string]interface{}); ok {
if textVal, ok := itemMap["text"].(string); ok {
sb.WriteString(textVal)
}
}
}
return sb.String()
}
b, err := json.Marshal(m.Content)
if err == nil {
return string(b)
}
return fmt.Sprintf("%v", m.Content)
}
type ChatCompletionRequest struct {
Model string `json:"model"`
Messages []ChatMessage `json:"messages"`
Tools []Tool `json:"tools,omitempty"`
ToolChoice interface{} `json:"tool_choice,omitempty"`
Stream bool `json:"stream"`
MaxTokens int `json:"max_tokens"`
MaxCompletionTokens int `json:"max_completion_tokens"`
Temperature *float64 `json:"temperature,omitempty"`
TopP *float64 `json:"top_p,omitempty"`
ReasoningEffort string `json:"reasoning_effort,omitempty"`
}
type ChatCompletionResponseChoice struct {
Index int `json:"index"`
Message ChatMessage `json:"message"`
FinishReason string `json:"finish_reason"`
}
type Usage struct {
PromptTokens int `json:"prompt_tokens"`
CompletionTokens int `json:"completion_tokens"`
TotalTokens int `json:"total_tokens"`
}
type ChatCompletionResponse struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
Model string `json:"model"`
Choices []ChatCompletionResponseChoice `json:"choices"`
Usage Usage `json:"usage"`
}
type StreamDelta struct {
Role string `json:"role,omitempty"`
Content string `json:"content,omitempty"`
ReasoningContent string `json:"reasoning_content,omitempty"`
ToolCalls []ToolCall `json:"tool_calls,omitempty"`
}
type StreamChoice struct {
Index int `json:"index"`
Delta StreamDelta `json:"delta"`
FinishReason *string `json:"finish_reason,omitempty"`
}
type StreamResponse struct {
ID string `json:"id"`
Object string `json:"object"`
Created int64 `json:"created"`
Model string `json:"model"`
Choices []StreamChoice `json:"choices"`
}
// ---------------------------------------------------------------------------
// SOCKS5 proxy client (RFC 1928 / RFC 1929)
// ---------------------------------------------------------------------------
type SOCKS5Config struct {
Address string
Username string
Password string
}
func ParseSOCKS5URL(proxyURL string) (*SOCKS5Config, error) {
cleanURL := strings.TrimSpace(proxyURL)
if cleanURL == "" {
return nil, nil
}
if strings.HasPrefix(cleanURL, "socks5://") {
cleanURL = strings.TrimPrefix(cleanURL, "socks5://")
} else if strings.HasPrefix(cleanURL, "socks5h://") {
cleanURL = strings.TrimPrefix(cleanURL, "socks5h://")
}
cfg := &SOCKS5Config{}
if atIdx := strings.LastIndex(cleanURL, "@"); atIdx != -1 {
userPass := cleanURL[:atIdx]
cfg.Address = cleanURL[atIdx+1:]
if colonIdx := strings.Index(userPass, ":"); colonIdx != -1 {
cfg.Username = userPass[:colonIdx]
cfg.Password = userPass[colonIdx+1:]
} else {
cfg.Username = userPass
}
} else {
cfg.Address = cleanURL
}
if !strings.Contains(cfg.Address, ":") {
cfg.Address = cfg.Address + ":1080"
}
return cfg, nil
}
func DialSOCKS5(ctx context.Context, proxyURL, targetAddr string) (net.Conn, error) {
cfg, err := ParseSOCKS5URL(proxyURL)
if err != nil {
return nil, fmt.Errorf("invalid socks5 proxy configuration: %w", err)
}
if cfg == nil {
var d net.Dialer
return d.DialContext(ctx, "tcp", targetAddr)
}
var d net.Dialer
conn, err := d.DialContext(ctx, "tcp", cfg.Address)
if err != nil {
return nil, fmt.Errorf("failed to connect to socks5 proxy at %s: %w", cfg.Address, err)
}
deadline, ok := ctx.Deadline()
if !ok {
deadline = time.Now().Add(30 * time.Second)
}
conn.SetDeadline(deadline)
defer conn.SetDeadline(time.Time{})
// 1. Negotiation Greeting (RFC 1928)
var greeting []byte
if cfg.Username != "" {
greeting = []byte{0x05, 0x02, 0x00, 0x02}
} else {
greeting = []byte{0x05, 0x01, 0x00}
}
if _, err := conn.Write(greeting); err != nil {
conn.Close()
return nil, fmt.Errorf("failed to write socks5 greeting: %w", err)
}
resp := make([]byte, 2)
if _, err := io.ReadFull(conn, resp); err != nil {
conn.Close()
return nil, fmt.Errorf("failed to read socks5 greeting response: %w", err)
}
if resp[0] != 0x05 {
conn.Close()
return nil, fmt.Errorf("unsupported socks version: 0x%02x", resp[0])
}
// 2. Authentication if required (RFC 1929)
if resp[1] == 0x02 {
if cfg.Username == "" {
conn.Close()
return nil, fmt.Errorf("socks5 proxy requires authentication, but no credentials provided")
}
uLen := byte(len(cfg.Username))
pLen := byte(len(cfg.Password))
authReq := []byte{0x01, uLen}
authReq = append(authReq, []byte(cfg.Username)...)
authReq = append(authReq, pLen)
authReq = append(authReq, []byte(cfg.Password)...)
if _, err := conn.Write(authReq); err != nil {
conn.Close()
return nil, fmt.Errorf("failed to send socks5 authentication: %w", err)
}
authResp := make([]byte, 2)
if _, err := io.ReadFull(conn, authResp); err != nil {
conn.Close()
return nil, fmt.Errorf("failed to read socks5 auth response: %w", err)
}
if authResp[1] != 0x00 {
conn.Close()
return nil, fmt.Errorf("socks5 authentication failed with status 0x%02x", authResp[1])
}
} else if resp[1] != 0x00 {
conn.Close()
return nil, fmt.Errorf("socks5 proxy rejected authentication methods: 0x%02x", resp[1])
}
// 3. Connection Request (CONNECT command)
host, portStr, err := net.SplitHostPort(targetAddr)
if err != nil {
conn.Close()
return nil, fmt.Errorf("invalid target address %s: %w", targetAddr, err)
}
port, err := strconv.Atoi(portStr)
if err != nil || port < 1 || port > 65535 {
conn.Close()
return nil, fmt.Errorf("invalid port in target address %s", targetAddr)
}
reqBuf := []byte{0x05, 0x01, 0x00}
ip := net.ParseIP(host)
if ip4 := ip.To4(); ip4 != nil {
reqBuf = append(reqBuf, 0x01)
reqBuf = append(reqBuf, ip4...)
} else if ip6 := ip.To16(); ip6 != nil {
reqBuf = append(reqBuf, 0x04)
reqBuf = append(reqBuf, ip6...)
} else {
reqBuf = append(reqBuf, 0x03, byte(len(host)))
reqBuf = append(reqBuf, []byte(host)...)
}
reqBuf = append(reqBuf, byte(port>>8), byte(port&0xFF))
if _, err := conn.Write(reqBuf); err != nil {
conn.Close()
return nil, fmt.Errorf("failed to send socks5 connect request: %w", err)
}
// 4. Connection Response
respHdr := make([]byte, 4)
if _, err := io.ReadFull(conn, respHdr); err != nil {
conn.Close()
return nil, fmt.Errorf("failed to read socks5 connect response: %w", err)
}
if respHdr[1] != 0x00 {
conn.Close()
return nil, fmt.Errorf("socks5 connect failed with reply code 0x%02x", respHdr[1])
}
switch respHdr[3] {
case 0x01:
bnd := make([]byte, 6)
io.ReadFull(conn, bnd)
case 0x03:
lenBuf := make([]byte, 1)
io.ReadFull(conn, lenBuf)
bnd := make([]byte, int(lenBuf[0])+2)
io.ReadFull(conn, bnd)
case 0x04:
bnd := make([]byte, 18)
io.ReadFull(conn, bnd)
}
return conn, nil
}
// ---------------------------------------------------------------------------
// Helper utilities
// ---------------------------------------------------------------------------
func GenerateUUID() string {
var b [16]byte
_, err := rand.Read(b[:])
if err != nil {
return fmt.Sprintf("%d", time.Now().UnixNano())
}
b[6] = (b[6] & 0x0f) | 0x40
b[8] = (b[8] & 0x3f) | 0x80
return fmt.Sprintf("%x-%x-%x-%x-%x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:])
}
func GenerateSessionHash() string {
const chars = "abcdefghijklmnopqrstuvwxyz0123456789"
var b [12]byte
rand.Read(b[:])
var sb strings.Builder
for _, v := range b {
sb.WriteByte(chars[int(v)%len(chars)])
}
return sb.String()
}
func FibonacciDelay(attempt int) time.Duration {
if attempt <= 0 {
return 1 * time.Second
}
a, b := 1, 1
for i := 1; i < attempt; i++ {
a, b = b, a+b
}
return time.Duration(a) * time.Second
}
func DoWithFibonacciRetry(client *http.Client, makeReq func() (*http.Request, error), maxRetries int) (*http.Response, error) {
var lastErr error
for attempt := 1; attempt <= maxRetries; attempt++ {
req, err := makeReq()
if err != nil {
return nil, err
}
resp, err := client.Do(req)
if err == nil && resp.StatusCode == http.StatusOK {
return resp, nil
}
if resp != nil {
respBody, _ := io.ReadAll(resp.Body)
resp.Body.Close()
lastErr = fmt.Errorf("HTTP status %d: %s", resp.StatusCode, string(respBody))
if _, ok := extractFailedGeneration(string(respBody)); ok {
break
}
} else {
lastErr = err
}
if attempt < maxRetries {
delay := FibonacciDelay(attempt)
time.Sleep(delay)
}
}
return nil, fmt.Errorf("request failed after %d retries: %v", maxRetries, lastErr)
}
func EnableCORS(w http.ResponseWriter) {
w.Header().Set("Access-Control-Allow-Origin", "*")
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization, api-key, X-User-Agent, X-Gradio-Space, X-Space-URL")
}
func ResolveMaxTokens(req ChatCompletionRequest) int {
mt := req.MaxTokens
if mt == 0 && req.MaxCompletionTokens > 0 {
mt = req.MaxCompletionTokens
}
if mt <= 0 {
mt = 131072
}
return mt
}
func EffectiveUserAgent(r *http.Request) string {
if r != nil {
if c := r.Header.Get("X-User-Agent"); c != "" {
return c
}
}
if ConfiguredUserAgent != "" {
return ConfiguredUserAgent
}
return DefaultUserAgent
}
// ---------------------------------------------------------------------------
// Tool and message processing
// ---------------------------------------------------------------------------
func BuildToolInstruction(tools []Tool) string {
if len(tools) == 0 {
return ""
}
toolsBytes, _ := json.MarshalIndent(tools, "", " ")
sampleFnName := ""
for _, t := range tools {
if fnMap, ok := t.Function.(map[string]interface{}); ok {
if n, ok := fnMap["name"].(string); ok && n != "" {
sampleFnName = n
break
}
}
}
if sampleFnName == "" {
sampleFnName = "example_tool"
}
return fmt.Sprintf(`# Tool Calling Instructions
You are equipped with external tools to assist with user queries.
You have access to the following tools:
%s
When the user asks a question or makes a request that can be fulfilled, assisted, or answered using any of the tools above, you MUST call the appropriate tool.
DO NOT refuse to answer, and DO NOT claim that you lack real-time access, live data, or tool capabilities when a tool is provided for that purpose.
## Tool Calling Syntax
To call a tool, you MUST output a block formatted as:
{"name": "%s", "arguments": {...}}
## Rules:
1. If an available tool is relevant to the user's request, invoking the tool is MANDATORY.
2. When calling a tool, your ENTIRE output must consist ONLY of the block. Do not add introductory text, commentary, or conversational filler.
3. If multiple tools are required, output each in its own block.
4. Arguments must be a valid JSON object strictly matching the tool's parameter definitions.
5. If no tools are relevant to the user's inquiry, respond normally with plain text.
6. When tool execution results are provided to you in subsequent turns (via blocks or tool role messages), formulate your final answer to the user based on those results.`, string(toolsBytes), sampleFnName)
}
func TransformMessages(req ChatCompletionRequest) (processed []ChatMessage, toolInstruction string, hasSystem bool) {
// 1. Build lookup map from tool_call_id to function name across all assistant messages
toolIDToName := make(map[string]string)
for _, msg := range req.Messages {
for _, tc := range msg.ToolCalls {
if tc.ID != "" && tc.Function.Name != "" {
toolIDToName[tc.ID] = tc.Function.Name
}
}
}
toolInstruction = BuildToolInstruction(req.Tools)
// 2. Process and coalesce messages preserving turn parity
var staged []ChatMessage
for i := 0; i < len(req.Messages); i++ {
msg := req.Messages[i]
contentStr := msg.GetContentString()
switch msg.Role {
case "system":
hasSystem = true
staged = append(staged, ChatMessage{Role: "system", Content: contentStr})
case "assistant":
var sb strings.Builder
if contentStr != "" {
sb.WriteString(contentStr)
}
for _, tc := range msg.ToolCalls {
if sb.Len() > 0 {
sb.WriteString("\n")
}
args := tc.Function.Arguments
if strings.TrimSpace(args) == "" {
args = "{}"
}
sb.WriteString(fmt.Sprintf("\n{\"name\": %q, \"arguments\": %s}\n", tc.Function.Name, args))
}
staged = append(staged, ChatMessage{
Role: "assistant",
Content: sb.String(),
ReasoningContent: msg.ReasoningContent,
ToolCalls: msg.ToolCalls,
})
case "tool", "function":
// Gather consecutive tool returns into a coalesced turn
var toolResponses []string
j := i
for j < len(req.Messages) && (req.Messages[j].Role == "tool" || req.Messages[j].Role == "function") {
tMsg := req.Messages[j]
tContent := tMsg.GetContentString()
tName := tMsg.Name
if tName == "" && tMsg.ToolCallID != "" {
if mapped, ok := toolIDToName[tMsg.ToolCallID]; ok {
tName = mapped
} else {
tName = tMsg.ToolCallID
}
}
var contentJSON []byte
if json.Valid([]byte(tContent)) {
contentJSON = []byte(tContent)
} else {
contentJSON, _ = json.Marshal(tContent)
}
toolResponses = append(toolResponses, fmt.Sprintf("\n{\"name\": %q, \"content\": %s}\n", tName, string(contentJSON)))
j++
}
i = j - 1 // advance loop
promptSuffix := "Please answer the user's request based on the tool result."
if len(toolResponses) > 1 {
promptSuffix = "Please answer the user's request based on the tool results."
}
coalesced := strings.Join(toolResponses, "\n") + "\n\n" + promptSuffix
staged = append(staged, ChatMessage{
Role: "user",
Content: coalesced,
})
default: // "user" or other roles
staged = append(staged, ChatMessage{Role: msg.Role, Content: contentStr})
}
}
// 3. Inject tool instructions into system prompt
if toolInstruction != "" {
if hasSystem {
systemInjected := false
for i, m := range staged {
if m.Role == "system" {
staged[i].Content = m.GetContentString() + "\n\n" + strings.TrimSpace(toolInstruction)
systemInjected = true
break
}
}
if !systemInjected {
staged = append([]ChatMessage{
{Role: "system", Content: strings.TrimSpace(toolInstruction)},
}, staged...)
}
} else {
staged = append([]ChatMessage{
{Role: "system", Content: strings.TrimSpace(toolInstruction)},
}, staged...)
hasSystem = true
}
}
return staged, toolInstruction, hasSystem
}
func cleanJSONBlock(input string) string {
s := strings.TrimSpace(input)
if strings.HasPrefix(s, "```") {
lines := strings.Split(s, "\n")
if len(lines) >= 2 {
if strings.HasPrefix(lines[len(lines)-1], "```") {
lines = lines[1 : len(lines)-1]
} else {
lines = lines[1:]
}
s = strings.TrimSpace(strings.Join(lines, "\n"))
}
}
return s
}
func sanitizeJSONValue(v interface{}) interface{} {
switch val := v.(type) {
case string:
return strings.TrimSpace(val)
case map[string]interface{}:
res := make(map[string]interface{})
for k, item := range val {
res[k] = sanitizeJSONValue(item)
}
return res
case []interface{}:
res := make([]interface{}, len(val))
for i, item := range val {
res[i] = sanitizeJSONValue(item)
}
return res
default:
return v
}
}
type ToolTagPair struct {
Start string
End string
}
var ToolTagPairs = []ToolTagPair{
{Start: "", End: ""},
{Start: "", End: ""},
{Start: "", End: ""},
{Start: "[TOOL_CALLS]", End: "[/TOOL_CALLS]"},
}
func getToolStartPrefixes() []string {
seen := make(map[string]bool)
var prefixes []string
for _, pair := range ToolTagPairs {
for i := 1; i <= len(pair.Start); i++ {
pref := pair.Start[:i]
if !seen[pref] {
seen[pref] = true
prefixes = append(prefixes, pref)
}
}
}
return prefixes
}
var toolStartPrefixes = getToolStartPrefixes()
func repairToolCallJSON(input string) (ToolCall, bool) {
s := strings.TrimSpace(input)
reName := regexp.MustCompile(`"(?:name|function|action|call)"\s*:\s*"([^"]+)"`)
matches := reName.FindStringSubmatch(s)
if len(matches) < 2 {
return ToolCall{}, false
}
fnName := matches[1]
reArgsObj := regexp.MustCompile(`"(?:arguments|parameters|args|input)"\s*:\s*(\{[\s\S]*\})`)
argMatches := reArgsObj.FindStringSubmatch(s)
argsStr := "{}"
if len(argMatches) >= 2 {
candidate := argMatches[1]
var dummy map[string]interface{}
if json.Unmarshal([]byte(candidate), &dummy) == nil {
argsStr = candidate
}
} else {
reArgsStr := regexp.MustCompile(`"(?:arguments|parameters|args|input)"\s*:\s*"((?:\\.|[^"\\])*)"`)
strMatches := reArgsStr.FindStringSubmatch(s)
if len(strMatches) >= 2 {
var unescaped string
if json.Unmarshal([]byte(`"`+strMatches[1]+`"`), &unescaped) == nil {
argsStr = unescaped
}
}
}
return ToolCall{
ID: "call_" + GenerateUUID()[:8],
Type: "function",
Function: ToolCallFunction{
Name: fnName,
Arguments: argsStr,
},
}, true
}
func parseSingleToolCall(jsonStr string) (ToolCall, bool) {
cleaned := cleanJSONBlock(jsonStr)
var raw map[string]interface{}
if err := json.Unmarshal([]byte(cleaned), &raw); err == nil {
sanitizedRaw, ok := sanitizeJSONValue(raw).(map[string]interface{})
if !ok {
sanitizedRaw = raw
}
for _, wrapperKey := range []string{"function", "function_call", "tool_call"} {
if fnObj, ok := sanitizedRaw[wrapperKey].(map[string]interface{}); ok {
if nameVal, ok := fnObj["name"].(string); ok && nameVal != "" {
argsStr := "{}"
var argsVal interface{}
if a, hasA := fnObj["arguments"]; hasA {
argsVal = a
} else if p, hasP := fnObj["parameters"]; hasP {
argsVal = p
} else if args, hasArgs := fnObj["args"]; hasArgs {
argsVal = args
}
if argsVal != nil {
if s, isStr := argsVal.(string); isStr {
var innerObj interface{}
if json.Unmarshal([]byte(s), &innerObj) == nil {
b, _ := json.Marshal(sanitizeJSONValue(innerObj))
argsStr = string(b)
} else {
argsStr = strings.TrimSpace(s)
}
} else {
b, _ := json.Marshal(sanitizeJSONValue(argsVal))
argsStr = string(b)
}
}
return ToolCall{
ID: "call_" + GenerateUUID()[:8],
Type: "function",
Function: ToolCallFunction{
Name: nameVal,
Arguments: argsStr,
},
}, true
}
}
}
nameVal := ""
for _, key := range []string{"name", "function", "action", "call"} {
if n, ok := sanitizedRaw[key].(string); ok && n != "" {
nameVal = n
break
}
}
if nameVal != "" {
argsStr := "{}"
var argsVal interface{}
for _, key := range []string{"arguments", "parameters", "args", "input"} {
if a, ok := sanitizedRaw[key]; ok {
argsVal = a
break
}
}
if argsVal != nil {
if s, isStr := argsVal.(string); isStr {
var innerObj interface{}
if json.Unmarshal([]byte(s), &innerObj) == nil {
b, _ := json.Marshal(sanitizeJSONValue(innerObj))
argsStr = string(b)
} else {
argsStr = strings.TrimSpace(s)
}
} else {
b, _ := json.Marshal(sanitizeJSONValue(argsVal))
argsStr = string(b)
}
} else {
argsMap := make(map[string]interface{})
for k, v := range sanitizedRaw {
if k != "name" && k != "function" && k != "type" && k != "action" && k != "call" {
argsMap[k] = v
}
}
if len(argsMap) > 0 {
b, _ := json.Marshal(sanitizeJSONValue(argsMap))
argsStr = string(b)
}
}
return ToolCall{
ID: "call_" + GenerateUUID()[:8],
Type: "function",
Function: ToolCallFunction{
Name: nameVal,
Arguments: argsStr,
},
}, true
}
}
return repairToolCallJSON(cleaned)
}
func parseMultipleToolCalls(raw string) ([]ToolCall, bool) {
cleaned := cleanJSONBlock(raw)
if cleaned == "" {
return nil, false
}
// 1. Direct JSON array: [{"name":...}, ...]
var rawList []interface{}
if err := json.Unmarshal([]byte(cleaned), &rawList); err == nil {
var calls []ToolCall
for _, item := range rawList {
b, err := json.Marshal(item)
if err == nil {
if tc, ok := parseSingleToolCall(string(b)); ok {
calls = append(calls, tc)
}
}
}
if len(calls) > 0 {
return calls, true
}
}
// 2. Wrapper object with tool_calls / calls array
var rawMap map[string]interface{}
if err := json.Unmarshal([]byte(cleaned), &rawMap); err == nil {
for _, listKey := range []string{"tool_calls", "calls", "functions"} {
if subArr, ok := rawMap[listKey].([]interface{}); ok && len(subArr) > 0 {
var calls []ToolCall
for _, item := range subArr {
b, err := json.Marshal(item)
if err == nil {
if tc, ok := parseSingleToolCall(string(b)); ok {
calls = append(calls, tc)
}
}
}
if len(calls) > 0 {
return calls, true
}
}
}
}
// 3. Single tool call
if tc, ok := parseSingleToolCall(cleaned); ok {
return []ToolCall{tc}, true
}
return nil, false
}
func parseXMLToolCall(block string) ([]ToolCall, bool) {
inner := strings.TrimSpace(block)
for _, pair := range ToolTagPairs {
inner = strings.ReplaceAll(inner, pair.Start, "")
inner = strings.ReplaceAll(inner, pair.End, "")
}
inner = cleanJSONBlock(inner)
if calls, ok := parseMultipleToolCalls(inner); ok && len(calls) > 0 {
return calls, true
}
var fnName string
if strings.Contains(inner, "") && strings.Contains(inner, "") {
nStart := strings.Index(inner, "") + len("")
nEnd := strings.Index(inner, "")
if nStart < nEnd {
fnName = strings.TrimSpace(inner[nStart:nEnd])
}
}
var argsStr string
if strings.Contains(inner, "") && strings.Contains(inner, "") {
aStart := strings.Index(inner, "") + len("")
aEnd := strings.Index(inner, "")
if aStart < aEnd {
argsStr = strings.TrimSpace(inner[aStart:aEnd])
}
}
if fnName != "" {
if argsStr == "" {
argsStr = "{}"
}
return []ToolCall{{
ID: "call_" + GenerateUUID()[:8],
Type: "function",
Function: ToolCallFunction{
Name: fnName,
Arguments: argsStr,
},
}}, true
}
return nil, false
}
func ExtractToolCallBlocks(content string) (blocks []string, remaining string) {
remaining = content
for _, pair := range ToolTagPairs {
for strings.Contains(remaining, pair.Start) {
sIdx := strings.Index(remaining, pair.Start)
rest := remaining[sIdx+len(pair.Start):]
relNextSIdx := strings.Index(rest, pair.Start)
var nextSIdx int
if relNextSIdx != -1 {
nextSIdx = sIdx + len(pair.Start) + relNextSIdx
} else {
nextSIdx = -1
}
relEIdx := strings.Index(rest, pair.End)
var eIdx int
if relEIdx != -1 {
eIdx = sIdx + len(pair.Start) + relEIdx
} else {
eIdx = -1
}
var blockText string
if eIdx != -1 && (nextSIdx == -1 || eIdx < nextSIdx) {
blockEndPos := eIdx + len(pair.End)
blockText = remaining[sIdx:blockEndPos]
remaining = strings.TrimSpace(remaining[:sIdx] + remaining[blockEndPos:])
} else if nextSIdx != -1 {
blockEndPos := nextSIdx
blockText = remaining[sIdx:blockEndPos]
remaining = strings.TrimSpace(remaining[:sIdx] + remaining[blockEndPos:])
} else {
blockText = remaining[sIdx:]
remaining = strings.TrimSpace(remaining[:sIdx])
}
blocks = append(blocks, blockText)
}
}
return blocks, remaining
}
func DetectToolCalls(content string) ([]ToolCall, string, bool) {
blocks, remaining := ExtractToolCallBlocks(content)
var calls []ToolCall
for _, block := range blocks {
if tcs, ok := parseXMLToolCall(block); ok {
calls = append(calls, tcs...)
}
}
if len(calls) > 0 {
return calls, remaining, true
}
if tcs, ok := parseMultipleToolCalls(strings.TrimSpace(content)); ok && len(tcs) > 0 {
return tcs, "", true
}
return nil, content, false
}
func ExtractThinking(content string) (string, string) {
if strings.Contains(content, "") && strings.Contains(content, "") {
start := strings.Index(content, "")
end := strings.Index(content, "")
if start < end {
reasoning := content[start+len("") : end]
rem := content[:start] + content[end+len(""):]
rem = strings.TrimPrefix(rem, "\n\n")
rem = strings.TrimPrefix(rem, "\n")
return rem, reasoning
}
}
return content, ""
}
// ---------------------------------------------------------------------------
// Response framing & streamer
// ---------------------------------------------------------------------------
type FinalOutput struct {
Content interface{}
ReasoningContent string
ToolCalls []ToolCall
FinishReason string
}
func WriteCompletionResponse(w http.ResponseWriter, completionID string, created int64, model string, out FinalOutput) {
finish := out.FinishReason
if finish == "" {
finish = "stop"
}
resp := ChatCompletionResponse{
ID: completionID,
Object: "chat.completion",
Created: created,
Model: model,
Choices: []ChatCompletionResponseChoice{
{
Index: 0,
Message: ChatMessage{
Role: "assistant",
Content: out.Content,
ReasoningContent: out.ReasoningContent,
ToolCalls: out.ToolCalls,
},
FinishReason: finish,
},
},
Usage: Usage{
PromptTokens: 0,
CompletionTokens: 0,
TotalTokens: 0,
},
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
}
func extractFailedGeneration(raw string) (string, bool) {
var obj map[string]interface{}
if err := json.Unmarshal([]byte(raw), &obj); err == nil {
if fg, ok := obj["failed_generation"].(string); ok && fg != "" {
return strings.TrimSpace(fg), true
}
if errVal, ok := obj["error"]; ok {
if errMap, ok := errVal.(map[string]interface{}); ok {
if fg, ok := errMap["failed_generation"].(string); ok && fg != "" {
return strings.TrimSpace(fg), true
}
} else if errStr, ok := errVal.(string); ok {
if fg, ok := extractFailedGeneration(errStr); ok {
return fg, true
}
}
}
}
keyIdx := strings.Index(raw, `"failed_generation"`)
if keyIdx == -1 {
keyIdx = strings.Index(raw, `'failed_generation'`)
}
if keyIdx == -1 {
keyIdx = strings.Index(raw, `failed_generation`)
}
if keyIdx == -1 {
return "", false
}
colonIdx := strings.Index(raw[keyIdx:], ":")
if colonIdx == -1 {
return "", false
}
valStart := keyIdx + colonIdx + 1
for valStart < len(raw) && (raw[valStart] == ' ' || raw[valStart] == '\t' || raw[valStart] == '\r' || raw[valStart] == '\n') {
valStart++
}
if valStart >= len(raw) {
return "", false
}
firstChar := raw[valStart]
var candidate string
if firstChar == '\'' || firstChar == '"' {
quoteChar := firstChar
var b strings.Builder
escaped := false
for i := valStart + 1; i < len(raw); i++ {
ch := raw[i]
if escaped {
switch ch {
case 'n':
b.WriteByte('\n')
case 'r':
b.WriteByte('\r')
case 't':
b.WriteByte('\t')
case '\\':
b.WriteByte('\\')
case '\'':
b.WriteByte('\'')
case '"':
b.WriteByte('"')
default:
b.WriteByte('\\')
b.WriteByte(ch)
}
escaped = false
} else if ch == '\\' {
escaped = true
} else if ch == quoteChar {
break
} else {
b.WriteByte(ch)
}
}
candidate = strings.TrimSpace(b.String())
} else if firstChar == '{' || firstChar == '[' {
openChar := firstChar
closeChar := byte('}')
if openChar == '[' {
closeChar = ']'
}
depth := 0
inStr := false
var strQuote byte
escaped := false
endIdx := -1
for i := valStart; i < len(raw); i++ {
ch := raw[i]
if inStr {
if escaped {
escaped = false
} else if ch == '\\' {
escaped = true
} else if ch == strQuote {
inStr = false
}
} else {
if ch == '"' || ch == '\'' {
inStr = true
strQuote = ch
} else if ch == openChar {
depth++
} else if ch == closeChar {
depth--
if depth == 0 {
endIdx = i + 1
break
}
}
}
}
if endIdx != -1 {
candidate = strings.TrimSpace(raw[valStart:endIdx])
}
}
if candidate != "" {
if strings.Contains(candidate, `\"`) {
candidate = strings.ReplaceAll(candidate, `\"`, `"`)
}
if strings.Contains(candidate, `\n`) {
candidate = strings.ReplaceAll(candidate, `\n`, "\n")
}
return candidate, true
}
return "", false
}
func extractGradioErrorMessage(dataStr string) string {
var errObj map[string]interface{}
if err := json.Unmarshal([]byte(dataStr), &errObj); err == nil {
if e, ok := errObj["error"].(string); ok && e != "" {
return e
}
if m, ok := errObj["message"].(string); ok && m != "" {
return m
}
if eNull, ok := errObj["error"]; ok && eNull == nil {
if t, ok := errObj["title"].(string); ok && t != "" {
return t
}
return "internal space error (check Gradio inputs/types)"
}
}
clean := strings.TrimSpace(dataStr)
if clean != "" && clean != "null" {
return clean
}
return "unknown upstream Gradio error"
}
type Streamer struct {
w http.ResponseWriter
flusher http.Flusher
id string
created int64
model string
started bool
}
func NewStreamer(w http.ResponseWriter, flusher http.Flusher, id string, created int64, model string) *Streamer {
return &Streamer{w: w, flusher: flusher, id: id, created: created, model: model}
}
func (s *Streamer) EnsureStarted() {
if s.started {
return
}
s.w.Header().Set("Content-Type", "text/event-stream")
s.w.Header().Set("Cache-Control", "no-cache")
s.w.Header().Set("Connection", "keep-alive")
s.started = true
s.Role()
}
func (s *Streamer) Role() {
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{Role: "assistant"})
}
func (s *Streamer) Reasoning(text string) {
s.EnsureStarted()
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{ReasoningContent: text})
}
func (s *Streamer) Content(text string) {
s.EnsureStarted()
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{Content: text})
}
func (s *Streamer) ToolCallDelta(tc ToolCall) {
s.EnsureStarted()
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{ToolCalls: []ToolCall{tc}})
}
func (s *Streamer) Finish(reason string) {
if !s.started {
return
}
sendStreamChunk(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{}, &reason)
}
func (s *Streamer) Done() {
if !s.started {
return
}
fmt.Fprintf(s.w, "data: [DONE]\n\n")
if s.flusher != nil {
s.flusher.Flush()
}
}
func (s *Streamer) Error(errMsg string) {
s.EnsureStarted()
errChunk := map[string]interface{}{
"error": map[string]interface{}{
"message": errMsg,
"type": "upstream_error",
"code": 502,
},
}
b, _ := json.Marshal(errChunk)
fmt.Fprintf(s.w, "data: %s\n\n", b)
if s.flusher != nil {
s.flusher.Flush()
}
}
func sendStreamDelta(w http.ResponseWriter, flusher http.Flusher, completionID string, createdTime int64, modelName string, delta StreamDelta) {
sendStreamChunk(w, flusher, completionID, createdTime, modelName, delta, nil)
}
func sendStreamChunk(w http.ResponseWriter, flusher http.Flusher, completionID string, createdTime int64, modelName string, delta StreamDelta, finishReason *string) {
chunk := StreamResponse{
ID: completionID,
Object: "chat.completion.chunk",
Created: createdTime,
Model: modelName,
Choices: []StreamChoice{
{
Index: 0,
Delta: delta,
FinishReason: finishReason,
},
},
}
b, _ := json.Marshal(chunk)
fmt.Fprintf(w, "data: %s\n\n", b)
if flusher != nil {
flusher.Flush()
}
}
// ---------------------------------------------------------------------------
// Stateful thinking tag filter for streaming
// ---------------------------------------------------------------------------
type StreamThinkingFilter struct {
inThinking bool
buf string
}
func NewStreamThinkingFilter() *StreamThinkingFilter {
return &StreamThinkingFilter{}
}
func hasPrefixOf(target string, prefixes []string) int {
maxMatch := 0
for _, p := range prefixes {
if strings.HasSuffix(target, p) && len(p) > maxMatch {
maxMatch = len(p)
}
}
return maxMatch
}
func hasSuffixPrefixOf(target string, tag string) int {
maxMatch := 0
for i := 1; i < len(tag); i++ {
p := tag[:i]
if strings.HasSuffix(target, p) && len(p) > maxMatch {
maxMatch = len(p)
}
}
return maxMatch
}
func (f *StreamThinkingFilter) Feed(chunk string, onContent func(string), onReasoning func(string)) {
f.buf += chunk
thinkStartTag := ""
thinkEndTag := ""
thinkStartPrefixes := []string{"<", " 0 {
if !f.inThinking {
if idx := strings.Index(f.buf, thinkStartTag); idx != -1 {
before := f.buf[:idx]
if before != "" {
onContent(before)
}
f.inThinking = true
f.buf = f.buf[idx+len(thinkStartTag):]
} else if matchLen := hasPrefixOf(f.buf, thinkStartPrefixes); matchLen > 0 {
safe := f.buf[:len(f.buf)-matchLen]
if safe != "" {
onContent(safe)
}
f.buf = f.buf[len(f.buf)-matchLen:]
break
} else {
onContent(f.buf)
f.buf = ""
break
}
} else {
if idx := strings.Index(f.buf, thinkEndTag); idx != -1 {
before := f.buf[:idx]
if before != "" {
onReasoning(before)
}
f.inThinking = false
f.buf = f.buf[idx+len(thinkEndTag):]
f.buf = strings.TrimPrefix(f.buf, "\n\n")
f.buf = strings.TrimPrefix(f.buf, "\n")
} else if matchLen := hasPrefixOf(f.buf, thinkEndPrefixes); matchLen > 0 {
safe := f.buf[:len(f.buf)-matchLen]
if safe != "" {
onReasoning(safe)
}
f.buf = f.buf[len(f.buf)-matchLen:]
break
} else {
onReasoning(f.buf)
f.buf = ""
break
}
}
}
}
func (f *StreamThinkingFilter) Flush(onContent func(string), onReasoning func(string)) {
if len(f.buf) > 0 {
if f.inThinking {
onReasoning(f.buf)
} else {
onContent(f.buf)
}
f.buf = ""
}
}
// ---------------------------------------------------------------------------
// Stateful tool call tag filter for streaming
// ---------------------------------------------------------------------------
type StreamToolCallFilter struct {
inToolCall bool
buf string
toolCallBuf string
toolIndex int
emittedCall bool
activePair ToolTagPair
activeEndTag string
}
func NewStreamToolCallFilter() *StreamToolCallFilter {
return &StreamToolCallFilter{}
}
func (f *StreamToolCallFilter) Feed(chunk string, onContent func(string), onToolCall func(ToolCall)) {
f.buf += chunk
for len(f.buf) > 0 {
if !f.inToolCall {
earliestIdx := -1
var matchedPair ToolTagPair
for _, pair := range ToolTagPairs {
if idx := strings.Index(f.buf, pair.Start); idx != -1 {
if earliestIdx == -1 || idx < earliestIdx {
earliestIdx = idx
matchedPair = pair
}
}
}
if earliestIdx != -1 {
before := f.buf[:earliestIdx]
if before != "" {
onContent(before)
}
f.inToolCall = true
f.activePair = matchedPair
f.activeEndTag = matchedPair.End
f.buf = f.buf[earliestIdx+len(matchedPair.Start):]
} else if matchLen := hasPrefixOf(f.buf, toolStartPrefixes); matchLen > 0 {
safe := f.buf[:len(f.buf)-matchLen]
if safe != "" {
onContent(safe)
}
f.buf = f.buf[len(f.buf)-matchLen:]
break
} else {
onContent(f.buf)
f.buf = ""
break
}
} else {
if idx := strings.Index(f.buf, f.activeEndTag); idx != -1 {
f.toolCallBuf += f.buf[:idx]
f.buf = f.buf[idx+len(f.activeEndTag):]
f.inToolCall = false
if tcs, ok := parseMultipleToolCalls(f.toolCallBuf); ok && len(tcs) > 0 {
for _, tc := range tcs {
idxCopy := f.toolIndex
tc.Index = &idxCopy
f.toolIndex++
f.emittedCall = true
onToolCall(tc)
}
} else if tcs2, ok2 := parseXMLToolCall(f.activePair.Start + f.toolCallBuf + f.activePair.End); ok2 && len(tcs2) > 0 {
for _, tc := range tcs2 {
idxCopy := f.toolIndex
tc.Index = &idxCopy
f.toolIndex++
f.emittedCall = true
onToolCall(tc)
}
} else {
onContent(f.activePair.Start + f.toolCallBuf + f.activePair.End)
}
f.toolCallBuf = ""
} else if matchLen := hasSuffixPrefixOf(f.buf, f.activeEndTag); matchLen > 0 {
safe := f.buf[:len(f.buf)-matchLen]
f.toolCallBuf += safe
f.buf = f.buf[len(f.buf)-matchLen:]
break
} else {
f.toolCallBuf += f.buf
f.buf = ""
break
}
}
}
}
func (f *StreamToolCallFilter) Flush(onContent func(string), onToolCall func(ToolCall)) {
if f.inToolCall && len(f.toolCallBuf) > 0 {
if tcs, ok := parseMultipleToolCalls(f.toolCallBuf); ok && len(tcs) > 0 {
for _, tc := range tcs {
idxCopy := f.toolIndex
tc.Index = &idxCopy
f.toolIndex++
f.emittedCall = true
onToolCall(tc)
}
} else if tcs2, ok2 := parseXMLToolCall(f.activePair.Start + f.toolCallBuf + f.activePair.End); ok2 && len(tcs2) > 0 {
for _, tc := range tcs2 {
idxCopy := f.toolIndex
tc.Index = &idxCopy
f.toolIndex++
f.emittedCall = true
onToolCall(tc)
}
} else {
onContent(f.activePair.Start + f.toolCallBuf)
}
f.toolCallBuf = ""
}
if len(f.buf) > 0 {
onContent(f.buf)
f.buf = ""
}
}
// ---------------------------------------------------------------------------
// Universal Gradio space inspector & metadata discovery
// ---------------------------------------------------------------------------
type GradioParamInfo struct {
Label string `json:"label"`
ParameterName string `json:"parameter_name"`
ParameterDefault interface{} `json:"parameter_default,omitempty"`
Component string `json:"component"`
}
type GradioEndpointInfo struct {
Parameters []GradioParamInfo `json:"parameters"`
Returns []GradioParamInfo `json:"returns"`
APIVisibility string `json:"api_visibility"`
Description string `json:"description"`
}
type GradioAPIInfoResponse struct {
NamedEndpoints map[string]GradioEndpointInfo `json:"named_endpoints"`
UnnamedEndpoints map[string]GradioEndpointInfo `json:"unnamed_endpoints"`
}
type GradioComponent struct {
ID int `json:"id"`
Type string `json:"type"`
Props map[string]interface{} `json:"props"`
SkipAPI bool `json:"skip_api"`
}
type GradioDependencyTypes struct {
Generator bool `json:"generator"`
Cancel bool `json:"cancel"`
}
type GradioDependency struct {
ID int `json:"id"`
APIName interface{} `json:"api_name"`
Inputs []int `json:"inputs"`
Outputs []int `json:"outputs"`
Queue interface{} `json:"queue"`
Types GradioDependencyTypes `json:"types"`
APIVisibility string `json:"api_visibility"`
}
type GradioConfigResponse struct {
Version string `json:"version"`
APIPrefix string `json:"api_prefix"`
Mode string `json:"mode"`
Title string `json:"title"`
Components []GradioComponent `json:"components"`
Dependencies []GradioDependency `json:"dependencies"`
}
type HFSpaceCardData struct {
Title string `json:"title"`
ShortDescription string `json:"short_description"`
}
type HFSpaceInfoResponse struct {
ID string `json:"id"`
Models []string `json:"models"`
CardData HFSpaceCardData `json:"cardData"`
}
type SpaceParamMapping struct {
InputIndex int
ComponentID int
ParamType string // "message", "history", "system_prompt", "temperature", "max_tokens", "top_p", "state", "other"
DefaultValue interface{}
}
type SpaceDiscovery struct {
SpaceURL string
Title string
Models []string
PrimaryModel string
APIPrefix string // e.g. "/gradio_api" or ""
Endpoint string // e.g. "/chat_fn" or "/chat"
CleanEndpoint string // e.g. "chat_fn" or "chat"
Protocol string // "call", "queue", "predict"
TotalInputs int
ParamMappings []SpaceParamMapping
HistoryIndex int // -1 if none
MessageIndex int // index for user message text
SystemIndex int // -1 if none
DefaultSystemPrompt string // default space system prompt if present
TempIndex int // -1 if none
MaxTokensIndex int // -1 if none
TopPIndex int // -1 if none
ThinkLevelIndex int // -1 if none
FunctionsJSONIndex int // -1 if none
PreservedThinkingIndex int // -1 if none
IsHunyuan3 bool
HistoryFormat string // "messages", "pairs", "none"
LastDiscovered time.Time
}
func (d *SpaceDiscovery) GetModelList() []ModelItem {
now := time.Now().Unix()
var items []ModelItem
seen := make(map[string]bool)
for _, m := range d.Models {
if m != "" && !seen[m] {
seen[m] = true
items = append(items, ModelItem{
ID: m,
Object: "model",
Created: now,
OwnedBy: "gradio",
})
}
}
if d.PrimaryModel != "" && !seen[d.PrimaryModel] {
seen[d.PrimaryModel] = true
items = append(items, ModelItem{
ID: d.PrimaryModel,
Object: "model",
Created: now,
OwnedBy: "gradio",
})
}
if len(items) == 0 {
items = append(items, ModelItem{
ID: "default",
Object: "model",
Created: now,
OwnedBy: "gradio",
})
}
return items
}
func NewDefaultSpaceDiscovery(spaceURL string) *SpaceDiscovery {
cleanURL := strings.TrimRight(spaceURL, "/")
if cleanURL != "" && !strings.HasPrefix(cleanURL, "http://") && !strings.HasPrefix(cleanURL, "https://") {
cleanURL = "https://" + cleanURL
}
return &SpaceDiscovery{
SpaceURL: cleanURL,
APIPrefix: "/gradio_api",
Endpoint: "/chat_fn",
CleanEndpoint: "chat_fn",
Protocol: "call",
TotalInputs: 1,
HistoryIndex: -1,
MessageIndex: 0,
SystemIndex: -1,
DefaultSystemPrompt: "",
TempIndex: -1,
MaxTokensIndex: -1,
TopPIndex: -1,
ThinkLevelIndex: -1,
FunctionsJSONIndex: -1,
PreservedThinkingIndex: -1,
HistoryFormat: "messages",
LastDiscovered: time.Now(),
}
}
// InspectSpace queries Gradio's /gradio_api/info, /config, and HuggingFace Space APIs
// to build an adaptive schema mapping for any Gradio space.
func InspectSpace(client *http.Client, rawURL, userAgent string) (*SpaceDiscovery, error) {
cleanURL := strings.TrimRight(rawURL, "/")
if !strings.HasPrefix(cleanURL, "http://") && !strings.HasPrefix(cleanURL, "https://") {
cleanURL = "https://" + cleanURL
}
discovery := NewDefaultSpaceDiscovery(cleanURL)
// 1. Try fetching /gradio_api/info or /info
var infoResp GradioAPIInfoResponse
infoFetched := false
for _, path := range []string{"/gradio_api/info", "/info"} {
infoURL := cleanURL + path
req, err := http.NewRequest("GET", infoURL, nil)
if err == nil {
req.Header.Set("User-Agent", userAgent)
resp, err := client.Do(req)
if err == nil && resp.StatusCode == http.StatusOK {
if json.NewDecoder(resp.Body).Decode(&infoResp) == nil {
infoFetched = true
if path == "/gradio_api/info" {
discovery.APIPrefix = "/gradio_api"
} else {
discovery.APIPrefix = ""
}
}
resp.Body.Close()
if infoFetched {
break
}
} else if resp != nil {
resp.Body.Close()
}
}
}
// 2. Try fetching /config
var configResp GradioConfigResponse
configFetched := false
for _, path := range []string{"/config", "/gradio_api/config"} {
cfgURL := cleanURL + path
req, err := http.NewRequest("GET", cfgURL, nil)
if err == nil {
req.Header.Set("User-Agent", userAgent)
resp, err := client.Do(req)
if err == nil && resp.StatusCode == http.StatusOK {
if json.NewDecoder(resp.Body).Decode(&configResp) == nil {
configFetched = true
if configResp.APIPrefix != "" {
discovery.APIPrefix = configResp.APIPrefix
}
if configResp.Title != "" {
discovery.Title = configResp.Title
}
}
resp.Body.Close()
if configFetched {
break
}
} else if resp != nil {
resp.Body.Close()
}
}
}
// 3. Inspect Hugging Face Space Metadata if hosted on HF
parsedURL, _ := url.Parse(cleanURL)
if parsedURL != nil && (strings.HasSuffix(parsedURL.Host, ".hf.space") || strings.Contains(parsedURL.Host, "huggingface.co")) {
subdomain := strings.TrimSuffix(parsedURL.Host, ".hf.space")
var owner, name string
dashIdx := strings.Index(subdomain, "-")
if dashIdx != -1 {
owner = subdomain[:dashIdx]
name = subdomain[dashIdx+1:]
}
if owner != "" && name != "" {
hfAPIURL := fmt.Sprintf("https://huggingface.co/api/spaces/%s/%s", owner, name)
req, err := http.NewRequest("GET", hfAPIURL, nil)
if err == nil {
req.Header.Set("User-Agent", userAgent)
resp, err := client.Do(req)
if err == nil && resp.StatusCode == http.StatusOK {
var hfResp HFSpaceInfoResponse
if json.NewDecoder(resp.Body).Decode(&hfResp) == nil {
for _, m := range hfResp.Models {
discovery.Models = append(discovery.Models, m)
cleanM := strings.TrimPrefix(m, "openai/")
cleanM = strings.TrimPrefix(cleanM, "models/")
if cleanM != m {
discovery.Models = append(discovery.Models, cleanM)
}
}
if hfResp.CardData.Title != "" && discovery.Title == "" {
discovery.Title = hfResp.CardData.Title
}
if len(discovery.Models) > 0 {
discovery.PrimaryModel = discovery.Models[0]
}
}
resp.Body.Close()
} else if resp != nil {
resp.Body.Close()
}
}
}
}
// Fallback model names if not discovered
if len(discovery.Models) == 0 {
if parsedURL != nil && strings.HasSuffix(parsedURL.Host, ".hf.space") {
sub := strings.TrimSuffix(parsedURL.Host, ".hf.space")
parts := strings.Split(sub, "-")
if len(parts) > 1 {
cleanModel := strings.Join(parts[1:], "-")
discovery.Models = append(discovery.Models, cleanModel)
discovery.PrimaryModel = cleanModel
}
}
}
if discovery.PrimaryModel == "" {
if len(discovery.Models) > 0 {
discovery.PrimaryModel = discovery.Models[0]
} else {
discovery.PrimaryModel = "gradio-chat"
discovery.Models = append(discovery.Models, "gradio-chat")
}
}
// 4. Score and select the best chat endpoint
bestEndpoint := ""
bestScore := -1000
var bestEndpointInfo *GradioEndpointInfo
if infoFetched && len(infoResp.NamedEndpoints) > 0 {
for epName, epInfo := range infoResp.NamedEndpoints {
score := 0
lowerName := strings.ToLower(epName)
if strings.Contains(lowerName, "chat") {
score += 100
}
if strings.Contains(lowerName, "predict") || strings.Contains(lowerName, "respond") || strings.Contains(lowerName, "generate") {
score += 50
}
for _, p := range epInfo.Parameters {
pLower := strings.ToLower(p.ParameterName)
if strings.Contains(pLower, "message") || strings.Contains(pLower, "text") || strings.Contains(pLower, "prompt") {
score += 40
}
if strings.Contains(pLower, "history") || strings.Contains(pLower, "chat") {
score += 20
}
}
if score > bestScore {
bestScore = score
bestEndpoint = epName
epCopy := epInfo
bestEndpointInfo = &epCopy
}
}
}
if bestEndpoint != "" {
discovery.Endpoint = bestEndpoint
discovery.CleanEndpoint = strings.TrimPrefix(bestEndpoint, "/")
}
// 5. Correlate with config.dependencies to determine exact input count & state padding
compMap := make(map[int]GradioComponent)
if configFetched {
for _, comp := range configResp.Components {
compMap[comp.ID] = comp
}
var matchingDep *GradioDependency
cleanTarget := strings.TrimPrefix(discovery.Endpoint, "/")
for _, dep := range configResp.Dependencies {
depAPIName := ""
if s, ok := dep.APIName.(string); ok {
depAPIName = strings.TrimPrefix(s, "/")
}
if depAPIName == cleanTarget {
depCopy := dep
matchingDep = &depCopy
break
}
}
if matchingDep != nil {
discovery.TotalInputs = len(matchingDep.Inputs)
for idx, compID := range matchingDep.Inputs {
mapping := SpaceParamMapping{
InputIndex: idx,
ComponentID: compID,
ParamType: "other",
}
if comp, exists := compMap[compID]; exists {
cType := strings.ToLower(comp.Type)
switch cType {
case "textbox", "multimodaltextbox":
if discovery.MessageIndex == 0 && idx == 0 {
mapping.ParamType = "message"
} else if discovery.SystemIndex == -1 {
mapping.ParamType = "system_prompt"
discovery.SystemIndex = idx
if comp.Props != nil {
if val, ok := comp.Props["value"].(string); ok && strings.TrimSpace(val) != "" {
discovery.DefaultSystemPrompt = strings.TrimSpace(val)
}
}
}
case "state":
mapping.ParamType = "state"
if idx == 1 && len(matchingDep.Inputs) == 2 {
// Standard Gradio ChatInterface: [textbox, state]
// Component 13 is state
}
case "slider", "number":
label := ""
if comp.Props != nil {
if l, ok := comp.Props["label"].(string); ok {
label = strings.ToLower(l)
}
}
if strings.Contains(label, "temp") {
mapping.ParamType = "temperature"
discovery.TempIndex = idx
} else if strings.Contains(label, "max") || strings.Contains(label, "token") {
mapping.ParamType = "max_tokens"
discovery.MaxTokensIndex = idx
} else if strings.Contains(label, "top_p") {
mapping.ParamType = "top_p"
discovery.TopPIndex = idx
}
}
}
discovery.ParamMappings = append(discovery.ParamMappings, mapping)
}
} else if bestEndpointInfo != nil {
discovery.TotalInputs = len(bestEndpointInfo.Parameters)
}
}
// Check parameters in bestEndpointInfo for input indices and history support
if bestEndpointInfo != nil {
if discovery.TotalInputs < len(bestEndpointInfo.Parameters) {
discovery.TotalInputs = len(bestEndpointInfo.Parameters)
}
for idx, p := range bestEndpointInfo.Parameters {
pName := strings.ToLower(p.ParameterName)
if strings.Contains(pName, "system") {
discovery.SystemIndex = idx
if p.ParameterDefault != nil && discovery.DefaultSystemPrompt == "" {
if defStr, ok := p.ParameterDefault.(string); ok && strings.TrimSpace(defStr) != "" {
discovery.DefaultSystemPrompt = strings.TrimSpace(defStr)
}
}
} else if strings.Contains(pName, "history") || strings.Contains(pName, "chat") {
discovery.HistoryIndex = idx
} else if strings.Contains(pName, "message") || (strings.Contains(pName, "prompt") && !strings.Contains(pName, "system")) || strings.Contains(pName, "text") {
discovery.MessageIndex = idx
} else if strings.Contains(pName, "think_level") || strings.Contains(pName, "thinking_level") {
discovery.ThinkLevelIndex = idx
} else if strings.Contains(pName, "functions") || strings.Contains(pName, "tools") {
discovery.FunctionsJSONIndex = idx
} else if strings.Contains(pName, "preserved") {
discovery.PreservedThinkingIndex = idx
} else if strings.Contains(pName, "temp") {
discovery.TempIndex = idx
} else if strings.Contains(pName, "token") {
discovery.MaxTokensIndex = idx
} else if strings.Contains(pName, "top_p") {
discovery.TopPIndex = idx
}
}
}
if discovery.FunctionsJSONIndex != -1 || discovery.ThinkLevelIndex != -1 || strings.Contains(cleanURL, "hy3") || strings.Contains(cleanURL, "hunyuan") {
discovery.IsHunyuan3 = true
discovery.Models = append(discovery.Models, "hy3", "hunyuan3", "tencent/Hy3")
if discovery.PrimaryModel == "gradio-chat" || discovery.PrimaryModel == "" {
discovery.PrimaryModel = "hy3"
}
}
// Ensure total inputs is at least 1
if discovery.TotalInputs < 1 {
discovery.TotalInputs = 1
}
return discovery, nil
}
// ---------------------------------------------------------------------------
// Universal Gradio gateway engine
// ---------------------------------------------------------------------------
type GradioJoinResponse struct {
EventID string `json:"event_id"`
}
type GradioGateway struct {
mu sync.RWMutex
defaultURL string
client *http.Client
discoveries map[string]*SpaceDiscovery
proxyURL string
}
func NewGradioGateway(defaultSpaceURL, proxyURL string, timeout time.Duration) *GradioGateway {
cleanDefault := strings.TrimRight(defaultSpaceURL, "/")
if !strings.HasPrefix(cleanDefault, "http://") && !strings.HasPrefix(cleanDefault, "https://") {
cleanDefault = "https://" + cleanDefault
}
transport := &http.Transport{
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
if proxyURL != "" {
return DialSOCKS5(ctx, proxyURL, addr)
}
var d net.Dialer
return d.DialContext(ctx, network, addr)
},
MaxIdleConns: 100,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 15 * time.Second,
}
gw := &GradioGateway{
defaultURL: cleanDefault,
client: &http.Client{Transport: transport, Timeout: timeout},
discoveries: make(map[string]*SpaceDiscovery),
proxyURL: proxyURL,
}
// Pre-discover the default space
disc, err := InspectSpace(gw.client, cleanDefault, DefaultUserAgent)
if err == nil && disc != nil {
gw.discoveries[cleanDefault] = disc
}
return gw
}
func (g *GradioGateway) GetDiscovery(spaceURL, userAgent string) *SpaceDiscovery {
target := spaceURL
if target == "" {
target = g.defaultURL
}
cleanTarget := strings.TrimRight(target, "/")
g.mu.RLock()
disc, exists := g.discoveries[cleanTarget]
g.mu.RUnlock()
if exists && disc != nil && time.Since(disc.LastDiscovered) < 30*time.Minute {
return disc
}
g.mu.Lock()
defer g.mu.Unlock()
// Double-check under lock
disc, exists = g.discoveries[cleanTarget]
if exists && disc != nil && time.Since(disc.LastDiscovered) < 30*time.Minute {
return disc
}
newDisc, err := InspectSpace(g.client, cleanTarget, userAgent)
if err == nil && newDisc != nil {
g.discoveries[cleanTarget] = newDisc
return newDisc
}
if disc != nil {
return disc
}
// Fallback discovery
fallback := NewDefaultSpaceDiscovery(cleanTarget)
fallback.TotalInputs = 2
fallback.PrimaryModel = "gradio-chat"
fallback.Models = []string{"gradio-chat"}
g.discoveries[cleanTarget] = fallback
return fallback
}
// BuildGradioPayload packages OpenAI messages and parameters into the target Gradio input array.
func (g *GradioGateway) BuildGradioPayload(disc *SpaceDiscovery, req ChatCompletionRequest) ([]interface{}, error) {
var transformed []ChatMessage
var toolInstruction string
if disc.IsHunyuan3 && disc.FunctionsJSONIndex != -1 {
for _, msg := range req.Messages {
transformed = append(transformed, ChatMessage{
Role: msg.Role,
Content: msg.GetContentString(),
ReasoningContent: msg.ReasoningContent,
ToolCalls: msg.ToolCalls,
ToolCallID: msg.ToolCallID,
Name: msg.Name,
})
}
} else {
// When no native tool calling support is detected, augment system prompt
transformed, toolInstruction, _ = TransformMessages(req)
}
hasClientSystem := false
for _, m := range req.Messages {
if m.Role == "system" {
hasClientSystem = true
break
}
}
var systemPromptStr string
var historyArray []map[string]interface{}
var lastUserMessage string
var nonSystem []ChatMessage
for _, m := range transformed {
cStr := m.GetContentString()
if m.Role == "system" {
if systemPromptStr == "" {
systemPromptStr = cStr
} else {
systemPromptStr += "\n\n" + cStr
}
} else {
nonSystem = append(nonSystem, m)
}
}
// If the space has a DefaultSystemPrompt and client provided no system message,
// retain and augment the default system prompt:
if !hasClientSystem && disc.DefaultSystemPrompt != "" {
if systemPromptStr != "" {
systemPromptStr = disc.DefaultSystemPrompt + "\n\n" + systemPromptStr
} else {
systemPromptStr = disc.DefaultSystemPrompt
}
}
// If the space has NO native system prompt input (disc.SystemIndex == -1),
// but we have system instructions (from system message or tool instructions):
if disc.SystemIndex == -1 && systemPromptStr != "" && len(nonSystem) > 0 {
// If the space supports conversation history, prepend system instructions to the first turn
if disc.HistoryIndex != -1 {
nonSystem[0].Content = systemPromptStr + "\n\n" + nonSystem[0].GetContentString()
}
}
if len(nonSystem) > 0 {
for i := 0; i < len(nonSystem)-1; i++ {
m := nonSystem[i]
cStr := m.GetContentString()
item := map[string]interface{}{"role": m.Role}
switch m.Role {
case "assistant":
if cStr != "" {
item["content"] = cStr
} else {
item["content"] = nil
}
if m.ReasoningContent != "" {
item["reasoning_content"] = m.ReasoningContent
}
if len(m.ToolCalls) > 0 {
item["tool_calls"] = m.ToolCalls
}
case "tool", "function":
item["role"] = "tool"
item["content"] = cStr
toolID := m.ToolCallID
if toolID == "" {
toolID = m.Name
}
if toolID != "" {
item["tool_call_id"] = toolID
}
if m.Name != "" {
item["name"] = m.Name
}
default:
item["content"] = cStr
}
historyArray = append(historyArray, item)
}
lastMsg := nonSystem[len(nonSystem)-1]
lastContent := lastMsg.GetContentString()
if lastMsg.Role == "tool" || lastMsg.Role == "function" {
toolName := lastMsg.Name
if toolName == "" {
toolName = lastMsg.ToolCallID
}
if disc.IsHunyuan3 {
toolItem := map[string]interface{}{
"role": "tool",
"content": lastContent,
}
if lastMsg.ToolCallID != "" {
toolItem["tool_call_id"] = lastMsg.ToolCallID
} else if toolName != "" {
toolItem["tool_call_id"] = toolName
}
if lastMsg.Name != "" {
toolItem["name"] = lastMsg.Name
}
historyArray = append(historyArray, toolItem)
lastUserMessage = "Please proceed based on the tool results."
} else {
if toolName != "" {
lastUserMessage = fmt.Sprintf("Tool result for %s: %s", toolName, lastContent)
} else {
lastUserMessage = lastContent
}
}
} else {
lastUserMessage = lastContent
}
} else if systemPromptStr != "" {
lastUserMessage = systemPromptStr
}
var promptMessageText string
if disc.HistoryIndex != -1 {
if disc.SystemIndex == -1 && len(nonSystem) > 1 && toolInstruction != "" {
promptMessageText = fmt.Sprintf("[System Directive: Tool calling mode active. If relevant, output a block.]\n\n%s", lastUserMessage)
} else {
promptMessageText = lastUserMessage
}
} else {
// Single message space: compose multi-turn history into the prompt
if len(nonSystem) <= 1 {
if systemPromptStr != "" && len(nonSystem) == 1 {
promptMessageText = systemPromptStr + "\n\n" + lastUserMessage
} else if systemPromptStr != "" {
promptMessageText = systemPromptStr
} else {
promptMessageText = lastUserMessage
}
} else {
var sb strings.Builder
if systemPromptStr != "" {
sb.WriteString("# Instructions\n" + systemPromptStr + "\n\n")
}
sb.WriteString("# Conversation History\n")
for i := 0; i < len(nonSystem)-1; i++ {
m := nonSystem[i]
roleLabel := "User"
if m.Role == "assistant" {
roleLabel = "Assistant"
}
sb.WriteString(fmt.Sprintf("%s: %s\n\n", roleLabel, m.GetContentString()))
}
lastRoleLabel := "User"
if len(nonSystem) > 0 && nonSystem[len(nonSystem)-1].Role == "assistant" {
lastRoleLabel = "Assistant"
}
sb.WriteString(fmt.Sprintf("# Current Request\n%s: %s", lastRoleLabel, lastUserMessage))
promptMessageText = sb.String()
}
}
// Allocate input array matching TotalInputs
totalInputs := disc.TotalInputs
if totalInputs < 1 {
totalInputs = 1
}
data := make([]interface{}, totalInputs)
// Populate mapped fields
msgIdx := disc.MessageIndex
if msgIdx >= 0 && msgIdx < len(data) {
data[msgIdx] = promptMessageText
}
if disc.HistoryIndex >= 0 && disc.HistoryIndex < len(data) {
if disc.HistoryFormat == "pairs" {
var pairs [][]string
for i := 0; i < len(historyArray); i += 2 {
u := ""
a := ""
if i < len(historyArray) {
u, _ = historyArray[i]["content"].(string)
}
if i+1 < len(historyArray) {
a, _ = historyArray[i+1]["content"].(string)
}
pairs = append(pairs, []string{u, a})
}
data[disc.HistoryIndex] = pairs
} else {
data[disc.HistoryIndex] = historyArray
}
}
if disc.SystemIndex >= 0 && disc.SystemIndex < len(data) {
data[disc.SystemIndex] = systemPromptStr
}
if disc.ThinkLevelIndex >= 0 && disc.ThinkLevelIndex < len(data) {
thinkLevel := "high"
if req.ReasoningEffort != "" {
effort := strings.ToLower(req.ReasoningEffort)
switch effort {
case "none", "off", "no_think", "0":
thinkLevel = "no_think"
case "low", "1":
thinkLevel = "low"
case "medium", "high", "2", "3":
thinkLevel = "high"
default:
thinkLevel = effort
}
}
data[disc.ThinkLevelIndex] = thinkLevel
}
if disc.TempIndex >= 0 && disc.TempIndex < len(data) {
if req.Temperature != nil {
data[disc.TempIndex] = *req.Temperature
} else if disc.IsHunyuan3 {
data[disc.TempIndex] = nil
} else {
data[disc.TempIndex] = 0.7
}
}
if disc.MaxTokensIndex >= 0 && disc.MaxTokensIndex < len(data) {
data[disc.MaxTokensIndex] = ResolveMaxTokens(req)
}
if disc.TopPIndex >= 0 && disc.TopPIndex < len(data) {
if req.TopP != nil {
data[disc.TopPIndex] = *req.TopP
} else if disc.IsHunyuan3 {
data[disc.TopPIndex] = 0
} else {
data[disc.TopPIndex] = 1.0
}
}
if disc.FunctionsJSONIndex >= 0 && disc.FunctionsJSONIndex < len(data) {
functionsJSONStr := ""
if len(req.Tools) > 0 {
b, err := json.Marshal(req.Tools)
if err == nil {
functionsJSONStr = string(b)
}
}
data[disc.FunctionsJSONIndex] = functionsJSONStr
}
return data, nil
}
// GradioOutputFrame holds parsed elements from a Gradio SSE output chunk
type GradioOutputFrame struct {
Content string
Reasoning string
ToolCalls []ToolCall
OK bool
}
// ParseGradioStreamOutput extracts structured content, reasoning, and tool calls from Gradio output
func ParseGradioStreamOutput(rawJSON string) GradioOutputFrame {
var frame GradioOutputFrame
var val interface{}
if err := json.Unmarshal([]byte(rawJSON), &val); err != nil {
return frame
}
switch v := val.(type) {
case string:
frame.Content = v
frame.OK = true
return frame
case []interface{}:
if len(v) == 0 {
return frame
}
// Check if v[0] is an inner slice (e.g. Hy3: [[content, reasoning, tool_calls, history]])
if inner, ok := v[0].([]interface{}); ok {
if len(inner) >= 2 {
s0, ok0 := inner[0].(string)
s1, ok1 := inner[1].(string)
if ok0 && ok1 {
frame.Content = s0
frame.Reasoning = s1
if len(inner) >= 3 {
if tcSlice, ok := inner[2].([]interface{}); ok && len(tcSlice) > 0 {
b, err := json.Marshal(tcSlice)
if err == nil {
var tcs []ToolCall
if err := json.Unmarshal(b, &tcs); err == nil {
frame.ToolCalls = tcs
}
}
}
}
frame.OK = true
return frame
}
// Check if inner is a chat pair: ["user msg", "assistant msg"]
if len(inner) == 2 {
if aStr, ok := inner[1].(string); ok {
frame.Content = aStr
frame.OK = true
return frame
}
}
}
// Check if inner is a list of chat message maps: [{"role":..., "content":...}, ...]
// or list of pairs: [["u", "a"], ...]
if len(inner) > 0 {
lastItem := inner[len(inner)-1]
if m, ok := lastItem.(map[string]interface{}); ok {
if c, ok := m["content"].(string); ok {
frame.Content = c
frame.OK = true
}
if r, ok := m["reasoning_content"].(string); ok {
frame.Reasoning = r
}
if tcsRaw, ok := m["tool_calls"].([]interface{}); ok && len(tcsRaw) > 0 {
b, err := json.Marshal(tcsRaw)
if err == nil {
var tcs []ToolCall
if err := json.Unmarshal(b, &tcs); err == nil {
frame.ToolCalls = tcs
}
}
}
if frame.OK {
return frame
}
} else if pair, ok := lastItem.([]interface{}); ok && len(pair) >= 2 {
if aStr, ok := pair[1].(string); ok {
frame.Content = aStr
frame.OK = true
return frame
}
}
}
}
// Check if v[0] is string (standard single output e.g. ["content", null])
if s, ok := v[0].(string); ok {
frame.Content = s
frame.OK = true
return frame
}
// Check if v is a flat list of messages: [{"role": "assistant", ...}]
lastItem := v[len(v)-1]
if m, ok := lastItem.(map[string]interface{}); ok {
if c, ok := m["content"].(string); ok {
frame.Content = c
frame.OK = true
}
if r, ok := m["reasoning_content"].(string); ok {
frame.Reasoning = r
}
if tcsRaw, ok := m["tool_calls"].([]interface{}); ok && len(tcsRaw) > 0 {
b, err := json.Marshal(tcsRaw)
if err == nil {
var tcs []ToolCall
if err := json.Unmarshal(b, &tcs); err == nil {
frame.ToolCalls = tcs
}
}
}
if frame.OK {
return frame
}
}
case map[string]interface{}:
for _, key := range []string{"text", "content", "response", "data", "value"} {
if s, ok := v[key].(string); ok {
frame.Content = s
frame.OK = true
break
}
}
if r, ok := v["reasoning_content"].(string); ok {
frame.Reasoning = r
} else if r, ok := v["reasoning"].(string); ok {
frame.Reasoning = r
}
if tcsRaw, ok := v["tool_calls"].([]interface{}); ok && len(tcsRaw) > 0 {
b, err := json.Marshal(tcsRaw)
if err == nil {
var tcs []ToolCall
if err := json.Unmarshal(b, &tcs); err == nil {
frame.ToolCalls = tcs
}
}
}
if frame.OK {
return frame
}
}
return frame
}
// ExtractTextFromGradioOutput extracts the assistant text string from Gradio output chunks (compatibility wrapper)
func ExtractTextFromGradioOutput(rawJSON string) (string, bool) {
frame := ParseGradioStreamOutput(rawJSON)
if frame.OK {
return frame.Content, true
}
return "", false
}
// ExecuteChatCompletion handles both streaming and non-streaming requests.
func (g *GradioGateway) ExecuteChatCompletion(w http.ResponseWriter, r *http.Request, req ChatCompletionRequest) error {
effUA := EffectiveUserAgent(r)
// Target space selection: check request headers or fallback to default space
spaceURL := g.defaultURL
if hdr := r.Header.Get("X-Gradio-Space"); hdr != "" {
spaceURL = hdr
} else if hdr := r.Header.Get("X-Space-URL"); hdr != "" {
spaceURL = hdr
}
disc := g.GetDiscovery(spaceURL, effUA)
modelName := req.Model
if modelName == "" {
modelName = disc.PrimaryModel
}
gradioData, err := g.BuildGradioPayload(disc, req)
if err != nil {
return fmt.Errorf("failed to build Gradio payload: %w", err)
}
payloadMap := map[string]interface{}{"data": gradioData}
jsonPayload, err := json.Marshal(payloadMap)
if err != nil {
return fmt.Errorf("failed to encode request: %w", err)
}
completionID := "chatcmpl-" + GenerateUUID()
createdTime := time.Now().Unix()
// 1. Submit to /call/{endpoint}
callURL := fmt.Sprintf("%s%s/call/%s", disc.SpaceURL, disc.APIPrefix, disc.CleanEndpoint)
makeCallReq := func() (*http.Request, error) {
r, err := http.NewRequest("POST", callURL, bytes.NewBuffer(jsonPayload))
if err != nil {
return nil, err
}
r.Header.Set("Content-Type", "application/json")
r.Header.Set("User-Agent", effUA)
return r, nil
}
resp, err := DoWithFibonacciRetry(g.client, makeCallReq, 5)
if err != nil {
if fg, ok := extractFailedGeneration(err.Error()); ok {
if tcs, _, has := DetectToolCalls(fg); has && len(tcs) > 0 {
if !req.Stream {
WriteCompletionResponse(w, completionID, createdTime, modelName, FinalOutput{
ToolCalls: tcs,
FinishReason: "tool_calls",
})
return nil
}
flusher, _ := w.(http.Flusher)
streamer := NewStreamer(w, flusher, completionID, createdTime, modelName)
for i, tc := range tcs {
iCopy := i
tc.Index = &iCopy
streamer.ToolCallDelta(tc)
}
streamer.Finish("tool_calls")
streamer.Done()
return nil
}
}
// If call failed, try without APIPrefix or try /call/v2
altCallURL := fmt.Sprintf("%s/call/%s", disc.SpaceURL, disc.CleanEndpoint)
makeAltReq := func() (*http.Request, error) {
r, err := http.NewRequest("POST", altCallURL, bytes.NewBuffer(jsonPayload))
if err != nil {
return nil, err
}
r.Header.Set("Content-Type", "application/json")
r.Header.Set("User-Agent", effUA)
return r, nil
}
resp, err = DoWithFibonacciRetry(g.client, makeAltReq, 3)
if err != nil {
if fg, ok := extractFailedGeneration(err.Error()); ok {
if tcs, _, has := DetectToolCalls(fg); has && len(tcs) > 0 {
if !req.Stream {
WriteCompletionResponse(w, completionID, createdTime, modelName, FinalOutput{
ToolCalls: tcs,
FinishReason: "tool_calls",
})
return nil
}
flusher, _ := w.(http.Flusher)
streamer := NewStreamer(w, flusher, completionID, createdTime, modelName)
for i, tc := range tcs {
iCopy := i
tc.Index = &iCopy
streamer.ToolCallDelta(tc)
}
streamer.Finish("tool_calls")
streamer.Done()
return nil
}
}
return fmt.Errorf("upstream Gradio call error: %w", err)
}
}
defer resp.Body.Close()
var joinRes GradioJoinResponse
if err := json.NewDecoder(resp.Body).Decode(&joinRes); err != nil || joinRes.EventID == "" {
return fmt.Errorf("failed to parse Gradio event ID from response")
}
// 2. Connect to Gradio SSE EventStream
streamURL := fmt.Sprintf("%s%s/call/%s/%s", disc.SpaceURL, disc.APIPrefix, disc.CleanEndpoint, joinRes.EventID)
makeStreamReq := func() (*http.Request, error) {
r, err := http.NewRequest("GET", streamURL, nil)
if err != nil {
return nil, err
}
r.Header.Set("Accept", "text/event-stream")
r.Header.Set("User-Agent", effUA)
return r, nil
}
streamResp, err := DoWithFibonacciRetry(g.client, makeStreamReq, 5)
if err != nil {
if fg, ok := extractFailedGeneration(err.Error()); ok {
if tcs, _, has := DetectToolCalls(fg); has && len(tcs) > 0 {
if !req.Stream {
WriteCompletionResponse(w, completionID, createdTime, modelName, FinalOutput{
ToolCalls: tcs,
FinishReason: "tool_calls",
})
return nil
}
flusher, _ := w.(http.Flusher)
streamer := NewStreamer(w, flusher, completionID, createdTime, modelName)
for i, tc := range tcs {
iCopy := i
tc.Index = &iCopy
streamer.ToolCallDelta(tc)
}
streamer.Finish("tool_calls")
streamer.Done()
return nil
}
}
return fmt.Errorf("upstream Gradio stream error: %w", err)
}
defer streamResp.Body.Close()
// 3. Handle Non-Streaming vs Streaming
if !req.Stream {
reader := bufio.NewReader(streamResp.Body)
var latestFrame GradioOutputFrame
currentEvent := ""
for {
line, err := reader.ReadString('\n')
if err != nil {
break
}
line = strings.TrimRight(line, "\r\n")
if strings.HasPrefix(line, "event: ") {
currentEvent = strings.TrimPrefix(line, "event: ")
continue
}
if strings.HasPrefix(line, "data: ") {
dataStr := strings.TrimPrefix(line, "data: ")
if currentEvent == "error" {
if fg, ok := extractFailedGeneration(dataStr); ok {
if tcs, _, has := DetectToolCalls(fg); has && len(tcs) > 0 {
latestFrame = GradioOutputFrame{
ToolCalls: tcs,
OK: true,
}
break
}
}
errMsg := extractGradioErrorMessage(dataStr)
log.Printf("Upstream Gradio error: %s", errMsg)
return fmt.Errorf("upstream Gradio error: %s", errMsg)
}
if frame := ParseGradioStreamOutput(dataStr); frame.OK {
latestFrame = frame
}
if currentEvent == "complete" {
break
}
}
}
if !latestFrame.OK {
return fmt.Errorf("upstream Gradio space returned empty or unparseable response")
}
cleanText := latestFrame.Content
reasoning := latestFrame.Reasoning
toolCalls := latestFrame.ToolCalls
hasTools := len(toolCalls) > 0
if reasoning == "" {
cleanText, reasoning = ExtractThinking(cleanText)
}
if !hasTools {
toolCalls, cleanText, hasTools = DetectToolCalls(cleanText)
}
finishReason := "stop"
var finalContent interface{} = cleanText
if hasTools && len(toolCalls) > 0 {
finishReason = "tool_calls"
if strings.TrimSpace(cleanText) == "" {
finalContent = nil
}
}
WriteCompletionResponse(w, completionID, createdTime, modelName, FinalOutput{
Content: finalContent,
ReasoningContent: reasoning,
ToolCalls: toolCalls,
FinishReason: finishReason,
})
return nil
}
// 4. Streaming Mode
flusher, _ := w.(http.Flusher)
streamer := NewStreamer(w, flusher, completionID, createdTime, modelName)
thinkFilter := NewStreamThinkingFilter()
toolFilter := NewStreamToolCallFilter()
reader := bufio.NewReader(streamResp.Body)
var prevContent string
var prevReasoning string
prevToolArgs := make(map[int]string)
nativeReasoningSeen := false
nativeToolCallsSeen := false
currentEvent := ""
var streamErr error
for {
line, err := reader.ReadString('\n')
if err != nil {
break
}
line = strings.TrimRight(line, "\r\n")
if strings.HasPrefix(line, "event: ") {
currentEvent = strings.TrimPrefix(line, "event: ")
continue
}
if strings.HasPrefix(line, "data: ") {
dataStr := strings.TrimPrefix(line, "data: ")
if currentEvent == "error" {
if fg, ok := extractFailedGeneration(dataStr); ok {
if tcs, _, has := DetectToolCalls(fg); has && len(tcs) > 0 {
for i, tc := range tcs {
iCopy := i
tc.Index = &iCopy
streamer.ToolCallDelta(tc)
}
nativeToolCallsSeen = true
break
}
}
errMsg := extractGradioErrorMessage(dataStr)
log.Printf("Upstream Gradio error: %s", errMsg)
if !streamer.started {
return fmt.Errorf("upstream Gradio error: %s", errMsg)
}
streamer.Error(errMsg)
streamErr = fmt.Errorf("upstream Gradio error: %s", errMsg)
break
}
frame := ParseGradioStreamOutput(dataStr)
if frame.OK {
// 1. Native reasoning handling
if frame.Reasoning != "" || nativeReasoningSeen {
nativeReasoningSeen = true
var deltaReasoning string
if strings.HasPrefix(frame.Reasoning, prevReasoning) {
deltaReasoning = frame.Reasoning[len(prevReasoning):]
} else if prevReasoning == "" {
deltaReasoning = frame.Reasoning
} else {
deltaReasoning = frame.Reasoning
}
prevReasoning = frame.Reasoning
if deltaReasoning != "" {
streamer.Reasoning(deltaReasoning)
}
}
// 2. Native tool calls handling
if len(frame.ToolCalls) > 0 {
nativeToolCallsSeen = true
for idx, tc := range frame.ToolCalls {
prevArgs, started := prevToolArgs[idx]
currArgs := tc.Function.Arguments
idxCopy := idx
if !started {
tcDelta := ToolCall{
Index: &idxCopy,
ID: tc.ID,
Type: tc.Type,
Function: ToolCallFunction{
Name: tc.Function.Name,
Arguments: currArgs,
},
}
streamer.ToolCallDelta(tcDelta)
prevToolArgs[idx] = currArgs
} else if len(currArgs) > len(prevArgs) {
var argDelta string
if strings.HasPrefix(currArgs, prevArgs) {
argDelta = currArgs[len(prevArgs):]
} else {
argDelta = currArgs[len(prevArgs):]
}
if argDelta != "" {
tcDelta := ToolCall{
Index: &idxCopy,
Function: ToolCallFunction{
Arguments: argDelta,
},
}
streamer.ToolCallDelta(tcDelta)
}
prevToolArgs[idx] = currArgs
}
}
}
// 3. Content handling
currentText := frame.Content
var delta string
if strings.HasPrefix(currentText, prevContent) {
delta = currentText[len(prevContent):]
} else if prevContent == "" {
delta = currentText
} else {
delta = currentText
}
prevContent = currentText
if delta != "" {
if nativeReasoningSeen || nativeToolCallsSeen {
if nativeReasoningSeen && nativeToolCallsSeen {
streamer.Content(delta)
} else if nativeReasoningSeen {
toolFilter.Feed(delta, func(cleanChunk string) {
if cleanChunk != "" {
streamer.Content(cleanChunk)
}
}, func(tc ToolCall) {
streamer.ToolCallDelta(tc)
})
} else {
thinkFilter.Feed(delta, func(contentChunk string) {
if contentChunk != "" {
streamer.Content(contentChunk)
}
}, func(reasoningChunk string) {
if reasoningChunk != "" {
streamer.Reasoning(reasoningChunk)
}
})
}
} else {
thinkFilter.Feed(delta, func(contentChunk string) {
toolFilter.Feed(contentChunk, func(cleanChunk string) {
if cleanChunk != "" {
streamer.Content(cleanChunk)
}
}, func(tc ToolCall) {
streamer.ToolCallDelta(tc)
})
}, func(reasoningChunk string) {
if reasoningChunk != "" {
streamer.Reasoning(reasoningChunk)
}
})
}
}
}
if currentEvent == "complete" {
break
}
}
}
if streamErr != nil {
return nil
}
// Flush remaining tokens in filters if used
if !nativeReasoningSeen {
thinkFilter.Flush(func(contentChunk string) {
if !nativeToolCallsSeen {
toolFilter.Feed(contentChunk, func(cleanChunk string) {
if cleanChunk != "" {
streamer.Content(cleanChunk)
}
}, func(tc ToolCall) {
streamer.ToolCallDelta(tc)
})
} else if contentChunk != "" {
streamer.Content(contentChunk)
}
}, func(reasoningChunk string) {
if reasoningChunk != "" {
streamer.Reasoning(reasoningChunk)
}
})
}
if !nativeToolCallsSeen {
toolFilter.Flush(func(cleanChunk string) {
if cleanChunk != "" {
streamer.Content(cleanChunk)
}
}, func(tc ToolCall) {
streamer.ToolCallDelta(tc)
})
}
if !streamer.started {
return fmt.Errorf("upstream Gradio space closed stream without sending content")
}
if nativeToolCallsSeen || toolFilter.emittedCall {
streamer.Finish("tool_calls")
} else {
streamer.Finish("stop")
}
streamer.Done()
return nil
}
// ---------------------------------------------------------------------------
// HTTP routes & main server
// ---------------------------------------------------------------------------
func main() {
spaceFlag := flag.String("space", DefaultSpaceURL, "Target Gradio Space URL")
flag.StringVar(spaceFlag, "url", DefaultSpaceURL, "Alias for -space")
portFlag := flag.Int("port", 8080, "Gateway HTTP server port")
hostFlag := flag.String("host", "0.0.0.0", "Gateway HTTP server host")
socksFlag := flag.String("socks", "", "Optional SOCKS5 proxy URL (e.g. socks5://127.0.0.1:1080)")
flag.StringVar(socksFlag, "proxy", "", "Alias for -socks")
flag.StringVar(socksFlag, "socks5", "", "Alias for -socks")
uaFlag := flag.String("user-agent", "", "Custom User-Agent header")
flag.StringVar(uaFlag, "ua", "", "Alias for -user-agent")
timeoutFlag := flag.Int("timeout", 300, "Upstream timeout in seconds")
flag.Parse()
// Environment variable fallbacks
if envSpace := os.Getenv("GRADIO_SPACE_URL"); envSpace != "" && *spaceFlag == DefaultSpaceURL {
*spaceFlag = envSpace
}
if *socksFlag == "" {
for _, envName := range []string{"ALL_PROXY", "all_proxy", "SOCKS5_PROXY", "socks5_proxy", "SOCKS_PROXY", "socks_proxy"} {
if p := os.Getenv(envName); p != "" {
*socksFlag = p
break
}
}
}
if *uaFlag != "" {
ConfiguredUserAgent = *uaFlag
}
gateway := NewGradioGateway(*spaceFlag, *socksFlag, time.Duration(*timeoutFlag)*time.Second)
mux := http.NewServeMux()
// Health and Info
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
EnableCORS(w)
if r.Method == "OPTIONS" {
w.WriteHeader(http.StatusOK)
return
}
if r.URL.Path != "/" {
http.NotFound(w, r)
return
}
disc := gateway.GetDiscovery("", EffectiveUserAgent(r))
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]interface{}{
"status": "running",
"service": "gr2gw",
"space_url": disc.SpaceURL,
"title": disc.Title,
"endpoint": disc.Endpoint,
"primary_model": disc.PrimaryModel,
"models": disc.Models,
"total_inputs": disc.TotalInputs,
"history_format": disc.HistoryFormat,
})
})
// Models list
handleModels := func(w http.ResponseWriter, r *http.Request) {
EnableCORS(w)
if r.Method == "OPTIONS" {
w.WriteHeader(http.StatusOK)
return
}
spaceURL := gateway.defaultURL
if hdr := r.Header.Get("X-Gradio-Space"); hdr != "" {
spaceURL = hdr
} else if hdr := r.Header.Get("X-Space-URL"); hdr != "" {
spaceURL = hdr
}
disc := gateway.GetDiscovery(spaceURL, EffectiveUserAgent(r))
resp := ModelsResponse{
Object: "list",
Data: disc.GetModelList(),
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
}
mux.HandleFunc("/models", handleModels)
mux.HandleFunc("/v1/models", handleModels)
// Chat completions
handleCompletions := func(w http.ResponseWriter, r *http.Request) {
EnableCORS(w)
if r.Method == "OPTIONS" {
w.WriteHeader(http.StatusOK)
return
}
if r.Method != "POST" {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req ChatCompletionRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
http.Error(w, fmt.Sprintf("Invalid JSON request: %v", err), http.StatusBadRequest)
return
}
if err := gateway.ExecuteChatCompletion(w, r, req); err != nil {
log.Printf("Chat completion error: %v", err)
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusBadGateway)
json.NewEncoder(w).Encode(map[string]interface{}{
"error": map[string]interface{}{
"message": err.Error(),
"type": "upstream_error",
"code": http.StatusBadGateway,
},
})
return
}
}
mux.HandleFunc("/chat/completions", handleCompletions)
mux.HandleFunc("/v1/chat/completions", handleCompletions)
addr := fmt.Sprintf("%s:%d", *hostFlag, *portFlag)
log.Printf("gr2gw listening on %s (target space: %s)", addr, *spaceFlag)
if *socksFlag != "" {
log.Printf("Using SOCKS5 proxy: %s", *socksFlag)
}
server := &http.Server{
Addr: addr,
Handler: mux,
ReadTimeout: time.Duration(*timeoutFlag+30) * time.Second,
WriteTimeout: time.Duration(*timeoutFlag+30) * time.Second,
}
if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed {
log.Fatalf("Server failed: %v", err)
}
}