Skip to content

Commit aa64095

Browse files
committed
fix: repair Grok Gateway WebSocket handshake
1 parent 97bb4fc commit aa64095

15 files changed

Lines changed: 859 additions & 15 deletions

‎.github/workflows/build.yml‎

Lines changed: 19 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -2,9 +2,14 @@ name: Build and Release
22

33
on:
44
workflow_dispatch:
5+
inputs:
6+
version:
7+
description: "Release version without the v prefix"
8+
required: true
9+
type: string
510
push:
6-
paths:
7-
- VERSION
11+
tags:
12+
- "v*"
813

914
permissions:
1015
contents: write
@@ -16,13 +21,21 @@ jobs:
1621
version: ${{ steps.rev.outputs.version }}
1722
tag: ${{ steps.rev.outputs.tag }}
1823
steps:
19-
- name: Checkout
20-
uses: actions/checkout@v4
21-
- name: Read version
24+
- name: Resolve version from tag
2225
id: rev
2326
shell: bash
27+
env:
28+
INPUT_VERSION: ${{ inputs.version }}
2429
run: |
25-
version="$(tr -d '[:space:]' < VERSION)"
30+
if [[ -n "${INPUT_VERSION}" ]]; then
31+
version="${INPUT_VERSION#v}"
32+
else
33+
version="${GITHUB_REF_NAME#v}"
34+
fi
35+
if [[ -z "${version}" || "${version}" == "main" ]]; then
36+
echo "A release tag or workflow_dispatch version is required" >&2
37+
exit 1
38+
fi
2639
echo "version=${version}" >> "${GITHUB_OUTPUT}"
2740
echo "tag=v${version}" >> "${GITHUB_OUTPUT}"
2841

‎.github/workflows/build_docker.yml‎

Lines changed: 19 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,9 +4,16 @@ on:
44
push:
55
branches:
66
- main
7+
tags:
8+
- "v*"
79
paths-ignore:
810
- README.md
911
workflow_dispatch:
12+
inputs:
13+
version:
14+
description: "Image version without the v prefix"
15+
required: true
16+
type: string
1017

1118
permissions:
1219
contents: read
@@ -22,11 +29,19 @@ jobs:
2229
- name: Checkout
2330
uses: actions/checkout@v4
2431

25-
- name: Read version
32+
- name: Resolve version
2633
id: version
2734
shell: bash
35+
env:
36+
INPUT_VERSION: ${{ inputs.version }}
2837
run: |
29-
version="$(tr -d '[:space:]' < VERSION)"
38+
if [[ -n "${INPUT_VERSION}" ]]; then
39+
version="${INPUT_VERSION#v}"
40+
elif [[ "${GITHUB_REF_TYPE}" == "tag" ]]; then
41+
version="${GITHUB_REF_NAME#v}"
42+
else
43+
version="${GITHUB_SHA::12}"
44+
fi
3045
echo "version=${version}" >> "${GITHUB_OUTPUT}"
3146
echo "tag=v${version}" >> "${GITHUB_OUTPUT}"
3247
echo "sha_short=${GITHUB_SHA::12}" >> "${GITHUB_OUTPUT}"
@@ -52,6 +67,8 @@ jobs:
5267
platforms: linux/amd64,linux/arm64,linux/arm/v7
5368
push: true
5469
provenance: false
70+
build-args: |
71+
VERSION=${{ steps.version.outputs.version }}
5572
tags: |
5673
${{ env.GHCR_REPO }}:latest
5774
${{ env.GHCR_REPO }}:${{ steps.version.outputs.version }}

‎Dockerfile‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,11 @@
11
FROM --platform=$BUILDPLATFORM golang:1.26-alpine AS builder
22
ARG TARGETOS TARGETARCH
3+
ARG VERSION=dev
34
WORKDIR /app
45
COPY go.mod go.sum ./
56
RUN go mod download
67
COPY . .
7-
RUN CGO_ENABLED=0 GOOS=${TARGETOS} GOARCH=${TARGETARCH} go build -trimpath -ldflags="-s -w -X main.Version=$(cat VERSION)" -o grok2api .
8+
RUN CGO_ENABLED=0 GOOS=${TARGETOS} GOARCH=${TARGETARCH} go build -trimpath -ldflags="-s -w -X main.Version=${VERSION}" -o grok2api .
89

910
FROM alpine:3.21
1011
RUN apk add --no-cache ca-certificates tzdata

‎VERSION‎

Lines changed: 0 additions & 1 deletion
This file was deleted.

‎internal/api/gateway_chat.go‎

