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