mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:18:18 +08:00
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:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user