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.
This commit is contained in:
IanShaw027
2026-08-08 08:48:38 +08:00
parent 85b65284ec
commit 35faaa6d21
6 changed files with 111 additions and 8 deletions
+7 -1
View File
@@ -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)
+34
View File
@@ -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)
@@ -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)
+26 -6
View File
@@ -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
@@ -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")
+10 -1
View File
@@ -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 {