diff --git a/backend/internal/handler/gateway_web_search.go b/backend/internal/handler/gateway_web_search.go index a0a249f25..741cc322f 100644 --- a/backend/internal/handler/gateway_web_search.go +++ b/backend/internal/handler/gateway_web_search.go @@ -28,12 +28,8 @@ const ( ) func (h *GatewayHandler) WebSearch(c *gin.Context) { - type webSearchReq struct { - Query string `json:"query" binding:"required"` - MaxResults int `json:"max_results"` - } - - var req webSearchReq + isXSearch := c.GetBool("grok_x_search_endpoint") + var req grokStandaloneSearchRequest if err := c.ShouldBindJSON(&req); err != nil { c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{ "type": "invalid_request_error", @@ -41,7 +37,28 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { }}) return } - req.MaxResults = normalizeGrokWebSearchMaxResults(req.MaxResults) + query := strings.TrimSpace(req.Query) + if query == "" { + query = strings.TrimSpace(req.Input) + } + if query == "" { + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{ + "type": "invalid_request_error", + "message": "query is required", + }}) + return + } + req.Query = query + maxResults := 0 + if req.MaxResults != nil { + maxResults = *req.MaxResults + } + maxResults = normalizeGrokWebSearchMaxResults(maxResults) + searchModel := resolveGrokStandaloneSearchModel() + searchLabel := "web_search" + if isXSearch { + searchLabel = "x_search" + } apiKey, ok := middleware2.GetAPIKeyFromContext(c) if !ok || apiKey == nil { @@ -55,7 +72,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { if apiKey.Group == nil || apiKey.Group.Platform != "grok" { c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{ "type": "invalid_request_error", - "message": "web search is only supported for grok groups", + "message": searchLabel + " is only supported for grok groups", }}) return } @@ -79,7 +96,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { "role": "user", "content": req.Query, }}, }) - if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, xai.DefaultTextModel, auditBody); decision != nil && !decision.AllowNextStage { + if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, searchModel, auditBody); decision != nil && !decision.AllowNextStage { status := decision.HTTPStatus if status == 0 { status = http.StatusForbidden @@ -123,7 +140,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { // First attempt + up to 3 failover accounts (max 4 total). for attempt := 0; attempt < 4; attempt++ { selected, selectErr := h.gatewayService.SelectAccountWithLoadAwareness( - c.Request.Context(), groupID, "", xai.DefaultTextModel, failedAccounts, "", 0, + c.Request.Context(), groupID, "", searchModel, failedAccounts, "", 0, ) if selectErr != nil { if attempt == 0 { @@ -159,7 +176,11 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { account = selected.Account accountReleaseFunc = release - nativeResp, providerName, err = h.doGrokNativeWebSearch(c.Request.Context(), c, account, req.Query, req.MaxResults) + if isXSearch { + nativeResp, providerName, err = h.doGrokNativeXSearch(c.Request.Context(), c, account, req, searchModel, maxResults) + } else { + nativeResp, providerName, err = h.doGrokNativeWebSearch(c.Request.Context(), c, account, req.Query, maxResults, searchModel) + } if err == nil { break } @@ -198,7 +219,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) // Request IDs are billing idempotency keys, so they must be unique per invocation. // Query/IP/UA hashes would collapse repeated identical searches into one charge. - searchRequestID := "web_search:" + uuid.NewString() + searchRequestID := searchLabel + ":" + uuid.NewString() if apiKey.Group != nil { if p := apiKey.Group.GetSearchPricePer1k(); p != nil && *p == 0 { logger.L().With( @@ -211,7 +232,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ Result: &service.ForwardResult{ RequestID: searchRequestID, - Model: "grok-web-search", + Model: "grok-" + strings.ReplaceAll(searchLabel, "_", "-"), SearchCount: 1, Duration: 0, }, @@ -240,7 +261,7 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) { "query": req.Query, "results": nativeResp.Results, "provider": providerName, - "max_results": req.MaxResults, + "max_results": maxResults, }) } @@ -299,13 +320,13 @@ func (h *GatewayHandler) acquireWebSearchAccountSlot( // doGrokNativeWebSearch executes web search using the Grok account's native capability // by calling the responses endpoint with web_search tool, then normalizes sources to unified format. -func (h *GatewayHandler) doGrokNativeWebSearch(ctx context.Context, c *gin.Context, account *service.Account, query string, maxResults int) (*websearch.SearchResponse, string, error) { +func (h *GatewayHandler) doGrokNativeWebSearch(ctx context.Context, c *gin.Context, account *service.Account, query string, maxResults int, model string) (*websearch.SearchResponse, string, error) { maxResults = normalizeGrokWebSearchMaxResults(maxResults) // Build a minimal responses request that triggers Grok web search tool. // Ask for structured metadata because xAI action.sources commonly contains URLs only. searchBody := map[string]any{ - "model": xai.DefaultTextModel, + "model": xai.ResolveDefaultTextModel(model), "input": buildGrokWebSearchPrompt(query, maxResults), "tools": []map[string]any{{"type": "web_search"}}, "include": []string{"web_search_call.action.sources"}, @@ -329,6 +350,23 @@ func (h *GatewayHandler) doGrokNativeWebSearch(ctx context.Context, c *gin.Conte }, "grok-native", nil } +func (h *GatewayHandler) doGrokNativeXSearch(ctx context.Context, c *gin.Context, account *service.Account, req grokStandaloneSearchRequest, model string, maxResults int) (*websearch.SearchResponse, string, error) { + maxResults = normalizeGrokWebSearchMaxResults(maxResults) + bodyBytes, err := buildGrokXSearchResponsesBody(req, model) + if err != nil { + return nil, "", err + } + respBytes, err := h.gatewayService.DoGrokNativeResponsesJSON(ctx, account, bodyBytes) + if err != nil { + return nil, "", err + } + results := extractGrokWebSearchSources(respBytes, maxResults) + return &websearch.SearchResponse{ + Results: results, + Query: req.Query, + }, "grok-native", nil +} + func normalizeGrokWebSearchMaxResults(maxResults int) int { if maxResults <= 0 { return defaultGrokWebSearchResults @@ -377,7 +415,8 @@ func extractGrokWebSearchSources(body []byte, maxResults int) []websearch.Search output := gjson.GetBytes(body, "output") output.ForEach(func(_, item gjson.Result) bool { - if item.Get("type").String() == "web_search_call" { + callType := item.Get("type").String() + if callType == "web_search_call" || callType == "x_search_call" { sources := item.Get("action.sources") if sources.IsArray() { sources.ForEach(func(_, src gjson.Result) bool { diff --git a/backend/internal/handler/openai_x_search.go b/backend/internal/handler/openai_x_search.go new file mode 100644 index 000000000..2c4e9dcda --- /dev/null +++ b/backend/internal/handler/openai_x_search.go @@ -0,0 +1,66 @@ +package handler + +import ( + "encoding/json" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/gin-gonic/gin" +) + +type grokStandaloneSearchRequest struct { + Query string `json:"query"` + Input string `json:"input"` + MaxResults *int `json:"max_results"` + AllowedXHandles []string `json:"allowed_x_handles"` + ExcludedXHandles []string `json:"excluded_x_handles"` + FromDate string `json:"from_date"` + ToDate string `json:"to_date"` + EnableImageUnderstanding *bool `json:"enable_image_understanding"` + EnableVideoUnderstanding *bool `json:"enable_video_understanding"` +} + +// XSearch marks the standalone endpoint so WebSearch can use native x_search +// while retaining its dedicated per-call billing contract. +func (h *GatewayHandler) XSearch(c *gin.Context) { + c.Set("grok_x_search_endpoint", true) + h.WebSearch(c) +} + +func resolveGrokStandaloneSearchModel() string { + return xai.ResolveDefaultTextModel(xai.RuntimeModelMappingOptions().DefaultText) +} + +func buildGrokXSearchResponsesBody(req grokStandaloneSearchRequest, model string) ([]byte, error) { + input := strings.TrimSpace(req.Query) + if input == "" { + input = strings.TrimSpace(req.Input) + } + tool := map[string]any{"type": "x_search"} + if len(req.AllowedXHandles) > 0 { + tool["allowed_x_handles"] = req.AllowedXHandles + } + if len(req.ExcludedXHandles) > 0 { + tool["excluded_x_handles"] = req.ExcludedXHandles + } + if strings.TrimSpace(req.FromDate) != "" { + tool["from_date"] = strings.TrimSpace(req.FromDate) + } + if strings.TrimSpace(req.ToDate) != "" { + tool["to_date"] = strings.TrimSpace(req.ToDate) + } + if req.EnableImageUnderstanding != nil { + tool["enable_image_understanding"] = *req.EnableImageUnderstanding + } + if req.EnableVideoUnderstanding != nil { + tool["enable_video_understanding"] = *req.EnableVideoUnderstanding + } + return json.Marshal(map[string]any{ + "model": xai.ResolveDefaultTextModel(model), + "input": input, + "tools": []map[string]any{tool}, + "tool_choice": "required", + "store": false, + "stream": false, + }) +} diff --git a/backend/internal/handler/openai_x_search_test.go b/backend/internal/handler/openai_x_search_test.go new file mode 100644 index 000000000..d5cabdba8 --- /dev/null +++ b/backend/internal/handler/openai_x_search_test.go @@ -0,0 +1,56 @@ +package handler + +import ( + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestBuildGrokXSearchResponsesBody(t *testing.T) { + t.Parallel() + understandImages := true + understandVideos := false + body, err := buildGrokXSearchResponsesBody(grokStandaloneSearchRequest{ + Query: "latest posts from xAI", + AllowedXHandles: []string{"xai"}, + ExcludedXHandles: []string{"spam"}, + FromDate: "2026-08-01", + ToDate: "2026-08-10", + EnableImageUnderstanding: &understandImages, + EnableVideoUnderstanding: &understandVideos, + }, xai.DefaultTextModel) + require.NoError(t, err) + require.Equal(t, xai.DefaultTextModel, gjson.GetBytes(body, "model").String()) + require.Equal(t, "latest posts from xAI", gjson.GetBytes(body, "input").String()) + require.Equal(t, "required", gjson.GetBytes(body, "tool_choice").String()) + require.Equal(t, "x_search", gjson.GetBytes(body, "tools.0.type").String()) + require.Equal(t, "xai", gjson.GetBytes(body, "tools.0.allowed_x_handles.0").String()) + require.Equal(t, "spam", gjson.GetBytes(body, "tools.0.excluded_x_handles.0").String()) + require.Equal(t, "2026-08-01", gjson.GetBytes(body, "tools.0.from_date").String()) + require.Equal(t, "2026-08-10", gjson.GetBytes(body, "tools.0.to_date").String()) + require.True(t, gjson.GetBytes(body, "tools.0.enable_image_understanding").Bool()) + require.False(t, gjson.GetBytes(body, "tools.0.enable_video_understanding").Bool()) + require.False(t, gjson.GetBytes(body, "store").Bool()) + require.False(t, gjson.GetBytes(body, "stream").Bool()) +} + +func TestBuildGrokXSearchResponsesBodyAcceptsInputAlias(t *testing.T) { + t.Parallel() + body, err := buildGrokXSearchResponsesBody(grokStandaloneSearchRequest{Input: "latest posts from xAI"}, xai.DefaultTextModel) + require.NoError(t, err) + require.Equal(t, "latest posts from xAI", gjson.GetBytes(body, "input").String()) +} + +func TestResolveGrokStandaloneSearchModelUsesRuntimeDefault(t *testing.T) { + original := xai.RuntimeModelMappingOptions() + t.Cleanup(func() { xai.SetRuntimeModelMappingOptions(original) }) + xai.SetRuntimeModelMappingOptions(xai.ModelMappingOptions{DefaultText: "grok-4.6"}) + + model := resolveGrokStandaloneSearchModel() + body, err := buildGrokXSearchResponsesBody(grokStandaloneSearchRequest{Query: "latest posts from xAI"}, model) + require.NoError(t, err) + require.Equal(t, "grok-4.6", model) + require.Equal(t, model, gjson.GetBytes(body, "model").String()) +} diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index 8d5fcd6a5..6d99d2c7e 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -315,6 +315,14 @@ func RegisterGatewayRoutes( } h.Gateway.WebSearch(c) }) + gateway.POST("/x_search", 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": "X Search API is not supported for this platform"}}) + return + } + h.Gateway.XSearch(c) + }) } // Gemini 原生 API 兼容层(Gemini SDK/CLI 直连) @@ -443,6 +451,14 @@ func RegisterGatewayRoutes( } h.Gateway.WebSearch(c) }) + r.POST("/x_search", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, 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": "X Search API is not supported for this platform"}}) + return + } + h.Gateway.XSearch(c) + }) // Antigravity 模型列表 r.GET("/antigravity/models", gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.Gateway.AntigravityModels) diff --git a/backend/internal/server/routes/prompt_audit_route_coverage_test.go b/backend/internal/server/routes/prompt_audit_route_coverage_test.go index f03cf5b31..af4272cb8 100644 --- a/backend/internal/server/routes/prompt_audit_route_coverage_test.go +++ b/backend/internal/server/routes/prompt_audit_route_coverage_test.go @@ -48,6 +48,7 @@ func TestEveryGatewayPOSTRouteIsClassifiedForPromptAuditCoverage(t *testing.T) { "/models/*modelAction": {"gemini_v1beta_handler.go"}, "/tts": {"grok_audio.go"}, "/web_search": {"gateway_web_search.go"}, + "/x_search": {"gateway_web_search.go"}, } excluded := map[string]string{ "/messages/count_tokens": "tokenization only; it does not execute a model request",