diff --git a/internal/aistudio/search_alias.go b/internal/aistudio/search_alias.go new file mode 100644 index 0000000..7cb5aec --- /dev/null +++ b/internal/aistudio/search_alias.go @@ -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 +} diff --git a/internal/aistudio/search_alias_test.go b/internal/aistudio/search_alias_test.go new file mode 100644 index 0000000..49db072 --- /dev/null +++ b/internal/aistudio/search_alias_test.go @@ -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") + } +} diff --git a/internal/api/gemini.go b/internal/api/gemini.go index f26e446..7a18571 100644 --- a/internal/api/gemini.go +++ b/internal/api/gemini.go @@ -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)) @@ -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 @@ -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()) @@ -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()) diff --git a/internal/api/openai.go b/internal/api/openai.go index 89eccbb..18fb9a5 100644 --- a/internal/api/openai.go +++ b/internal/api/openai.go @@ -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{ @@ -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()) diff --git a/internal/api/responses.go b/internal/api/responses.go index 783841d..840600f 100644 --- a/internal/api/responses.go +++ b/internal/api/responses.go @@ -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()) diff --git a/internal/api/search_alias.go b/internal/api/search_alias.go new file mode 100644 index 0000000..4d70404 --- /dev/null +++ b/internal/api/search_alias.go @@ -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 +} diff --git a/internal/api/search_alias_test.go b/internal/api/search_alias_test.go new file mode 100644 index 0000000..352f2a8 --- /dev/null +++ b/internal/api/search_alias_test.go @@ -0,0 +1,200 @@ +package api + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "reflect" + "testing" + + "github.com/Mag1cFall/AIStudio2API/internal/aistudio" +) + +func TestGeminiSearchAliasInjectsGoogleSearch(t *testing.T) { + var request geminiRequest + if err := json.Unmarshal([]byte(`{"contents":[{"parts":[{"text":"hello"}]}]}`), &request); err != nil { + t.Fatal(err) + } + generateRequest, err := request.toGenerateRequest("test-id", "gemini-3.8-flash-search") + if err != nil { + t.Fatal(err) + } + service := &searchAliasCaptureService{} + if _, err := (&server{service: service}).generate(context.Background(), generateRequest); err != nil { + t.Fatal(err) + } + assertResolvedSearchAlias(t, service.generated, true) +} + +func TestOpenAISearchAliasInjectsGoogleSearch(t *testing.T) { + var request chatRequest + if err := json.Unmarshal([]byte(`{"model":"gemini-3.8-flash-search","messages":[{"role":"user","content":"hello"}]}`), &request); err != nil { + t.Fatal(err) + } + generateRequest, err := request.toGenerateRequest("test-id") + if err != nil { + t.Fatal(err) + } + service := &searchAliasCaptureService{} + if _, err := (&server{service: service}).generate(context.Background(), generateRequest); err != nil { + t.Fatal(err) + } + assertResolvedSearchAlias(t, service.generated, true) +} + +func TestRegularModelDoesNotInjectGoogleSearch(t *testing.T) { + request := aistudio.GenerateRequest{Model: "gemini-3.8-flash"} + resolved, enabled := aistudio.ResolveSearchGenerateRequest(request) + if enabled || resolved.Model != request.Model || resolved.Tools.GoogleSearch != nil { + t.Fatalf("regular model changed: %#v, enabled=%v", resolved, enabled) + } +} + +func TestSearchAliasPreservesFunctionDeclarations(t *testing.T) { + function := aistudio.FunctionDeclaration{Name: "lookup", Parameters: json.RawMessage(`{"type":"object"}`)} + request := aistudio.GenerateRequest{ + Model: "gemini-3.8-flash-search", + Tools: aistudio.Tools{Functions: []aistudio.FunctionDeclaration{function}}, + } + resolved, enabled := aistudio.ResolveSearchGenerateRequest(request) + assertResolvedSearchAlias(t, resolved, enabled) + if !reflect.DeepEqual(resolved.Tools.Functions, []aistudio.FunctionDeclaration{function}) { + t.Fatalf("function declarations changed: %#v", resolved.Tools.Functions) + } +} + +func TestSearchAliasDoesNotDuplicateExistingGoogleSearch(t *testing.T) { + request := aistudio.GenerateRequest{ + Model: "gemini-3.8-flash-search", + Contents: []aistudio.Content{{Role: aistudio.RoleUser, Parts: []aistudio.Part{{Text: "hello"}}}}, + Tools: aistudio.Tools{GoogleSearch: &aistudio.GoogleSearchOptions{WebSearch: true}}, + } + resolved, enabled := aistudio.ResolveSearchGenerateRequest(request) + assertResolvedSearchAlias(t, resolved, enabled) + wire := encodeSearchAliasGenerateRequest(t, resolved) + tools, ok := wire[6].([]any) + if !ok || len(tools) != 1 { + t.Fatalf("wire search tools = %#v, want exactly one", wire[6]) + } +} + +func TestSearchAliasEncodeGenerateContentUsesBaseModelAndExistingSearchWire(t *testing.T) { + request, enabled := aistudio.ResolveSearchGenerateRequest(aistudio.GenerateRequest{ + Model: "gemini-3.8-flash-search", + Contents: []aistudio.Content{{Role: aistudio.RoleUser, Parts: []aistudio.Part{{Text: "hello"}}}}, + }) + assertResolvedSearchAlias(t, request, enabled) + wire := encodeSearchAliasGenerateRequest(t, request) + if wire[0] != "models/gemini-3.8-flash" { + t.Fatalf("wire model = %#v", wire[0]) + } + wantSearch := []any{nil, nil, nil, []any{nil, []any{[]any{}}}} + tools, ok := wire[6].([]any) + if !ok || len(tools) != 1 || !reflect.DeepEqual(tools[0], wantSearch) { + t.Fatalf("wire tools = %#v, want %#v", wire[6], []any{wantSearch}) + } +} + +func TestGeminiStreamGenerateContentResolvesSearchAlias(t *testing.T) { + service := &searchAliasCaptureService{} + s := &server{service: service} + request := aistudio.GenerateRequest{ + ID: "response-id", Model: "gemini-3.8-flash-search", + Contents: []aistudio.Content{{Role: aistudio.RoleUser, Parts: []aistudio.Part{{Text: "hello"}}}}, + } + recorder := httptest.NewRecorder() + httpRequest := httptest.NewRequest("POST", "/v1beta/models/gemini-3.8-flash-search:streamGenerateContent", nil) + s.handleGeminiGenerate(recorder, httpRequest, request, true) + assertResolvedSearchAlias(t, service.generated, true) +} + +func TestOpenAIModelsExposeEligibleSearchAlias(t *testing.T) { + service := &searchAliasCaptureService{models: []aistudio.Model{ + { + ID: "gemini-3.8-flash", Name: "Gemini 3.8 Flash", Methods: []string{"generateContent"}, + Capabilities: map[string]bool{"chat_model": true, "google_search": true}, + }, + { + ID: "gemini-no-search", Name: "Gemini No Search", Methods: []string{"generateContent"}, + Capabilities: map[string]bool{"chat_model": true}, + }, + }} + s := &server{service: service} + recorder := httptest.NewRecorder() + s.handleOpenAIModels(recorder, httptest.NewRequest(http.MethodGet, "/v1/models", nil)) + if recorder.Code != http.StatusOK { + t.Fatalf("status = %d, body=%s", recorder.Code, recorder.Body.String()) + } + var response struct { + Data []struct { + ID string `json:"id"` + } `json:"data"` + } + if err := json.Unmarshal(recorder.Body.Bytes(), &response); err != nil { + t.Fatal(err) + } + ids := make(map[string]bool, len(response.Data)) + for _, model := range response.Data { + ids[model.ID] = true + } + for _, id := range []string{"gemini-3.8-flash", "gemini-3.8-flash-search", "gemini-no-search"} { + if !ids[id] { + t.Errorf("missing model %q: %#v", id, ids) + } + } + if ids["gemini-no-search-search"] { + t.Fatalf("unsupported Search alias exposed: %#v", ids) + } +} + +func assertResolvedSearchAlias(t *testing.T, request aistudio.GenerateRequest, enabled bool) { + t.Helper() + if !enabled { + t.Fatal("Search alias was not recognized") + } + if request.Model != "gemini-3.8-flash" { + t.Fatalf("resolved model = %q", request.Model) + } + if request.Tools.GoogleSearch == nil || !request.Tools.GoogleSearch.WebSearch { + t.Fatalf("Google Search not injected: %#v", request.Tools) + } +} + +type searchAliasCaptureService struct { + generated aistudio.GenerateRequest + models []aistudio.Model +} + +func (service *searchAliasCaptureService) Models(context.Context) ([]aistudio.Model, error) { + return service.models, nil +} + +func (service *searchAliasCaptureService) CountTokens(context.Context, aistudio.TokenCountRequest) (aistudio.TokenCount, error) { + return aistudio.TokenCount{}, nil +} + +func (service *searchAliasCaptureService) Generate(_ context.Context, request aistudio.GenerateRequest) (<-chan aistudio.Event, error) { + service.generated = request + events := make(chan aistudio.Event, 1) + events <- aistudio.Event{Kind: aistudio.EventFinish, FinishReason: "stop"} + close(events) + return events, nil +} + +func encodeSearchAliasGenerateRequest(t *testing.T, request aistudio.GenerateRequest) []any { + t.Helper() + body, err := aistudio.EncodeGenerateContentRequest( + request, + aistudio.GenerationDefaults{MaxOutputTokens: 1024}, + aistudio.RequestContext{}, + ) + if err != nil { + t.Fatal(err) + } + var wire []any + if err := json.Unmarshal(body, &wire); err != nil { + t.Fatal(err) + } + return wire +}