// groqqer: Pure Go OpenAI-compatible LLM gateway for the Groq Streamlit space
// Reverse engineers https://dromerosm-groq-chatbot.hf.space via direct WebSocket Protobuf wire format
// Created by Luxferre in 2026, released into the public domain with no warranties
package main
import (
"bufio"
"bytes"
"context"
"crypto/rand"
"crypto/tls"
"encoding/base64"
"encoding/binary"
"encoding/json"
"errors"
"flag"
"fmt"
"io"
"log"
"net"
"net/http"
"net/url"
"os"
"os/signal"
"regexp"
"strconv"
"strings"
"sync"
"syscall"
"time"
)
var (
DefaultUserAgent = "Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/133.0.0.0 Safari/537.36"
DefaultTargetURL = "https://dromerosm-groq-chatbot.hf.space"
DefaultModel = "qwen/qwen3.6-27b"
)
// ---------------------------------------------------------------------------
// 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"`
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"`
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"`
}
// ---------------------------------------------------------------------------
// UUID & Utility Functions
// ---------------------------------------------------------------------------
func GenerateUUID() string {
var b [16]byte
_, _ = rand.Read(b[:])
b[6] = (b[6] & 0x0f) | 0x40
b[8] = (b[8] & 0x3f) | 0x80
return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:16])
}
// ---------------------------------------------------------------------------
// 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{})
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])
}
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])
}
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)
}
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
}
// ---------------------------------------------------------------------------
// Pure Go Protobuf Wire Formatter & Decoder
// ---------------------------------------------------------------------------
func encodeVarint(val uint64) []byte {
var buf []byte
for val >= 0x80 {
buf = append(buf, byte(val|0x80))
val >>= 7
}
buf = append(buf, byte(val))
return buf
}
func encodeTag(fieldNum int, wireType int) []byte {
return encodeVarint(uint64((fieldNum << 3) | wireType))
}
func encodeLengthDelimited(fieldNum int, data []byte) []byte {
tag := encodeTag(fieldNum, 2)
length := encodeVarint(uint64(len(data)))
res := append(tag, length...)
return append(res, data...)
}
func encodeString(fieldNum int, str string) []byte {
return encodeLengthDelimited(fieldNum, []byte(str))
}
type ProtoField struct {
Tag int
WireType int
Varint uint64
Data []byte
}
func decodeProtoFields(data []byte) []ProtoField {
var fields []ProtoField
r := bytes.NewReader(data)
for r.Len() > 0 {
rawTag, err := binary.ReadUvarint(r)
if err != nil {
break
}
fieldNum := int(rawTag >> 3)
wireType := int(rawTag & 0x07)
pf := ProtoField{Tag: fieldNum, WireType: wireType}
switch wireType {
case 0:
v, err := binary.ReadUvarint(r)
if err != nil {
return fields
}
pf.Varint = v
case 1:
buf := make([]byte, 8)
if _, err := io.ReadFull(r, buf); err != nil {
return fields
}
pf.Data = buf
case 2:
length, err := binary.ReadUvarint(r)
if err != nil {
return fields
}
buf := make([]byte, length)
if _, err := io.ReadFull(r, buf); err != nil {
return fields
}
pf.Data = buf
case 5:
buf := make([]byte, 4)
if _, err := io.ReadFull(r, buf); err != nil {
return fields
}
pf.Data = buf
default:
return fields
}
fields = append(fields, pf)
}
return fields
}
type ParsedForwardMsg struct {
IsFinished bool
DeltaPath []uint64
Markdown string
ChatInputID string
SelectboxID string
Options []string
AlertText string
}
func parseForwardMsg(data []byte) ParsedForwardMsg {
var res ParsedForwardMsg
fields := decodeProtoFields(data)
for _, f := range fields {
if f.Tag == 6 && (f.Varint == 0 || f.Varint == 1 || f.Varint == 3 || f.Varint == 4 || f.Varint == 5) {
res.IsFinished = true
}
if f.Tag == 2 && f.WireType == 2 { // ForwardMsgMetadata
for _, mf := range decodeProtoFields(f.Data) {
if mf.Tag == 2 { // delta_path
if mf.WireType == 2 { // packed uint32
r := bytes.NewReader(mf.Data)
for r.Len() > 0 {
v, _ := binary.ReadUvarint(r)
res.DeltaPath = append(res.DeltaPath, v)
}
} else if mf.WireType == 0 {
res.DeltaPath = append(res.DeltaPath, mf.Varint)
}
}
}
}
if f.Tag == 5 && f.WireType == 2 { // Delta
dFields := decodeProtoFields(f.Data)
for _, df := range dFields {
if df.Tag == 3 && df.WireType == 2 { // Element new_element
eFields := decodeProtoFields(df.Data)
for _, ef := range eFields {
if ef.Tag == 29 && ef.WireType == 2 { // Markdown
mFields := decodeProtoFields(ef.Data)
for _, mf := range mFields {
if mf.Tag == 1 && mf.WireType == 2 {
res.Markdown = string(mf.Data)
}
}
} else if ef.Tag == 49 && ef.WireType == 2 { // ChatInput
cFields := decodeProtoFields(ef.Data)
for _, cf := range cFields {
if cf.Tag == 1 && cf.WireType == 2 {
res.ChatInputID = string(cf.Data)
}
}
} else if ef.Tag == 25 && ef.WireType == 2 { // Selectbox
sFields := decodeProtoFields(ef.Data)
for _, sf := range sFields {
if sf.Tag == 1 && sf.WireType == 2 {
res.SelectboxID = string(sf.Data)
} else if sf.Tag == 4 && sf.WireType == 2 {
res.Options = append(res.Options, string(sf.Data))
}
}
} else if ef.Tag == 30 && ef.WireType == 2 { // Alert
aFields := decodeProtoFields(ef.Data)
for _, af := range aFields {
if af.Tag == 1 && af.WireType == 2 {
res.AlertText = string(af.Data)
}
}
}
}
}
}
}
}
return res
}
// ---------------------------------------------------------------------------
// RFC 6455 Pure Go WebSocket Client
// ---------------------------------------------------------------------------
func sendWSBinaryFrame(conn io.Writer, payload []byte) error {
var header []byte
header = append(header, 0x82) // Binary frame (0x02) | FIN (0x80)
length := len(payload)
maskKey := make([]byte, 4)
_, _ = rand.Read(maskKey)
if length < 126 {
header = append(header, byte(length|0x80))
} else if length < 65536 {
header = append(header, 126|0x80)
var b [2]byte
binary.BigEndian.PutUint16(b[:], uint16(length))
header = append(header, b[:]...)
} else {
header = append(header, 127|0x80)
var b [8]byte
binary.BigEndian.PutUint64(b[:], uint64(length))
header = append(header, b[:]...)
}
header = append(header, maskKey...)
masked := make([]byte, length)
for i := 0; i < length; i++ {
masked[i] = payload[i] ^ maskKey[i%4]
}
_, err := conn.Write(append(header, masked...))
return err
}
func sendWSPongFrame(conn io.Writer, payload []byte) error {
var header []byte
header = append(header, 0x8a) // Pong (0x0A) | FIN (0x80)
length := len(payload)
maskKey := make([]byte, 4)
_, _ = rand.Read(maskKey)
header = append(header, byte(length|0x80))
header = append(header, maskKey...)
masked := make([]byte, length)
for i := 0; i < length; i++ {
masked[i] = payload[i] ^ maskKey[i%4]
}
_, err := conn.Write(append(header, masked...))
return err
}
func readWSFrame(conn net.Conn, reader *bufio.Reader, timeout time.Duration) ([]byte, byte, error) {
if timeout > 0 {
_ = conn.SetReadDeadline(time.Now().Add(timeout))
} else {
_ = conn.SetReadDeadline(time.Time{})
}
b1, err := reader.ReadByte()
if err != nil {
return nil, 0, err
}
b2, err := reader.ReadByte()
if err != nil {
return nil, 0, err
}
opcode := b1 & 0x0f
isMasked := (b2 & 0x80) != 0
length := int(b2 & 0x7f)
if length == 126 {
var ext uint16
if err := binary.Read(reader, binary.BigEndian, &ext); err != nil {
return nil, 0, err
}
length = int(ext)
} else if length == 127 {
var ext uint64
if err := binary.Read(reader, binary.BigEndian, &ext); err != nil {
return nil, 0, err
}
length = int(ext)
}
var maskKey [4]byte
if isMasked {
if _, err := io.ReadFull(reader, maskKey[:]); err != nil {
return nil, 0, err
}
}
payload := make([]byte, length)
if _, err := io.ReadFull(reader, payload); err != nil {
return nil, 0, err
}
if isMasked {
for i := 0; i < length; i++ {
payload[i] ^= maskKey[i%4]
}
}
return payload, opcode, nil
}
// ---------------------------------------------------------------------------
// Streamlit Session & Client
// ---------------------------------------------------------------------------
type StreamlitSession struct {
conn net.Conn
reader *bufio.Reader
chatInputID string
selectboxID string
options []string
activeModel string
closed bool
mu sync.Mutex
}
func (s *StreamlitSession) Close() {
s.mu.Lock()
defer s.mu.Unlock()
if !s.closed {
s.closed = true
if s.conn != nil {
_ = s.conn.Close()
}
}
}
func dialStreamlitWebSocket(ctx context.Context, targetURL, proxyURL, userAgent string) (net.Conn, *bufio.Reader, error) {
u, err := url.Parse(targetURL)
if err != nil {
return nil, nil, fmt.Errorf("invalid target url: %w", err)
}
host := u.Hostname()
port := u.Port()
isTLS := u.Scheme == "https" || u.Scheme == "wss"
if port == "" {
if isTLS {
port = "443"
} else {
port = "80"
}
}
targetAddr := net.JoinHostPort(host, port)
var rawConn net.Conn
if proxyURL != "" {
rawConn, err = DialSOCKS5(ctx, proxyURL, targetAddr)
} else {
var d net.Dialer
rawConn, err = d.DialContext(ctx, "tcp", targetAddr)
}
if err != nil {
return nil, nil, fmt.Errorf("failed to connect to %s: %w", targetAddr, err)
}
var conn net.Conn = rawConn
if isTLS {
tlsConfig := &tls.Config{
ServerName: host,
}
tlsConn := tls.Client(rawConn, tlsConfig)
if err := tlsConn.HandshakeContext(ctx); err != nil {
rawConn.Close()
return nil, nil, fmt.Errorf("tls handshake failed: %w", err)
}
conn = tlsConn
}
keyBytes := make([]byte, 16)
_, _ = rand.Read(keyBytes)
secKey := base64.StdEncoding.EncodeToString(keyBytes)
originScheme := "https"
if !isTLS {
originScheme = "http"
}
origin := fmt.Sprintf("%s://%s", originScheme, host)
req := fmt.Sprintf(
"GET /_stcore/stream HTTP/1.1\r\n"+
"Host: %s\r\n"+
"Upgrade: websocket\r\n"+
"Connection: Upgrade\r\n"+
"Sec-WebSocket-Key: %s\r\n"+
"Sec-WebSocket-Version: 13\r\n"+
"Origin: %s\r\n"+
"User-Agent: %s\r\n\r\n",
host, secKey, origin, userAgent,
)
if _, err := conn.Write([]byte(req)); err != nil {
conn.Close()
return nil, nil, fmt.Errorf("failed to write websocket handshake request: %w", err)
}
reader := bufio.NewReader(conn)
statusLine, err := reader.ReadString('\n')
if err != nil || !strings.Contains(statusLine, "101") {
conn.Close()
return nil, nil, fmt.Errorf("websocket handshake rejected: %s", statusLine)
}
for {
line, err := reader.ReadString('\n')
if err != nil || strings.TrimSpace(line) == "" {
break
}
}
return conn, reader, nil
}
func OpenStreamlitSession(ctx context.Context, targetURL, proxyURL, userAgent string) (*StreamlitSession, error) {
conn, reader, err := dialStreamlitWebSocket(ctx, targetURL, proxyURL, userAgent)
if err != nil {
return nil, err
}
sess := &StreamlitSession{
conn: conn,
reader: reader,
}
// Initial rerun_script BackMsg (ClientState: query_string = "")
clientState := encodeString(1, "")
backMsg := encodeLengthDelimited(11, clientState)
if err := sendWSBinaryFrame(conn, backMsg); err != nil {
sess.Close()
return nil, fmt.Errorf("failed to send initial rerun BackMsg: %w", err)
}
// Read frames until script_finished
for {
payload, opcode, err := readWSFrame(conn, reader, 25*time.Second)
if err != nil {
sess.Close()
return nil, fmt.Errorf("error reading initial session frames: %w", err)
}
if opcode == 0x09 { // Ping
_ = sendWSPongFrame(conn, payload)
continue
}
parsed := parseForwardMsg(payload)
if parsed.ChatInputID != "" {
sess.chatInputID = parsed.ChatInputID
}
if parsed.SelectboxID != "" {
sess.selectboxID = parsed.SelectboxID
sess.options = parsed.Options
}
if parsed.IsFinished {
break
}
}
if sess.chatInputID == "" {
// Default fallback chatInput widget ID if not matched in stream
sess.chatInputID = "$$ID-ae617304c8297cf4ddb1a23ee392ce5b-None"
}
if sess.selectboxID == "" {
sess.selectboxID = "$$ID-625687b40312ff8d73403e7858a2f552-chat_model_selector"
}
return sess, nil
}
func (s *StreamlitSession) SwitchModel(modelID string) error {
if modelID == "" || s.selectboxID == "" {
return nil
}
cleanModel := strings.TrimSpace(modelID)
if idx := strings.Index(cleanModel, " ("); idx != -1 {
cleanModel = strings.TrimSpace(cleanModel[:idx])
}
// Match model against options
matchedModel := cleanModel
for _, opt := range s.options {
rawOpt := opt
if idx := strings.Index(rawOpt, " ("); idx != -1 {
rawOpt = strings.TrimSpace(rawOpt[:idx])
}
if strings.EqualFold(rawOpt, cleanModel) || strings.Contains(strings.ToLower(rawOpt), strings.ToLower(cleanModel)) {
matchedModel = rawOpt
break
}
}
if s.activeModel == matchedModel {
return nil
}
// Streamlit Selectbox sends string_value (field 6) in WidgetState
sbWidget := append(encodeString(1, s.selectboxID), encodeString(6, matchedModel)...)
wsData := encodeLengthDelimited(1, sbWidget)
wStates := encodeLengthDelimited(2, wsData)
csData := append(encodeString(1, ""), wStates...)
switchBackMsg := encodeLengthDelimited(11, csData)
if err := sendWSBinaryFrame(s.conn, switchBackMsg); err != nil {
return fmt.Errorf("failed to send model switch BackMsg: %w", err)
}
for {
payload, opcode, err := readWSFrame(s.conn, s.reader, 25*time.Second)
if err != nil {
return fmt.Errorf("error reading model switch frames: %w", err)
}
if opcode == 0x09 {
_ = sendWSPongFrame(s.conn, payload)
continue
}
parsed := parseForwardMsg(payload)
if parsed.ChatInputID != "" {
s.chatInputID = parsed.ChatInputID
}
if parsed.SelectboxID != "" {
s.selectboxID = parsed.SelectboxID
}
if parsed.IsFinished {
break
}
}
s.activeModel = matchedModel
return nil
}
type GatewayAlertError struct {
Alert string
IsRateLimit bool
RetryAfter int
}
func (e *GatewayAlertError) Error() string {
return e.Alert
}
func parseRetryAfter(alert string) int {
re := regexp.MustCompile(`(?i)(?:retry\s+after|retry\s+in)\s+(\d+)`)
match := re.FindStringSubmatch(alert)
if len(match) >= 2 {
sec, _ := strconv.Atoi(match[1])
if sec > 0 {
return sec
}
}
return 6
}
func isRateLimit(alert string) bool {
lower := strings.ToLower(alert)
return strings.Contains(lower, "rate limit") || strings.Contains(lower, "too many requests") || strings.Contains(lower, "429")
}
func resolveAlert(alerts []string) (alertText string, isRL bool, retryAfter int) {
for _, a := range alerts {
if isRateLimit(a) {
return a, true, parseRetryAfter(a)
}
}
for _, a := range alerts {
lower := strings.ToLower(a)
if !strings.Contains(lower, "the last response did not finish") {
return a, false, 0
}
}
if len(alerts) > 0 {
return alerts[len(alerts)-1], false, 0
}
return "", false, 0
}
func isStaticUIMarkdown(text, promptText string) bool {
trimmed := strings.TrimSpace(text)
if trimmed == "" || trimmed == strings.TrimSpace(promptText) {
return true
}
staticPrefixes := []string{
"Model:",
"Model :",
"- **Model ID:**",
"[Powered by Groq]",
"---",
"-----",
"Text models available",
"This model accepts and returns text",
"This model uses non-streaming responses",
"Select a vision-capable model",
"Attached document:",
"### Upload an Image",
"### Usage Summary",
"### Chat Interface",
"### Upload Context",
"### Comparison Context",
"A document is already attached",
"**Token Usage",
"**Latest API Rate Snapshot:**",
"**Request Max Tokens:**",
"**Conversation Total:**",
"**Important:**",
"Important:",
"Partial response: generation did not finish.",
}
for _, p := range staticPrefixes {
if strings.HasPrefix(trimmed, p) {
return true
}
}
if strings.Contains(trimmed, "**Important:**") ||
strings.Contains(trimmed, "Important:") ||
strings.Contains(trimmed, "older messages were left out of this request") ||
strings.Contains(trimmed, "Response limit adjusted to") {
return true
}
return false
}
func (s *StreamlitSession) SubmitPrompt(promptText string, onChunk func(delta string), onAlert func(alert string)) error {
// Construct BackMsg with chat_input_value and selected model in selectbox
chatInputVal := encodeString(1, promptText)
ciWidget := append(encodeString(1, s.chatInputID), encodeLengthDelimited(15, chatInputVal)...)
var wsData []byte
if s.selectboxID != "" && s.activeModel != "" {
sbWidget := append(encodeString(1, s.selectboxID), encodeString(6, s.activeModel)...)
wsData = append(encodeLengthDelimited(1, sbWidget), encodeLengthDelimited(1, ciWidget)...)
} else {
wsData = encodeLengthDelimited(1, ciWidget)
}
wStates := encodeLengthDelimited(2, wsData)
csData := append(encodeString(1, ""), wStates...)
promptBackMsg := encodeLengthDelimited(11, csData)
if err := sendWSBinaryFrame(s.conn, promptBackMsg); err != nil {
return fmt.Errorf("failed to send prompt BackMsg: %w", err)
}
lastLen := 0
assistantPathKey := ""
userEchoSeen := false
var alerts []string
for {
payload, opcode, err := readWSFrame(s.conn, s.reader, 90*time.Second)
if err != nil {
return fmt.Errorf("connection closed during response: %w", err)
}
if opcode == 0x09 { // Ping
_ = sendWSPongFrame(s.conn, payload)
continue
}
parsed := parseForwardMsg(payload)
if parsed.AlertText != "" {
alerts = append(alerts, parsed.AlertText)
if onAlert != nil {
onAlert(parsed.AlertText)
}
}
if parsed.Markdown != "" {
// Ignore sidebar elements (deltaPath[0] == 1)
if len(parsed.DeltaPath) > 0 && parsed.DeltaPath[0] == 1 {
continue
}
// In this Streamlit space:
// DeltaPath[0] == 0 is the main page.
// DeltaPath[1] <= 5 are the header widgets (0-4) and the user chat echo container (5).
// The assistant response is always DeltaPath[1] >= 6.
if len(parsed.DeltaPath) >= 2 && parsed.DeltaPath[0] == 0 && parsed.DeltaPath[1] <= 5 {
continue
}
// Metadata or captions inside message containers have index >= 1 (e.g. [0, 6, 1])
if len(parsed.DeltaPath) >= 3 && parsed.DeltaPath[len(parsed.DeltaPath)-1] != 0 {
continue
}
// Check if user echo arrived (fallback for environments without hierarchical delta paths)
trimmedMD := strings.TrimSpace(parsed.Markdown)
trimmedPrompt := strings.TrimSpace(promptText)
if !userEchoSeen {
if trimmedMD == trimmedPrompt ||
strings.HasPrefix(trimmedPrompt, trimmedMD) ||
strings.HasPrefix(trimmedMD, trimmedPrompt) {
userEchoSeen = true
continue
}
}
// If text is not static UI, stream assistant delta
if !isStaticUIMarkdown(parsed.Markdown, promptText) {
pathKey := fmt.Sprintf("%v", parsed.DeltaPath)
if assistantPathKey == "" {
assistantPathKey = pathKey
}
if pathKey == assistantPathKey {
if len(parsed.Markdown) > lastLen {
delta := parsed.Markdown[lastLen:]
lastLen = len(parsed.Markdown)
if onChunk != nil {
onChunk(delta)
}
}
}
}
}
if parsed.IsFinished {
break
}
}
if lastLen == 0 {
if alertText, isRL, retrySec := resolveAlert(alerts); alertText != "" {
return &GatewayAlertError{
Alert: alertText,
IsRateLimit: isRL,
RetryAfter: retrySec,
}
}
return fmt.Errorf("empty assistant response received from space")
}
return nil
}
// ---------------------------------------------------------------------------
// Tool and Message Processing
// ---------------------------------------------------------------------------
func BuildToolInstruction(tools []Tool, toolChoice interface{}) string {
if len(tools) == 0 {
return ""
}
toolsBytes, _ := json.Marshal(tools)
var sb strings.Builder
sb.WriteString("\n\n# Tool Calling Instructions\nYou have access to the following functions:\n\n")
sb.WriteString(string(toolsBytes))
sb.WriteString("\n\n\n")
sb.WriteString("When you need to call a function, respond ONLY with a block formatted exactly as follows:\n")
sb.WriteString("\n{\"name\": \"\", \"arguments\": {}}\n\n\n")
sb.WriteString("CRITICAL EXECUTION RULES:\n")
sb.WriteString("1. If you invoke a tool, output ONLY the block. Do not write conversational greetings, explanations, or filler text outside the tags.\n")
sb.WriteString("2. Put any reasoning or thought process inside ... tags.\n")
sb.WriteString("3. If multiple tools need to be called, output each inside its own ... block.\n")
sb.WriteString("4. When tool execution results are provided in blocks, inspect the output:\n")
sb.WriteString(" - If further steps or additional tools are needed to fulfill the user's request, emit the next block.\n")
sb.WriteString(" - If all needed information has been retrieved, synthesize the answers and deliver the final response to the user.\n")
sb.WriteString(" - Never stop or terminate the conversation prematurely while intermediate steps remain.\n")
if toolChoice != nil {
if choiceStr, ok := toolChoice.(string); ok {
if choiceStr == "required" {
sb.WriteString("\n\nYou MUST invoke at least one tool from the list above.")
}
} else if choiceMap, ok := toolChoice.(map[string]interface{}); ok {
if fnMap, ok := choiceMap["function"].(map[string]interface{}); ok {
if fnName, ok := fnMap["name"].(string); ok && fnName != "" {
sb.WriteString(fmt.Sprintf("\n\nYou MUST call the function %q.", fnName))
}
}
}
}
return sb.String()
}
func ExtractToolName(fn interface{}) string {
if fn == nil {
return ""
}
if fnMap, ok := fn.(map[string]interface{}); ok {
if name, ok := fnMap["name"].(string); ok {
return strings.TrimSpace(name)
}
}
if tcf, ok := fn.(ToolCallFunction); ok {
return strings.TrimSpace(tcf.Name)
}
b, err := json.Marshal(fn)
if err == nil {
var m map[string]interface{}
if json.Unmarshal(b, &m) == nil {
if name, ok := m["name"].(string); ok {
return strings.TrimSpace(name)
}
}
}
return ""
}
func GetAllowedToolNames(tools []Tool, toolChoice interface{}) map[string]bool {
if len(tools) == 0 {
return nil
}
if choiceStr, ok := toolChoice.(string); ok {
if choiceStr == "none" {
return nil
}
}
if choiceMap, ok := toolChoice.(map[string]interface{}); ok {
if fnMap, ok := choiceMap["function"].(map[string]interface{}); ok {
if name, ok := fnMap["name"].(string); ok && strings.TrimSpace(name) != "" {
return map[string]bool{strings.TrimSpace(name): true}
}
}
}
allowed := make(map[string]bool)
for _, t := range tools {
name := ExtractToolName(t.Function)
if name != "" {
allowed[name] = true
}
}
if len(allowed) == 0 {
return nil
}
return allowed
}
var toolCallTagRegex = regexp.MustCompile(`(?s).*?`)
func stripToolCallTags(s string) string {
return strings.TrimSpace(toolCallTagRegex.ReplaceAllString(s, ""))
}
func FormatPrompt(req ChatCompletionRequest) string {
toolInstruction := BuildToolInstruction(req.Tools, req.ToolChoice)
// Build mapping from tool_call_id to function name
toolCallIDToName := make(map[string]string)
for _, msg := range req.Messages {
if msg.Role == "assistant" {
for _, tc := range msg.ToolCalls {
if tc.ID != "" && tc.Function.Name != "" {
toolCallIDToName[tc.ID] = tc.Function.Name
}
}
}
}
// Find where trailing tool responses start (if any)
firstTrailingToolIdx := len(req.Messages)
for i := len(req.Messages) - 1; i >= 0; i-- {
role := req.Messages[i].Role
if role == "tool" || role == "function" {
firstTrailingToolIdx = i
} else {
break
}
}
var systemInstructions []string
var historyTurns []string
var currentTurn string
// Process messages before the trailing tool responses
endHistoryIdx := firstTrailingToolIdx
if firstTrailingToolIdx == len(req.Messages) && len(req.Messages) > 0 {
if req.Messages[len(req.Messages)-1].Role == "user" {
endHistoryIdx = len(req.Messages) - 1
currentTurn = req.Messages[len(req.Messages)-1].GetContentString()
}
}
for i := 0; i < endHistoryIdx; i++ {
msg := req.Messages[i]
contentStr := msg.GetContentString()
switch msg.Role {
case "system":
if contentStr != "" {
systemInstructions = append(systemInstructions, contentStr)
}
case "assistant":
cleanAssistant := strings.TrimSpace(contentStr)
if strings.HasPrefix(cleanAssistant, "Model:") || strings.HasPrefix(cleanAssistant, "Model :") {
cleanAssistant = ""
}
if len(msg.ToolCalls) > 0 {
cleanAssistant = stripToolCallTags(cleanAssistant)
}
var sb strings.Builder
if cleanAssistant != "" {
sb.WriteString(cleanAssistant)
}
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))
}
if sb.Len() > 0 {
historyTurns = append(historyTurns, "Assistant: "+sb.String())
}
case "tool", "function":
toolName := msg.Name
if resolved, ok := toolCallIDToName[msg.ToolCallID]; ok && resolved != "" {
toolName = resolved
} else if toolName == "" {
toolName = msg.ToolCallID
}
var contentJSON []byte
if json.Valid([]byte(contentStr)) {
contentJSON = []byte(contentStr)
} else {
contentJSON, _ = json.Marshal(contentStr)
}
jsonStr := string(contentJSON)
if len(jsonStr) > 2000 {
jsonStr = jsonStr[:1800] + `"... [truncated]"`
}
var turnText string
if msg.ToolCallID != "" {
turnText = fmt.Sprintf("\n{\"name\": %q, \"tool_call_id\": %q, \"content\": %s}\n", toolName, msg.ToolCallID, jsonStr)
} else {
turnText = fmt.Sprintf("\n{\"name\": %q, \"content\": %s}\n", toolName, jsonStr)
}
historyTurns = append(historyTurns, fmt.Sprintf("Tool Result (%s): %s", toolName, turnText))
case "user":
historyTurns = append(historyTurns, "User: "+contentStr)
}
}
// If we have trailing tool responses, group them together with a continuation directive
if firstTrailingToolIdx < len(req.Messages) {
var toolResSb strings.Builder
toolResSb.WriteString("[Tool Execution Results]\n")
for i := firstTrailingToolIdx; i < len(req.Messages); i++ {
msg := req.Messages[i]
toolName := msg.Name
if resolved, ok := toolCallIDToName[msg.ToolCallID]; ok && resolved != "" {
toolName = resolved
} else if toolName == "" {
toolName = msg.ToolCallID
}
contentStr := msg.GetContentString()
var contentJSON []byte
if json.Valid([]byte(contentStr)) {
contentJSON = []byte(contentStr)
} else {
contentJSON, _ = json.Marshal(contentStr)
}
jsonStr := string(contentJSON)
if len(jsonStr) > 3500 {
jsonStr = jsonStr[:3200] + `"... [truncated]"`
}
toolResSb.WriteString("\n")
if msg.ToolCallID != "" {
toolResSb.WriteString(fmt.Sprintf("{\"name\": %q, \"tool_call_id\": %q, \"content\": %s}\n", toolName, msg.ToolCallID, jsonStr))
} else {
toolResSb.WriteString(fmt.Sprintf("{\"name\": %q, \"content\": %s}\n", toolName, jsonStr))
}
toolResSb.WriteString("\n")
}
toolResSb.WriteString("\n[Next Steps Directive]\n")
toolResSb.WriteString("You have received the results of the tool execution(s) above.\n")
toolResSb.WriteString("- Analyze these results in the context of the user request and conversation history.\n")
toolResSb.WriteString("- If additional tool calls are needed to complete the task, emit the next block immediately.\n")
toolResSb.WriteString("- If all information needed to fulfill the request is now available, provide the final comprehensive response to the user.\n")
toolResSb.WriteString("- Do not terminate or stop without answering the user or taking the next step.")
currentTurn = toolResSb.String()
}
// Prune history turns if history exceeds budget (~12000 chars)
totalHistLen := 0
for _, t := range historyTurns {
totalHistLen += len(t)
}
if totalHistLen > 12000 && len(historyTurns) > 3 {
firstTurn := historyTurns[0]
var pruned []string
pruned = append(pruned, firstTurn)
pruned = append(pruned, "[... earlier conversation turns omitted for brevity ...]")
budget := 10000
var recent []string
currentBudget := 0
for i := len(historyTurns) - 1; i >= 1; i-- {
tLen := len(historyTurns[i])
if currentBudget+tLen > budget {
break
}
currentBudget += tLen
recent = append([]string{historyTurns[i]}, recent...)
}
pruned = append(pruned, recent...)
historyTurns = pruned
}
var promptBuilder strings.Builder
if len(systemInstructions) > 0 || toolInstruction != "" {
promptBuilder.WriteString("[System Instructions]\n")
for _, sys := range systemInstructions {
promptBuilder.WriteString(sys)
promptBuilder.WriteString("\n")
}
if toolInstruction != "" {
promptBuilder.WriteString(strings.TrimSpace(toolInstruction))
promptBuilder.WriteString("\n")
}
promptBuilder.WriteString("\n")
}
if len(historyTurns) > 0 {
promptBuilder.WriteString("[Conversation History]\n")
for _, turn := range historyTurns {
promptBuilder.WriteString(turn)
promptBuilder.WriteString("\n")
}
promptBuilder.WriteString("\n")
}
if currentTurn != "" {
if strings.HasPrefix(currentTurn, "[Tool Execution Results]") {
promptBuilder.WriteString(currentTurn)
} else {
if len(historyTurns) > 0 || len(systemInstructions) > 0 {
promptBuilder.WriteString("[User]\n")
}
promptBuilder.WriteString(currentTurn)
}
}
return strings.TrimSpace(promptBuilder.String())
}
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{}:
cleanMap := make(map[string]interface{})
for k, childV := range val {
cleanKey := strings.TrimSpace(k)
cleanMap[cleanKey] = sanitizeJSONValue(childV)
}
return cleanMap
case []interface{}:
cleanSlice := make([]interface{}, len(val))
for i, childV := range val {
cleanSlice[i] = sanitizeJSONValue(childV)
}
return cleanSlice
default:
return v
}
}
var toolNameRegex = regexp.MustCompile(`"\s*(?:name|function|action|call)\s*"\s*:\s*"\s*([^"]+?)\s*"`)
func repairToolCallJSON(jsonStr string) (ToolCall, bool) {
nameMatch := toolNameRegex.FindStringSubmatch(jsonStr)
if len(nameMatch) < 2 {
return ToolCall{}, false
}
nameVal := strings.TrimSpace(nameMatch[1])
argsStr := "{}"
argsKwList := []string{`"arguments"`, `" parameters "`, `"arguments "`, `" parameters"`, `"parameters"`, `"args"`, `"input"`}
argsIdx := -1
for _, kw := range argsKwList {
idx := strings.Index(jsonStr, kw)
if idx >= 0 {
argsIdx = idx + len(kw)
break
}
}
var targetStr string
if argsIdx >= 0 {
targetStr = strings.TrimSpace(jsonStr[argsIdx:])
if strings.HasPrefix(targetStr, ":") {
targetStr = strings.TrimSpace(targetStr[1:])
}
} else {
targetStr = jsonStr
}
if strings.HasPrefix(targetStr, "{") {
endIdx := strings.LastIndex(targetStr, "}")
if endIdx > 0 {
objCandidate := targetStr[:endIdx+1]
var testMap map[string]interface{}
if json.Unmarshal([]byte(objCandidate), &testMap) == nil {
b, _ := json.Marshal(sanitizeJSONValue(testMap))
return ToolCall{
ID: "call_" + GenerateUUID()[:8],
Type: "function",
Function: ToolCallFunction{
Name: nameVal,
Arguments: string(b),
},
}, true
}
}
} else if strings.HasPrefix(targetStr, `"`) {
endIdx := strings.LastIndex(targetStr, `"`)
if endIdx > 0 {
val := strings.TrimSpace(targetStr[1:endIdx])
b, _ := json.Marshal(val)
argsStr = string(b)
}
}
return ToolCall{
ID: "call_" + GenerateUUID()[:8],
Type: "function",
Function: ToolCallFunction{
Name: nameVal,
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 parseToolCallsFromText(input string) ([]ToolCall, bool) {
cleaned := cleanJSONBlock(strings.TrimSpace(input))
if cleaned == "" {
return nil, false
}
// Check if it is a JSON array of tool calls: [...]
if strings.HasPrefix(cleaned, "[") && strings.HasSuffix(cleaned, "]") {
var rawSlice []interface{}
if err := json.Unmarshal([]byte(cleaned), &rawSlice); err == nil && len(rawSlice) > 0 {
var calls []ToolCall
for _, item := range rawSlice {
b, _ := json.Marshal(item)
if tc, ok := parseSingleToolCall(string(b)); ok {
calls = append(calls, tc)
} else if tc, ok := repairToolCallJSON(string(b)); ok {
calls = append(calls, tc)
}
}
if len(calls) > 0 {
return calls, true
}
}
}
// Check 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)
if strings.HasPrefix(inner, "") {
inner = strings.TrimPrefix(inner, "")
}
if strings.HasSuffix(inner, "") {
inner = strings.TrimSuffix(inner, "")
}
inner = cleanJSONBlock(inner)
if calls, ok := parseToolCallsFromText(inner); ok {
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) {
var b []string
s := content
searchOffset := 0
var remainingBuilder strings.Builder
lastPos := 0
for {
relStart := strings.Index(s[searchOffset:], "")
if relStart == -1 {
break
}
startIdx := searchOffset + relStart
afterStart := startIdx + len("")
relEnd := strings.Index(s[afterStart:], "")
if relEnd == -1 {
// Unclosed , do not treat as block
searchOffset = afterStart
continue
}
if innerNext := strings.Index(s[afterStart:afterStart+relEnd], ""); innerNext != -1 {
searchOffset = afterStart + innerNext
continue
}
endIdx := afterStart + relEnd + len("")
blockText := s[startIdx:endIdx]
b = append(b, blockText)
if startIdx > lastPos {
remainingBuilder.WriteString(s[lastPos:startIdx])
}
lastPos = endIdx
searchOffset = endIdx
}
if lastPos < len(s) {
remainingBuilder.WriteString(s[lastPos:])
}
return b, strings.TrimSpace(remainingBuilder.String())
}
func DetectToolCalls(content string, allowedTools map[string]bool) ([]ToolCall, string, bool) {
if len(allowedTools) == 0 {
return nil, content, false
}
type validBlock struct {
startIdx int
endIdx int
calls []ToolCall
}
var validBlocks []validBlock
searchOffset := 0
for {
relStart := strings.Index(content[searchOffset:], "")
if relStart == -1 {
break
}
startIdx := searchOffset + relStart
afterStart := startIdx + len("")
relEnd := strings.Index(content[afterStart:], "")
if relEnd == -1 {
// Unclosed tag - leave as plain content
searchOffset = afterStart
continue
}
if innerNext := strings.Index(content[afterStart:afterStart+relEnd], ""); innerNext != -1 {
searchOffset = afterStart + innerNext
continue
}
endIdx := afterStart + relEnd + len("")
blockText := content[startIdx:endIdx]
if parsedCalls, ok := parseXMLToolCall(blockText); ok {
var callsForBlock []ToolCall
for _, c := range parsedCalls {
if allowedTools[c.Function.Name] {
callsForBlock = append(callsForBlock, c)
}
}
if len(callsForBlock) > 0 {
validBlocks = append(validBlocks, validBlock{
startIdx: startIdx,
endIdx: endIdx,
calls: callsForBlock,
})
searchOffset = endIdx
continue
}
}
// Not a valid tool call or function name not allowed - keep in content
searchOffset = endIdx
}
if len(validBlocks) > 0 {
var allCalls []ToolCall
var remainingBuilder strings.Builder
lastPos := 0
for _, vb := range validBlocks {
if vb.startIdx > lastPos {
remainingBuilder.WriteString(content[lastPos:vb.startIdx])
}
allCalls = append(allCalls, vb.calls...)
lastPos = vb.endIdx
}
if lastPos < len(content) {
remainingBuilder.WriteString(content[lastPos:])
}
for i := range allCalls {
idx := i
allCalls[i].Index = &idx
}
cleanRemaining := strings.TrimSpace(remainingBuilder.String())
return allCalls, cleanRemaining, true
}
// Direct JSON fallback (only if every parsed call is present in allowedTools)
trimmed := strings.TrimSpace(content)
if (strings.HasPrefix(trimmed, "{") && strings.HasSuffix(trimmed, "}")) ||
(strings.HasPrefix(trimmed, "[") && strings.HasSuffix(trimmed, "]")) ||
(strings.HasPrefix(trimmed, "```json") && strings.HasSuffix(trimmed, "```")) {
if directCalls, ok := parseToolCallsFromText(trimmed); ok {
var validCalls []ToolCall
for _, c := range directCalls {
if allowedTools[c.Function.Name] {
validCalls = append(validCalls, c)
}
}
if len(validCalls) > 0 && len(validCalls) == len(directCalls) {
for i := range validCalls {
idx := i
validCalls[i].Index = &idx
}
return validCalls, "", 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 strings.TrimSpace(rem), strings.TrimSpace(reasoning)
}
}
return content, ""
}
// ---------------------------------------------------------------------------
// Streaming Filters & Framing
// ---------------------------------------------------------------------------
type StreamThinkingFilter struct {
inThinking bool
buf string
}
func NewStreamThinkingFilter() *StreamThinkingFilter {
return &StreamThinkingFilter{}
}
func hasPrefixOf(target string, prefixes []string) int {
for _, p := range prefixes {
if strings.HasSuffix(target, p) {
return len(p)
}
}
return 0
}
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 = ""
}
}
type StreamToolCallFilter struct {
allowedTools map[string]bool
inToolCall bool
buf string
toolCallBuf string
toolIndex int
emittedCall bool
preambleBuf string
}
func NewStreamToolCallFilter(allowedTools map[string]bool) *StreamToolCallFilter {
return &StreamToolCallFilter{
allowedTools: allowedTools,
}
}
func (f *StreamToolCallFilter) appendPreambleOrContent(text string, onContent func(string)) {
if text == "" {
return
}
if f.emittedCall {
onContent(text)
return
}
if len(f.preambleBuf)+len(text) < 128 && !strings.Contains(f.preambleBuf+text, "\n\n") {
f.preambleBuf += text
} else {
if f.preambleBuf != "" {
onContent(f.preambleBuf)
f.preambleBuf = ""
}
onContent(text)
}
}
func (f *StreamToolCallFilter) Feed(chunk string, onContent func(string), onReasoning func(string), onToolCall func(ToolCall)) {
if len(f.allowedTools) == 0 {
if chunk != "" {
onContent(chunk)
}
return
}
f.buf += chunk
toolStartTag := ""
toolEndTag := ""
startPrefixes := []string{"<", " 0 {
if !f.inToolCall {
if idx := strings.Index(f.buf, toolStartTag); idx != -1 {
before := f.buf[:idx]
f.inToolCall = true
f.buf = f.buf[idx+len(toolStartTag):]
f.preambleBuf += before
} else if matchLen := hasPrefixOf(f.buf, startPrefixes); matchLen > 0 {
safe := f.buf[:len(f.buf)-matchLen]
f.buf = f.buf[len(f.buf)-matchLen:]
f.appendPreambleOrContent(safe, onContent)
break
} else {
safe := f.buf
f.buf = ""
f.appendPreambleOrContent(safe, onContent)
break
}
} else {
if idx := strings.Index(f.buf, toolEndTag); idx != -1 {
f.toolCallBuf += f.buf[:idx]
f.buf = f.buf[idx+len(toolEndTag):]
f.inToolCall = false
rawBlock := "" + f.toolCallBuf + ""
tcs, ok := parseXMLToolCall(rawBlock)
var validCalls []ToolCall
if ok {
for _, tc := range tcs {
if f.allowedTools[tc.Function.Name] {
validCalls = append(validCalls, tc)
}
}
}
if len(validCalls) > 0 {
if strings.TrimSpace(f.preambleBuf) != "" {
onReasoning(f.preambleBuf)
}
f.preambleBuf = ""
for _, tc := range validCalls {
idxCopy := f.toolIndex
tc.Index = &idxCopy
f.toolIndex++
f.emittedCall = true
onToolCall(tc)
}
} else {
// Tool call invalid or function name not allowed - flush as content
if f.preambleBuf != "" {
onContent(f.preambleBuf)
f.preambleBuf = ""
}
onContent(rawBlock)
}
f.toolCallBuf = ""
} else if matchLen := hasPrefixOf(f.buf, endPrefixes); 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), onReasoning func(string), onToolCall func(ToolCall)) {
if len(f.allowedTools) == 0 {
if f.buf != "" {
onContent(f.buf)
f.buf = ""
}
return
}
if f.inToolCall && len(f.toolCallBuf) > 0 {
rawBlock := "" + f.toolCallBuf + ""
tcs, ok := parseXMLToolCall(rawBlock)
var validCalls []ToolCall
if ok {
for _, tc := range tcs {
if f.allowedTools[tc.Function.Name] {
validCalls = append(validCalls, tc)
}
}
}
if len(validCalls) > 0 {
if strings.TrimSpace(f.preambleBuf) != "" {
onReasoning(f.preambleBuf)
}
f.preambleBuf = ""
for _, tc := range validCalls {
idxCopy := f.toolIndex
tc.Index = &idxCopy
f.toolIndex++
f.emittedCall = true
onToolCall(tc)
}
} else {
if f.preambleBuf != "" {
onContent(f.preambleBuf)
f.preambleBuf = ""
}
onContent("" + f.toolCallBuf)
}
f.toolCallBuf = ""
f.inToolCall = false
}
if f.emittedCall {
if strings.TrimSpace(f.preambleBuf) != "" {
onReasoning(f.preambleBuf)
}
f.preambleBuf = ""
} else {
if f.preambleBuf != "" {
onContent(f.preambleBuf)
f.preambleBuf = ""
}
}
if len(f.buf) > 0 {
onContent(f.buf)
f.buf = ""
}
}
type Streamer struct {
w http.ResponseWriter
flusher http.Flusher
id string
created int64
model string
}
func NewStreamer(w http.ResponseWriter, flusher http.Flusher, id string, created int64, model string) *Streamer {
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
w.Header().Set("Connection", "keep-alive")
return &Streamer{w: w, flusher: flusher, id: id, created: created, model: model}
}
func (s *Streamer) Role() {
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{Role: "assistant"}, "")
}
func (s *Streamer) Reasoning(text string) {
if text != "" {
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{ReasoningContent: text}, "")
}
}
func (s *Streamer) Content(text string) {
if text != "" {
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{Content: text}, "")
}
}
func (s *Streamer) ToolCall(tc ToolCall) {
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{ToolCalls: []ToolCall{tc}}, "")
}
func (s *Streamer) Finish(reason string) {
sendStreamDelta(s.w, s.flusher, s.id, s.created, s.model, StreamDelta{}, reason)
fmt.Fprintf(s.w, "data: [DONE]\n\n")
if s.flusher != nil {
s.flusher.Flush()
}
}
func sendStreamDelta(w http.ResponseWriter, flusher http.Flusher, id string, created int64, model string, delta StreamDelta, finishReason string) {
chunk := StreamResponse{
ID: id,
Object: "chat.completion.chunk",
Created: created,
Model: model,
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()
}
}
func WriteCompletionResponse(w http.ResponseWriter, completionID string, created int64, model string, content string, reasoning string, toolCalls []ToolCall, finishReason string) {
if finishReason == "" {
finishReason = "stop"
}
var contentVal interface{} = content
if len(toolCalls) > 0 {
finishReason = "tool_calls"
if reasoning == "" && strings.TrimSpace(content) != "" {
reasoning = strings.TrimSpace(content)
}
contentVal = nil
}
resp := ChatCompletionResponse{
ID: completionID,
Object: "chat.completion",
Created: created,
Model: model,
Choices: []ChatCompletionResponseChoice{
{
Index: 0,
Message: ChatMessage{
Role: "assistant",
Content: contentVal,
ReasoningContent: reasoning,
ToolCalls: toolCalls,
},
FinishReason: finishReason,
},
},
Usage: Usage{
PromptTokens: 0,
CompletionTokens: 0,
TotalTokens: 0,
},
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
}
type APIErrorDetail struct {
Message string `json:"message"`
Type string `json:"type"`
Param *string `json:"param"`
Code int `json:"code"`
}
type APIErrorResponse struct {
Error APIErrorDetail `json:"error"`
}
func WriteAPIError(w http.ResponseWriter, status int, message string, errType string, code int) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
resp := APIErrorResponse{
Error: APIErrorDetail{
Message: message,
Type: errType,
Code: code,
},
}
_ = json.NewEncoder(w).Encode(resp)
}
// ---------------------------------------------------------------------------
// Gateway Controller & Request Orchestration
// ---------------------------------------------------------------------------
type GroqqerGateway struct {
targetURL string
proxyURL string
userAgent string
defaultModel string
port int
models []ModelItem
modelsMu sync.RWMutex
}
func DefaultFallbackModels() []string {
return []string{
"qwen/qwen3.6-27b",
"qwen/qwen3.8-27b",
"openai/gpt-oss-120b",
"openai/gpt-oss-20b",
"groq/compound",
"groq/compound-mini",
"allam-2-7b",
"llama-3.3-70b-versatile",
"llama-3.1-8b-instant",
"meta-llama/llama-4-scout-17b-16e-instruct",
}
}
func (g *GroqqerGateway) RefreshModels() {
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second)
defer cancel()
sess, err := OpenStreamlitSession(ctx, g.targetURL, g.proxyURL, g.userAgent)
if err != nil {
log.Printf("warning: model discovery session failed: %v", err)
if len(g.GetModels()) == 0 {
g.setFallbackModels()
}
return
}
defer sess.Close()
now := time.Now().Unix()
var items []ModelItem
for _, opt := range sess.options {
rawOpt := opt
if idx := strings.Index(rawOpt, " ("); idx != -1 {
rawOpt = strings.TrimSpace(rawOpt[:idx])
}
if rawOpt != "" {
items = append(items, ModelItem{
ID: rawOpt,
Object: "model",
Created: now,
OwnedBy: "groq",
})
}
}
if len(items) == 0 {
g.setFallbackModels()
return
}
g.modelsMu.Lock()
g.models = items
g.modelsMu.Unlock()
}
func (g *GroqqerGateway) setFallbackModels() {
g.modelsMu.Lock()
defer g.modelsMu.Unlock()
now := time.Now().Unix()
var items []ModelItem
for _, id := range DefaultFallbackModels() {
items = append(items, ModelItem{
ID: id,
Object: "model",
Created: now,
OwnedBy: "groq",
})
}
g.models = items
}
func (g *GroqqerGateway) GetModels() []ModelItem {
g.modelsMu.RLock()
defer g.modelsMu.RUnlock()
if len(g.models) == 0 {
now := time.Now().Unix()
var items []ModelItem
for _, id := range DefaultFallbackModels() {
items = append(items, ModelItem{
ID: id,
Object: "model",
Created: now,
OwnedBy: "groq",
})
}
return items
}
out := make([]ModelItem, len(g.models))
copy(out, g.models)
return out
}
func (g *GroqqerGateway) MatchModel(requestedModel string) string {
if requestedModel == "" {
return g.defaultModel
}
models := g.GetModels()
for _, m := range models {
if strings.EqualFold(m.ID, requestedModel) {
return m.ID
}
}
// Try suffix match e.g. "qwen3.6-27b" for "qwen/qwen3.6-27b"
cleanReq := strings.ToLower(strings.TrimSpace(requestedModel))
if strings.Contains(cleanReq, "/") {
parts := strings.Split(cleanReq, "/")
cleanReq = parts[len(parts)-1]
}
for _, m := range models {
if strings.Contains(strings.ToLower(m.ID), cleanReq) {
return m.ID
}
}
return requestedModel
}
func (g *GroqqerGateway) HandleModels(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Access-Control-Allow-Origin", "*")
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type")
if r.Method == "OPTIONS" {
w.WriteHeader(http.StatusOK)
return
}
models := g.GetModels()
resp := ModelsResponse{
Object: "list",
Data: models,
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(resp)
}
func (g *GroqqerGateway) HandleChatCompletions(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Access-Control-Allow-Origin", "*")
w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type")
if r.Method == "OPTIONS" {
w.WriteHeader(http.StatusOK)
return
}
if r.Method != "POST" {
WriteAPIError(w, http.StatusMethodNotAllowed, "Method not allowed", "invalid_request_error", 405)
return
}
var req ChatCompletionRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
WriteAPIError(w, http.StatusBadRequest, "Invalid JSON payload: "+err.Error(), "invalid_request_error", 400)
return
}
if len(req.Messages) == 0 {
WriteAPIError(w, http.StatusBadRequest, "messages array must not be empty", "invalid_request_error", 400)
return
}
requestedModel := g.MatchModel(req.Model)
promptText := FormatPrompt(req)
completionID := "chatcmpl-" + GenerateUUID()
created := time.Now().Unix()
allowedTools := GetAllowedToolNames(req.Tools, req.ToolChoice)
log.Printf("Handling completion: model=%s, stream=%v, messages=%d", requestedModel, req.Stream, len(req.Messages))
ctx, cancel := context.WithTimeout(r.Context(), 180*time.Second)
defer cancel()
start := time.Now()
maxAttempts := 3
for attempt := 0; attempt < maxAttempts; attempt++ {
sess, err := OpenStreamlitSession(ctx, g.targetURL, g.proxyURL, g.userAgent)
if err != nil {
log.Printf("session connection error (attempt %d): %v", attempt+1, err)
if attempt < maxAttempts-1 {
time.Sleep(1 * time.Second)
continue
}
WriteAPIError(w, http.StatusBadGateway, "Failed to connect to target space: "+err.Error(), "api_error", 502)
return
}
if err := sess.SwitchModel(requestedModel); err != nil {
log.Printf("model switch warning: %v", err)
}
effectiveModel := sess.activeModel
if effectiveModel == "" {
effectiveModel = requestedModel
}
if req.Stream {
flusher, ok := w.(http.Flusher)
if !ok {
sess.Close()
WriteAPIError(w, http.StatusInternalServerError, "Streaming unsupported", "api_error", 500)
return
}
var streamer *Streamer
var thinkingFilter *StreamThinkingFilter
var toolFilter *StreamToolCallFilter
var streamStarted bool
var alertMsg string
initStreamer := func() {
if !streamStarted {
streamer = NewStreamer(w, flusher, completionID, created, effectiveModel)
streamer.Role()
thinkingFilter = NewStreamThinkingFilter()
toolFilter = NewStreamToolCallFilter(allowedTools)
streamStarted = true
}
}
err = sess.SubmitPrompt(promptText,
func(delta string) {
initStreamer()
thinkingFilter.Feed(delta,
func(cleanContent string) {
toolFilter.Feed(cleanContent,
func(userText string) {
streamer.Content(userText)
},
func(reasoning string) {
streamer.Reasoning(reasoning)
},
func(tc ToolCall) {
streamer.ToolCall(tc)
},
)
},
func(reasoning string) {
streamer.Reasoning(reasoning)
},
)
},
func(alert string) {
alertMsg = alert
},
)
sess.Close()
if err != nil {
var alertErr *GatewayAlertError
if errors.As(err, &alertErr) && !streamStarted && alertErr.IsRateLimit {
if attempt < maxAttempts-1 {
waitSec := alertErr.RetryAfter
if waitSec <= 0 {
waitSec = 6
}
if waitSec > 20 {
waitSec = 20
}
log.Printf("Rate limit hit on streaming %s. Waiting %d seconds before retry (attempt %d)...", effectiveModel, waitSec, attempt+1)
select {
case <-time.After(time.Duration(waitSec)*time.Second + 500*time.Millisecond):
continue
case <-ctx.Done():
WriteAPIError(w, http.StatusTooManyRequests, "Rate limit retry cancelled: context deadline exceeded", "rate_limit_error", 429)
return
}
}
w.Header().Set("Retry-After", strconv.Itoa(alertErr.RetryAfter))
WriteAPIError(w, http.StatusTooManyRequests, alertErr.Alert, "rate_limit_error", 429)
return
}
if !streamStarted {
if errors.As(err, &alertErr) {
lower := strings.ToLower(alertErr.Alert)
if strings.Contains(lower, "too large") || strings.Contains(lower, "cannot fit") {
WriteAPIError(w, http.StatusBadRequest, alertErr.Alert, "invalid_request_error", 400)
return
}
WriteAPIError(w, http.StatusBadGateway, alertErr.Alert, "api_error", 502)
return
}
WriteAPIError(w, http.StatusInternalServerError, "Streaming error: "+err.Error(), "api_error", 500)
return
}
if alertMsg != "" {
streamer.Content("\n\n" + alertMsg)
}
}
if streamStarted {
finishReason := "stop"
thinkingFilter.Flush(
func(cleanContent string) {
toolFilter.Feed(cleanContent,
func(userText string) {
streamer.Content(userText)
},
func(reasoning string) {
streamer.Reasoning(reasoning)
},
func(tc ToolCall) {
streamer.ToolCall(tc)
},
)
},
func(reasoning string) {
streamer.Reasoning(reasoning)
},
)
toolFilter.Flush(
func(userText string) {
streamer.Content(userText)
},
func(reasoning string) {
streamer.Reasoning(reasoning)
},
func(tc ToolCall) {
streamer.ToolCall(tc)
},
)
if toolFilter.emittedCall {
finishReason = "tool_calls"
}
streamer.Finish(finishReason)
log.Printf("Streaming completion finished for %s in %v", effectiveModel, time.Since(start))
return
}
WriteAPIError(w, http.StatusInternalServerError, "Empty response from space", "api_error", 500)
return
}
// Non-streaming completion
var fullBuilder strings.Builder
var alertMsg string
err = sess.SubmitPrompt(promptText,
func(delta string) {
fullBuilder.WriteString(delta)
},
func(alert string) {
alertMsg = alert
},
)
sess.Close()
_ = alertMsg
if err != nil {
var alertErr *GatewayAlertError
if errors.As(err, &alertErr) && alertErr.IsRateLimit {
if attempt < maxAttempts-1 {
waitSec := alertErr.RetryAfter
if waitSec <= 0 {
waitSec = 6
}
if waitSec > 20 {
waitSec = 20
}
log.Printf("Rate limit hit on %s. Waiting %d seconds before retry (attempt %d)...", effectiveModel, waitSec, attempt+1)
select {
case <-time.After(time.Duration(waitSec)*time.Second + 500*time.Millisecond):
continue
case <-ctx.Done():
WriteAPIError(w, http.StatusTooManyRequests, "Rate limit retry cancelled: context deadline exceeded", "rate_limit_error", 429)
return
}
}
w.Header().Set("Retry-After", strconv.Itoa(alertErr.RetryAfter))
WriteAPIError(w, http.StatusTooManyRequests, alertErr.Alert, "rate_limit_error", 429)
return
}
if errors.As(err, &alertErr) {
lower := strings.ToLower(alertErr.Alert)
if strings.Contains(lower, "too large") || strings.Contains(lower, "cannot fit") {
WriteAPIError(w, http.StatusBadRequest, alertErr.Alert, "invalid_request_error", 400)
return
}
WriteAPIError(w, http.StatusBadGateway, alertErr.Alert, "api_error", 502)
return
}
WriteAPIError(w, http.StatusInternalServerError, "Prompt execution error: "+err.Error(), "api_error", 500)
return
}
fullText := fullBuilder.String()
if strings.TrimSpace(fullText) == "" {
WriteAPIError(w, http.StatusInternalServerError, "The model returned an empty response", "api_error", 500)
return
}
cleanContent, reasoningContent := ExtractThinking(fullText)
toolCalls, remainingContent, hasTools := DetectToolCalls(cleanContent, allowedTools)
finishReason := "stop"
if hasTools {
finishReason = "tool_calls"
cleanContent = remainingContent
}
WriteCompletionResponse(w, completionID, created, effectiveModel, cleanContent, reasoningContent, toolCalls, finishReason)
log.Printf("Completion finished for %s in %v", effectiveModel, time.Since(start))
return
}
}
// ---------------------------------------------------------------------------
// Main Entrypoint & Signal Handling
// ---------------------------------------------------------------------------
func main() {
port := flag.Int("port", 8080, "HTTP server listening port")
targetURL := flag.String("target", DefaultTargetURL, "Groq Streamlit space URL")
defaultModel := flag.String("default-model", DefaultModel, "Default model ID if unspecified")
userAgent := flag.String("user-agent", DefaultUserAgent, "User-Agent string")
flag.StringVar(userAgent, "ua", DefaultUserAgent, "User-Agent string (alias)")
socksProxy := flag.String("socks", "", "SOCKS5 proxy (e.g. socks5://127.0.0.1:1080)")
flag.StringVar(socksProxy, "proxy", "", "Proxy address alias")
flag.StringVar(socksProxy, "socks5", "", "SOCKS5 proxy alias")
// Retain flags for backward compatibility (browser/Xvfb are no longer required)
_ = flag.String("browser", "", "Ignored: browser is no longer required (headless WebSocket mode)")
_ = flag.Bool("xvfb", false, "Ignored: Xvfb is no longer required")
_ = flag.Bool("no-xvfb", false, "Ignored: Xvfb is no longer required")
_ = flag.Bool("headless", true, "Ignored: always runs headless via direct WebSocket")
flag.Parse()
effectiveProxy := *socksProxy
if effectiveProxy == "" {
for _, envVar := range []string{"ALL_PROXY", "all_proxy", "SOCKS5_PROXY", "socks5_proxy", "SOCKS_PROXY", "socks_proxy"} {
if val := os.Getenv(envVar); val != "" {
effectiveProxy = val
break
}
}
}
log.Printf("Initializing groqqer gateway (pure Go WebSocket client, zero browser dependency)...")
gateway := &GroqqerGateway{
targetURL: *targetURL,
proxyURL: effectiveProxy,
userAgent: *userAgent,
defaultModel: *defaultModel,
port: *port,
}
// Initial discovery of live models in background
go func() {
log.Println("Discovering available models from Streamlit space...")
gateway.RefreshModels()
models := gateway.GetModels()
log.Printf("Ready with %d active models. Default: %s", len(models), gateway.defaultModel)
}()
mux := http.NewServeMux()
mux.HandleFunc("/v1/models", gateway.HandleModels)
mux.HandleFunc("/models", gateway.HandleModels)
mux.HandleFunc("/v1/chat/completions", gateway.HandleChatCompletions)
mux.HandleFunc("/chat/completions", gateway.HandleChatCompletions)
// Root status handler
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/" {
http.NotFound(w, r)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]interface{}{
"service": "groqqer",
"mode": "headless-websocket-protobuf",
"status": "running",
"target": *targetURL,
"default": *defaultModel,
"models": "/v1/models",
"endpoints": []string{"/v1/chat/completions", "/v1/models"},
})
})
server := &http.Server{
Addr: fmt.Sprintf(":%d", *port),
Handler: mux,
}
// Graceful shutdown handling
sigChan := make(chan os.Signal, 1)
signal.Notify(sigChan, os.Interrupt, syscall.SIGTERM)
go func() {
<-sigChan
log.Println("Shutting down groqqer...")
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_ = server.Shutdown(ctx)
os.Exit(0)
}()
log.Printf("groqqer gateway listening on http://127.0.0.1:%d", *port)
if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed {
log.Fatalf("server error: %v", err)
}
}