feat: 新增独立 /x_search,走原生 x_search 并沿用搜索计费

Grok 分组原先只有 /web_search,无法带 handle/日期过滤,也无法走
xAI 的 x_search tool。

- POST /x_search(仅 Grok 分组),复用 web_search 的审计、failover 与按次计费
- 上游 Responses 强制 x_search;计费模型记为 grok-x-search
This commit is contained in:
IanShaw027
2026-08-13 07:49:19 +08:00
parent 363cc4994b
commit 0de6d7e9ba
5 changed files with 195 additions and 17 deletions
+56 -17
View File
@@ -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 {
@@ -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,
})
}
@@ -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())
}
+16
View File
@@ -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)
@@ -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",