From 35faaa6d2100ae7ed94982d29c1e23882c709563 Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Sat, 8 Aug 2026 08:48:38 +0800 Subject: [PATCH] feat(grok): register custom-voices CRUD and audio download gateway routes Forward list/get/patch/delete and reference-audio paths with safe path segment encoding, method passthrough, and empty-body GET/DELETE handling. --- backend/internal/handler/grok_audio.go | 8 ++++- backend/internal/server/routes/gateway.go | 34 +++++++++++++++++++ .../internal/server/routes/gateway_test.go | 24 +++++++++++++ backend/internal/service/grok_audio.go | 32 +++++++++++++---- backend/internal/service/grok_audio_test.go | 10 ++++++ backend/internal/service/grok_upstream_url.go | 11 +++++- 6 files changed, 111 insertions(+), 8 deletions(-) diff --git a/backend/internal/handler/grok_audio.go b/backend/internal/handler/grok_audio.go index 4a76b0d900..cd1b1da912 100644 --- a/backend/internal/handler/grok_audio.go +++ b/backend/internal/handler/grok_audio.go @@ -279,7 +279,13 @@ func (h *OpenAIGatewayHandler) recordGrokVoiceUsage( } func readGrokVoiceGatewayBody(c *gin.Context) ([]byte, error) { - if c == nil || c.Request == nil || c.Request.Body == nil { + if c == nil || c.Request == nil { + return nil, errors.New("request body is required") + } + if c.Request.Body == nil { + if c.Request.Method == http.MethodGet || c.Request.Method == http.MethodDelete { + return nil, nil + } return nil, errors.New("request body is required") } return io.ReadAll(c.Request.Body) diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index fd8ce8d5c3..809f55d015 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -286,6 +286,23 @@ func RegisterGatewayRoutes( gateway.POST("/tts", voiceHandler("tts")) gateway.POST("/stt", voiceHandler("stt")) gateway.POST("/custom-voices", voiceHandler("custom-voices")) + customVoicePathHandler := func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Voice API is not supported for this platform"}}) + return + } + endpoint := "custom-voices/" + c.Param("voice_id") + if strings.HasSuffix(strings.TrimRight(c.Request.URL.Path, "/"), "/audio") { + endpoint += "/audio" + } + h.OpenAIGateway.GrokVoice(c, endpoint) + } + gateway.GET("/custom-voices", voiceHandler("custom-voices")) + gateway.GET("/custom-voices/:voice_id/audio", customVoicePathHandler) + gateway.GET("/custom-voices/:voice_id", customVoicePathHandler) + gateway.PATCH("/custom-voices/:voice_id", customVoicePathHandler) + gateway.DELETE("/custom-voices/:voice_id", customVoicePathHandler) gateway.GET("/realtime", func(c *gin.Context) { if getGroupPlatform(c) != service.PlatformGrok { service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) @@ -401,6 +418,23 @@ func RegisterGatewayRoutes( r.POST("/tts", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootVoiceHandler("tts")) r.POST("/stt", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootVoiceHandler("stt")) r.POST("/custom-voices", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootVoiceHandler("custom-voices")) + rootCustomVoicePathHandler := func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Voice API is not supported for this platform"}}) + return + } + endpoint := "custom-voices/" + c.Param("voice_id") + if strings.HasSuffix(strings.TrimRight(c.Request.URL.Path, "/"), "/audio") { + endpoint += "/audio" + } + h.OpenAIGateway.GrokVoice(c, endpoint) + } + r.GET("/custom-voices", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootVoiceHandler("custom-voices")) + r.GET("/custom-voices/:voice_id/audio", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootCustomVoicePathHandler) + r.GET("/custom-voices/:voice_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootCustomVoicePathHandler) + r.PATCH("/custom-voices/:voice_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootCustomVoicePathHandler) + r.DELETE("/custom-voices/:voice_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootCustomVoicePathHandler) r.GET("/realtime", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, func(c *gin.Context) { if getGroupPlatform(c) != service.PlatformGrok { service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) diff --git a/backend/internal/server/routes/gateway_test.go b/backend/internal/server/routes/gateway_test.go index 0fefb73adf..1c5ba37cb5 100644 --- a/backend/internal/server/routes/gateway_test.go +++ b/backend/internal/server/routes/gateway_test.go @@ -196,6 +196,30 @@ func TestGatewayRoutesGrokImagesAndVideosPathsAreRegistered(t *testing.T) { } } +func TestGatewayRoutesGrokCustomVoiceCRUDPathsAreRegistered(t *testing.T) { + router := newGatewayRoutesTestRouter(service.PlatformGrok) + registered := make(map[string]bool) + for _, route := range router.Routes() { + registered[route.Method+" "+route.Path] = true + } + for _, route := range []string{ + "POST /v1/custom-voices", + "GET /v1/custom-voices", + "GET /v1/custom-voices/:voice_id", + "PATCH /v1/custom-voices/:voice_id", + "DELETE /v1/custom-voices/:voice_id", + "GET /v1/custom-voices/:voice_id/audio", + "POST /custom-voices", + "GET /custom-voices", + "GET /custom-voices/:voice_id", + "PATCH /custom-voices/:voice_id", + "DELETE /custom-voices/:voice_id", + "GET /custom-voices/:voice_id/audio", + } { + require.True(t, registered[route], "%s should be registered", route) + } +} + func TestGatewayRoutesCompositeVideoLookupsUseGrokHandler(t *testing.T) { router := newGatewayRoutesTestRouter(service.PlatformComposite) diff --git a/backend/internal/service/grok_audio.go b/backend/internal/service/grok_audio.go index 2a9fa9030d..310a01fe36 100644 --- a/backend/internal/service/grok_audio.go +++ b/backend/internal/service/grok_audio.go @@ -22,7 +22,8 @@ var supportedGrokVoiceHTTPEndpoints = map[string]struct{}{ "custom-voices": {}, } -// ForwardGrokVoice forwards the official xAI Voice HTTP APIs (/tts, /stt, /custom-voices). +// ForwardGrokVoice forwards the official xAI Voice HTTP APIs (/tts, /stt, and +// the custom-voices CRUD/audio subresources). // The response is intentionally passed through because TTS returns audio bytes // while STT returns JSON and xAI may add format-specific headers. func (s *OpenAIGatewayService) ForwardGrokVoice(ctx context.Context, c *gin.Context, account *Account, endpoint string, body []byte, contentType string) (*OpenAIForwardResult, error) { @@ -33,9 +34,24 @@ func (s *OpenAIGatewayService) ForwardGrokVoice(ctx context.Context, c *gin.Cont return nil, fmt.Errorf("account platform %s is not supported for grok voice", account.Platform) } endpoint = strings.Trim(strings.TrimSpace(endpoint), "/") - if _, ok := supportedGrokVoiceHTTPEndpoints[endpoint]; !ok { + parts := strings.Split(endpoint, "/") + baseEndpoint := parts[0] + if _, ok := supportedGrokVoiceHTTPEndpoints[baseEndpoint]; !ok { return nil, fmt.Errorf("unsupported grok voice endpoint: %s", endpoint) } + if len(parts) > 1 && baseEndpoint != "custom-voices" { + return nil, fmt.Errorf("unsupported grok voice endpoint: %s", endpoint) + } + if baseEndpoint == "custom-voices" { + if len(parts) > 3 || (len(parts) == 3 && parts[2] != "audio") { + return nil, fmt.Errorf("unsupported grok voice endpoint: %s", endpoint) + } + } + for _, part := range parts[1:] { + if part == "" || part == "." || part == ".." || strings.ContainsAny(part, "?#\\") { + return nil, fmt.Errorf("invalid grok voice endpoint path") + } + } token, _, err := s.getRequestCredential(ctx, c, account) if err != nil { return nil, err @@ -46,7 +62,11 @@ func (s *OpenAIGatewayService) ForwardGrokVoice(ctx context.Context, c *gin.Cont } upstreamCtx, release := detachUpstreamContext(ctx) defer release() - req, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(body)) + method := http.MethodPost + if c != nil && c.Request != nil && strings.TrimSpace(c.Request.Method) != "" { + method = c.Request.Method + } + req, err := http.NewRequestWithContext(upstreamCtx, method, targetURL, bytes.NewReader(body)) if err != nil { return nil, err } @@ -81,11 +101,11 @@ func (s *OpenAIGatewayService) ForwardGrokVoice(ctx context.Context, c *gin.Cont return nil, err } writeGrokMediaResponse(c, resp, data, s.responseHeaderFilter) - audioUsage := estimateGrokVoiceAudioUsage(endpoint, body, contentType, data, time.Since(started)) + audioUsage := estimateGrokVoiceAudioUsage(baseEndpoint, body, contentType, data, time.Since(started)) return &OpenAIForwardResult{ RequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")), - Model: endpoint, - UpstreamModel: endpoint, + Model: baseEndpoint, + UpstreamModel: baseEndpoint, Duration: time.Since(started), AudioUsage: audioUsage, }, nil diff --git a/backend/internal/service/grok_audio_test.go b/backend/internal/service/grok_audio_test.go index 6647ccdff3..57a60431ad 100644 --- a/backend/internal/service/grok_audio_test.go +++ b/backend/internal/service/grok_audio_test.go @@ -42,6 +42,16 @@ func TestBuildGrokVoiceURL_RequiresEndpoint(t *testing.T) { require.Error(t, err) } +func TestBuildGrokVoiceURL_EncodesCustomVoicePathSegments(t *testing.T) { + account := &Account{Platform: PlatformGrok, Type: AccountTypeOAuth} + got, err := buildGrokVoiceURL(account, nil, "custom-voices/nlbqfwie/audio") + require.NoError(t, err) + require.Equal(t, xai.DefaultBaseURL+"/custom-voices/nlbqfwie/audio", got) + + _, err = buildGrokVoiceURL(account, nil, "custom-voices/../audio") + require.Error(t, err) +} + func TestForwardGrokVoice_RejectsNonGrok(t *testing.T) { svc := &OpenAIGatewayService{} _, err := svc.ForwardGrokVoice(context.Background(), nil, &Account{Platform: PlatformOpenAI}, "tts", []byte(`{}`), "application/json") diff --git a/backend/internal/service/grok_upstream_url.go b/backend/internal/service/grok_upstream_url.go index 7d6b495198..9574b3adbe 100644 --- a/backend/internal/service/grok_upstream_url.go +++ b/backend/internal/service/grok_upstream_url.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "net/url" "strings" "github.com/Wei-Shaw/sub2api/internal/config" @@ -164,7 +165,15 @@ func buildGrokVoiceURL(account *Account, cfg *config.Config, endpoint string) (s if ep == "" { return "", fmt.Errorf("voice endpoint is required") } - return strings.TrimRight(validated, "/") + "/" + ep, nil + parts := strings.Split(ep, "/") + encoded := make([]string, 0, len(parts)) + for _, part := range parts { + if strings.TrimSpace(part) == "" || part == "." || part == ".." { + return "", fmt.Errorf("invalid voice endpoint path") + } + encoded = append(encoded, url.PathEscape(part)) + } + return strings.TrimRight(validated, "/") + "/" + strings.Join(encoded, "/"), nil } func isGrokCLIProxyBaseURL(raw string) bool {