Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 87 additions & 0 deletions internal/aistudio/search_alias.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
package aistudio

import "strings"

const searchModelAliasSuffix = "-search"

// ResolveSearchModelAlias resolves the strict -search model suffix.
func ResolveSearchModelAlias(model string) (baseModel string, enableSearch bool) {
if !strings.HasSuffix(model, searchModelAliasSuffix) {
return model, false
}
baseModel = strings.TrimSuffix(model, searchModelAliasSuffix)
if baseModel == "" || baseModel == "models/" {
return model, false
}
return baseModel, true
}

// ResolveSearchGenerateRequest maps a Search alias to its real model and merges
// native Google Search into the request without replacing any existing tools.
func ResolveSearchGenerateRequest(request GenerateRequest) (GenerateRequest, bool) {
baseModel, enableSearch := ResolveSearchModelAlias(request.Model)
if !enableSearch {
return request, false
}
request.Model = baseModel
request.Tools = toolsWithWebSearch(request.Tools)
return request, true
}

func toolsWithWebSearch(tools Tools) Tools {
if tools.GoogleSearch != nil {
options := *tools.GoogleSearch
options.WebSearch = true
tools.GoogleSearch = &options
return tools
}
for _, name := range tools.Google {
if name == "google_search" {
return tools
}
}
tools.GoogleSearch = &GoogleSearchOptions{WebSearch: true}
return tools
}

// ModelsWithSearchAliases adds derived catalog entries for chat models that
// support both GenerateContent and native Google Search.
func ModelsWithSearchAliases(models []Model) []Model {
result := make([]Model, 0, len(models)*2)
knownIDs := make(map[string]struct{}, len(models)*2)
for _, model := range models {
knownIDs[model.ID] = struct{}{}
}
for _, model := range cloneModels(models) {
result = append(result, model)
if !searchAliasEligible(model) {
continue
}
aliasID := model.ID + searchModelAliasSuffix
if _, exists := knownIDs[aliasID]; exists {
continue
}
alias := cloneModels([]Model{model})[0]
alias.ID = aliasID
alias.Name = model.Name + " (Search)"
alias.Description = searchAliasDescription(model.Description)
result = append(result, alias)
knownIDs[aliasID] = struct{}{}
}
return result
}

func searchAliasEligible(model Model) bool {
return !strings.HasSuffix(model.ID, searchModelAliasSuffix) &&
model.Capabilities["chat_model"] &&
model.Capabilities["google_search"] &&
hasMethod(model, "generateContent")
}

func searchAliasDescription(description string) string {
const note = "Automatically enables Google Search."
if strings.TrimSpace(description) == "" {
return note
}
return description + " " + note
}
87 changes: 87 additions & 0 deletions internal/aistudio/search_alias_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
package aistudio

import (
"reflect"
"strings"
"testing"
)

func TestResolveSearchModelAliasStrictSuffix(t *testing.T) {
for _, test := range []struct {
model string
base string
search bool
}{
{model: "gemini-3.8-flash-search", base: "gemini-3.8-flash", search: true},
{model: "models/gemini-3.8-flash-search", base: "models/gemini-3.8-flash", search: true},
{model: "gemini-3.8-flash", base: "gemini-3.8-flash"},
{model: "gemini-search-preview", base: "gemini-search-preview"},
{model: "gemini-search-extra", base: "gemini-search-extra"},
{model: "gemini-search ", base: "gemini-search "},
} {
t.Run(test.model, func(t *testing.T) {
base, search := ResolveSearchModelAlias(test.model)
if base != test.base || search != test.search {
t.Fatalf("ResolveSearchModelAlias(%q) = (%q, %v), want (%q, %v)", test.model, base, search, test.base, test.search)
}
})
}
}