Lines changed: 94 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,94 @@
1+
package api
2+
3+
import (
4+
"context"
5+
"encoding/json"
6+
"net/http"
7+
"strings"
8+
"time"
9+
10+
"github.com/google/uuid"
11+
12+
"github.com/aurora-develop/grok2api/internal/account"
13+
"github.com/aurora-develop/grok2api/internal/grok"
14+
"github.com/aurora-develop/grok2api/internal/model"
15+
)
16+
17+
func usesGatewayChat(modelName string) bool {
18+
return modelName == "grok-4.6"
19+
}
20+
21+
func gatewayUpstreamModel(modelName string) string {
22+
if modelName == "grok-4.6" {
23+
return model.ModeExpert.ApiStr()
24+
}
25+
return modelName
26+
}
27+
28+
func (s *Server) runGatewayChatOnce(w http.ResponseWriter, r *http.Request, lease *account.Lease, message string, fileInputs []string, emitThink, stream bool, modelName string) error {
29+
ctx, cancel := context.WithTimeout(r.Context(), 5*time.Minute)
30+
defer cancel()
31+
32+
completionID := "chatcmpl-" + uuid.NewString()
33+
created := time.Now().Unix()
34+
events := s.Transport.StreamGatewayChat(ctx, lease.Token, gatewayUpstreamModel(modelName), message, fileInputs)
35+
var text, reasoning strings.Builder
36+
var conversationID, responseID string
37+
38+
if stream {
39+
sw := newSSEWriter(w)
40+
sw.writeComment("heartbeat")
41+
for ev := range events {
42+
if ev.Err != nil {
43+
sw.writeOpenAIError(ev.Err.Error(), "upstream_error", "", "")
44+
return nil
45+
}
46+
conversationID = ev.ConversationID
47+
responseID = ev.ResponseID
48+
switch ev.Kind {
49+
case grok.EventText:
50+
text.WriteString(ev.Content)
51+
sw.writeJSONData(makeStreamChunk(completionID, created, modelName, ev.Content, "", false))
52+
case grok.EventThinking:
53+
if emitThink {
54+
reasoning.WriteString(ev.Content)
55+
chunk := makeStreamChunk(completionID, created, modelName, "", ev.Content, false)
56+
chunk["choices"].([]any)[0].(map[string]any)["delta"] = map[string]any{"reasoning_content": ev.Content}
57+
sw.writeJSONData(chunk)
58+
}
59+
case grok.EventSoftStop:
60+
sw.writeJSONData(makeStreamChunk(completionID, created, modelName, "", "", true))
61+
}
62+
}
63+
sw.writeDone()
64+
if s.ConvTracker != nil && conversationID != "" && responseID != "" {
65+
s.ConvTracker.Set(lease.Token, conversationID, responseID)
66+
}
67+
return nil
68+
}
69+
70+
for ev := range events {
71+
if ev.Err != nil {
72+
return ev.Err
73+
}
74+
conversationID = ev.ConversationID
75+
responseID = ev.ResponseID
76+
switch ev.Kind {
77+
case grok.EventText:
78+
text.WriteString(ev.Content)
79+
case grok.EventThinking:
80+
if emitThink {
81+
reasoning.WriteString(ev.Content)
82+
}
83+
}
84+
}
85+
resp := makeChatResponse(completionID, created, modelName, text.String(), reasoning.String(), emitThink)
86+
body, _ := json.Marshal(resp)
87+
w.Header().Set("Content-Type", "application/json; charset=utf-8")
88+
w.WriteHeader(http.StatusOK)
89+
_, _ = w.Write(body)
90+
if s.ConvTracker != nil && conversationID != "" && responseID != "" {
91+
s.ConvTracker.Set(lease.Token, conversationID, responseID)
92+
}
93+
return nil
94+
}

‎internal/api/openai_chat.go‎

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -146,12 +146,15 @@ func (s *Server) runGrokChatWithRetry(c *gin.Context, req *chatCompletionRequest
146146
// runGrokChatOnce executes one chat attempt against grok.com.
147147
// For follow-up messages it uses /responses endpoint with conversation tracking.
148148
func (s *Server) runGrokChatOnce(w http.ResponseWriter, r *http.Request, lease *account.Lease, spec *model.Spec, message string, fileInputs []string, temp, topP float64, emitThink, stream bool, modelName string) error {
149+
if usesGatewayChat(modelName) {
150+
return s.runGatewayChatOnce(w, r, lease, message, fileInputs, emitThink, stream, modelName)
151+
}
149152
// Check if we have an active conversation for this token.
150153
convCtx := s.ConvTracker.Get(lease.Token)
151154
var (
152-
payload map[string]any
153-
targetURL string
154-
isNew bool
155+
payload map[string]any
156+
targetURL string
157+
isNew bool
155158
)
156159
if convCtx != nil && convCtx.ConversationID != "" && convCtx.LastResponseID != "" {
157160
// Follow-up message in existing conversation.
Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,18 @@
1+
package api
2+
3+
import "testing"
4+
5+
func TestGrok46UsesGatewayChat(t *testing.T) {
6+
if !usesGatewayChat("grok-4.6") {
7+
t.Fatal("grok-4.6 should use the WebSocket Gateway")
8+
}
9+
if usesGatewayChat("grok-4.20-auto") {
10+
t.Fatal("legacy grok model should keep the HTTP chat path")
11+
}
12+
}
13+
14+
func TestGatewayUpstreamModel(t *testing.T) {
15+
if got := gatewayUpstreamModel("grok-4.6"); got != "expert" {
16+
t.Fatalf("upstream model = %q", got)
17+
}
18+
}

0 commit comments

Comments
 (0)