func TestSearchAliasUnsupportedModelFailsToolValidation(t *testing.T) {
request, enabled := ResolveSearchGenerateRequest(GenerateRequest{Model: "gemini-no-search-search"})
if !enabled || request.Model != "gemini-no-search" {
t.Fatalf("resolved request = %#v, enabled=%v", request, enabled)
}
err := validateRequestedTools(request.Tools, Model{
ID: "gemini-no-search", Capabilities: map[string]bool{"chat_model": true},
})
if err == nil || !strings.Contains(err.Error(), "google_search") {
t.Fatalf("validation error = %v, want google_search capability error", err)
}
}

func TestModelsWithSearchAliases(t *testing.T) {
searchCapabilities := map[string]bool{"chat_model": true, "google_search": true, "function_declarations": true}
models := []Model{
{
ID: "gemini-3.8-flash", Name: "Gemini 3.8 Flash", Description: "Fast model.",
Methods: []string{"generateContent", "countTokens"}, InputTokenLimit: 100, OutputTokenLimit: 20,
Capabilities: searchCapabilities, CapabilityOptions: map[string][]string{"aliases": {"latest"}},
AccessModes: []int64{1}, Paid: true,
},
{ID: "gemini-no-search", Name: "No Search", Methods: []string{"generateContent"}, Capabilities: map[string]bool{"chat_model": true}},
{ID: "speech-search-capable", Name: "Speech", Methods: []string{"generateContent"}, Capabilities: map[string]bool{"google_search": true, "speech_route": true}},
{ID: "video-search-capable", Name: "Video", Methods: []string{"predictLongRunning"}, Capabilities: map[string]bool{"google_search": true, "video_route": true}},
{ID: "live-search-capable", Name: "Live", Methods: []string{"bidiGenerateContent"}, Capabilities: map[string]bool{"chat_model": true, "google_search": true, "live_route": true}},
}

got := ModelsWithSearchAliases(models)
if len(got) != len(models)+1 {
t.Fatalf("model count = %d, want %d: %#v", len(got), len(models)+1, got)
}
base, alias := got[0], got[1]
if base.ID != "gemini-3.8-flash" || alias.ID != "gemini-3.8-flash-search" || alias.Name != "Gemini 3.8 Flash (Search)" {
t.Fatalf("base/alias = %#v / %#v", base, alias)
}
if alias.InputTokenLimit != base.InputTokenLimit || alias.OutputTokenLimit != base.OutputTokenLimit || alias.Paid != base.Paid ||
!reflect.DeepEqual(alias.Methods, base.Methods) || !reflect.DeepEqual(alias.Capabilities, base.Capabilities) ||
!reflect.DeepEqual(alias.CapabilityOptions, base.CapabilityOptions) || !reflect.DeepEqual(alias.AccessModes, base.AccessModes) {
t.Fatalf("alias did not inherit base metadata: base=%#v alias=%#v", base, alias)
}
if !strings.Contains(alias.Description, "Automatically enables Google Search") {
t.Fatalf("alias description = %q", alias.Description)
}
for _, model := range got {
if model.ID == "gemini-no-search-search" || model.ID == "speech-search-capable-search" ||
model.ID == "video-search-capable-search" || model.ID == "live-search-capable-search" {
t.Fatalf("ineligible alias exposed: %q", model.ID)
}
}

alias.Capabilities["mutated"] = true
alias.CapabilityOptions["aliases"][0] = "mutated"
if base.Capabilities["mutated"] || base.CapabilityOptions["aliases"][0] != "latest" {
t.Fatal("alias metadata shares mutable state with base model")
}
}
8 changes: 5 additions & 3 deletions internal/api/gemini.go
Original file line number Diff line number Diff line change
Expand Up @@ -197,6 +197,7 @@ func (s *server) handleGeminiModels(w http.ResponseWriter, r *http.Request) {
}
return
}
models = aistudio.ModelsWithSearchAliases(models)
data := make([]map[string]any, 0, len(models))
for _, model := range models {
data = append(data, geminiModelObject(model))
Expand All @@ -213,7 +214,7 @@ func (s *server) handleGeminiModel(w http.ResponseWriter, r *http.Request) {
}
return
}
for _, model := range models {
for _, model := range aistudio.ModelsWithSearchAliases(models) {
if model.ID == modelID {
writeJSON(w, http.StatusOK, geminiModelObject(model))
return
Expand Down Expand Up @@ -749,9 +750,10 @@ func geminiEmptyObjectPresent(raw json.RawMessage, field string) (bool, error) {
}

func (s *server) handleGeminiCountTokens(w http.ResponseWriter, r *http.Request, request aistudio.GenerateRequest) {
count, err := s.service.CountTokens(r.Context(), aistudio.TokenCountRequest{
tokenRequest := resolveSearchTokenCountRequest(aistudio.TokenCountRequest{
Model: request.Model, System: request.System, Contents: request.Contents, Tools: request.Tools,
})
count, err := s.service.CountTokens(r.Context(), tokenRequest)
if err != nil {
if shouldWriteRequestError(r, err) {
writeGeminiError(w, statusFromError(err), geminiErrorStatus(err), err.Error())
Expand All @@ -762,7 +764,7 @@ func (s *server) handleGeminiCountTokens(w http.ResponseWriter, r *http.Request,
}

func (s *server) handleGeminiGenerate(w http.ResponseWriter, r *http.Request, request aistudio.GenerateRequest, stream bool) {
events, err := s.service.Generate(r.Context(), request)
events, err := s.generate(r.Context(), request)
if err != nil {
if shouldWriteRequestError(r, err) {
writeGeminiError(w, statusFromError(err), geminiErrorStatus(err), err.Error())
Expand Down
3 changes: 2 additions & 1 deletion internal/api/openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -85,6 +85,7 @@ func (s *server) handleOpenAIModels(w http.ResponseWriter, r *http.Request) {
writeAnthropicModels(w, models)
return
}
models = aistudio.ModelsWithSearchAliases(models)
data := make([]map[string]any, 0, len(models))
for _, model := range models {
item := map[string]any{
Expand Down Expand Up @@ -131,7 +132,7 @@ func (s *server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
writeOpenAIError(w, http.StatusBadRequest, "invalid_request", err.Error())
return
}
events, err := s.service.Generate(r.Context(), generateRequest)
events, err := s.generate(r.Context(), generateRequest)
if err != nil {
if shouldWriteRequestError(r, err) {
writeOpenAIError(w, statusFromError(err), openAIErrorCode(err), err.Error())
Expand Down
2 changes: 1 addition & 1 deletion internal/api/responses.go
Original file line number Diff line number Diff line change
Expand Up @@ -163,7 +163,7 @@ func (s *server) handleResponses(w http.ResponseWriter, r *http.Request) {
instructions = append(instructions, inlineInstructions...)
generateRequest.System = strings.Join(instructions, "\n")
}
events, err := s.service.Generate(r.Context(), generateRequest)
events, err := s.generate(r.Context(), generateRequest)
if err != nil {
if shouldWriteRequestError(r, err) {
writeOpenAIError(w, statusFromError(err), openAIErrorCode(err), err.Error())
Expand Down
26 changes: 26 additions & 0 deletions internal/api/search_alias.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
package api

import (
"context"

"github.com/Mag1cFall/AIStudio2API/internal/aistudio"
)

// generate resolves provider-independent model aliases at the shared API
// boundary before account routing and model capability validation.
func (s *server) generate(ctx context.Context, request aistudio.GenerateRequest) (<-chan aistudio.Event, error) {
request, _ = aistudio.ResolveSearchGenerateRequest(request)
return s.service.Generate(ctx, request)
}

func resolveSearchTokenCountRequest(request aistudio.TokenCountRequest) aistudio.TokenCountRequest {
generate, enabled := aistudio.ResolveSearchGenerateRequest(aistudio.GenerateRequest{
Model: request.Model,
Tools: request.Tools,
})
if enabled {
request.Model = generate.Model
request.Tools = generate.Tools
}
return request
}
Loading