diff --git a/cmd/fetch_antigravity_models/main.go b/cmd/fetch_antigravity_models/main.go index 6e34eda1..7fd97b59 100644 --- a/cmd/fetch_antigravity_models/main.go +++ b/cmd/fetch_antigravity_models/main.go @@ -42,6 +42,7 @@ const ( antigravitySandboxBaseURLDaily = "https://daily-cloudcode-pa.sandbox.googleapis.com" antigravityBaseURLProd = "https://cloudcode-pa.googleapis.com" antigravityModelsPath = "/v1internal:fetchAvailableModels" + maxFetchAttemptsPerEndpoint = 2 ) func init() { @@ -189,15 +190,24 @@ func main() { fmt.Printf("Model list saved to: %s\n", outputPath) } +func defaultAntigravityFetchBaseURLs() []string { + return []string{antigravityBaseURLDaily, antigravityBaseURLProd, antigravitySandboxBaseURLDaily} +} + func fetchModels(ctx context.Context, auth *coreauth.Auth) []modelEntry { - accessToken := metaStringValue(auth.Metadata, "access_token") + return fetchModelsFromBaseURLs(ctx, auth, defaultAntigravityFetchBaseURLs(), nil) +} + +func fetchModelsFromBaseURLs(ctx context.Context, auth *coreauth.Auth, baseURLs []string, client *http.Client) []modelEntry { + var accessToken string + if auth != nil { + accessToken = metaStringValue(auth.Metadata, "access_token") + } if accessToken == "" { fmt.Fprintln(os.Stderr, "error: no access token found in auth") return nil } - baseURLs := []string{antigravityBaseURLProd, antigravityBaseURLDaily, antigravitySandboxBaseURLDaily} - for _, baseURL := range baseURLs { modelsURL := baseURL + antigravityModelsPath @@ -211,78 +221,87 @@ func fetchModels(ctx context.Context, auth *coreauth.Auth) []modelEntry { payload = []byte(`{}`) } - httpReq, errReq := http.NewRequestWithContext(ctx, http.MethodPost, modelsURL, strings.NewReader(string(payload))) - if errReq != nil { - continue - } - httpReq.Close = true - httpReq.Header.Set("Content-Type", "application/json") - httpReq.Header.Set("Authorization", "Bearer "+accessToken) - httpReq.Header.Set("User-Agent", misc.AntigravityUserAgent()) - - httpClient := &http.Client{Timeout: 30 * time.Second} - if transport, _, errProxy := proxyutil.BuildHTTPTransport(auth.ProxyURL); errProxy == nil && transport != nil { - httpClient.Transport = transport - } - httpResp, errDo := httpClient.Do(httpReq) - if errDo != nil { - continue - } - - bodyBytes, errRead := io.ReadAll(httpResp.Body) - httpResp.Body.Close() - if errRead != nil { - continue - } - - if httpResp.StatusCode < http.StatusOK || httpResp.StatusCode >= http.StatusMultipleChoices { - continue - } - - result := gjson.GetBytes(bodyBytes, "models") - if !result.Exists() { - continue - } - - var models []modelEntry - - for originalName, modelData := range result.Map() { - modelID := strings.TrimSpace(originalName) - if modelID == "" { + for attempt := 1; attempt <= maxFetchAttemptsPerEndpoint; attempt++ { + httpReq, errReq := http.NewRequestWithContext(ctx, http.MethodPost, modelsURL, strings.NewReader(string(payload))) + if errReq != nil { continue } - // Skip internal/experimental models - switch modelID { - case "chat_20706", "chat_23310", "tab_flash_lite_preview", "tab_jump_flash_lite_preview", "gemini-2.5-flash-thinking", "gemini-2.5-pro": + httpReq.Close = true + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("Authorization", "Bearer "+accessToken) + httpReq.Header.Set("User-Agent", misc.AntigravityUserAgent()) + + httpClient := client + if httpClient == nil { + httpClient = &http.Client{Timeout: 30 * time.Second} + if auth != nil { + if transport, _, errProxy := proxyutil.BuildHTTPTransport(auth.ProxyURL); errProxy == nil && transport != nil { + httpClient.Transport = transport + } + } + } + httpResp, errDo := httpClient.Do(httpReq) + if errDo != nil { continue } - displayName := modelData.Get("displayName").String() - if displayName == "" { - displayName = modelID + bodyBytes, errRead := io.ReadAll(httpResp.Body) + if errClose := httpResp.Body.Close(); errClose != nil { + log.Errorf("response body close error: %v", errClose) + } + if errRead != nil { + continue } - entry := modelEntry{ - ID: modelID, - Object: "model", - OwnedBy: "antigravity", - Type: "antigravity", - DisplayName: displayName, - Name: modelID, - Description: displayName, + if httpResp.StatusCode < http.StatusOK || httpResp.StatusCode >= http.StatusMultipleChoices { + continue } - if maxTok := modelData.Get("maxTokens").Int(); maxTok > 0 { - entry.ContextLength = int(maxTok) + result := gjson.GetBytes(bodyBytes, "models") + if !result.Exists() { + continue } - if maxOut := modelData.Get("maxOutputTokens").Int(); maxOut > 0 { - entry.MaxCompletionTokens = int(maxOut) + + var models []modelEntry + + for originalName, modelData := range result.Map() { + modelID := strings.TrimSpace(originalName) + if modelID == "" { + continue + } + // Skip internal/experimental models + switch modelID { + case "chat_20706", "chat_23310", "tab_flash_lite_preview", "tab_jump_flash_lite_preview", "gemini-2.5-flash-thinking", "gemini-2.5-pro": + continue + } + + displayName := modelData.Get("displayName").String() + if displayName == "" { + displayName = modelID + } + + entry := modelEntry{ + ID: modelID, + Object: "model", + OwnedBy: "antigravity", + Type: "antigravity", + DisplayName: displayName, + Name: modelID, + Description: displayName, + } + + if maxTok := modelData.Get("maxTokens").Int(); maxTok > 0 { + entry.ContextLength = int(maxTok) + } + if maxOut := modelData.Get("maxOutputTokens").Int(); maxOut > 0 { + entry.MaxCompletionTokens = int(maxOut) + } + + models = append(models, entry) } - models = append(models, entry) + return models } - - return models } return nil diff --git a/cmd/fetch_antigravity_models/main_test.go b/cmd/fetch_antigravity_models/main_test.go new file mode 100644 index 00000000..db1d2b9b --- /dev/null +++ b/cmd/fetch_antigravity_models/main_test.go @@ -0,0 +1,127 @@ +package main + +import ( + "context" + "net/http" + "net/http/httptest" + "reflect" + "sync/atomic" + "testing" + + coreauth "github.com/router-for-me/CLIProxyAPI/v7/sdk/cliproxy/auth" +) + +func TestDefaultAntigravityFetchBaseURLs(t *testing.T) { + want := []string{ + antigravityBaseURLDaily, + antigravityBaseURLProd, + antigravitySandboxBaseURLDaily, + } + + got := defaultAntigravityFetchBaseURLs() + if !reflect.DeepEqual(got, want) { + t.Fatalf("defaultAntigravityFetchBaseURLs() = %#v, want %#v", got, want) + } +} + +func TestFetchModelsRetryPerEndpoint(t *testing.T) { + var endpoint1Calls atomic.Int32 + var endpoint2Calls atomic.Int32 + + server1 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + call := endpoint1Calls.Add(1) + if call == 1 { + http.Error(w, `{"error":"temporary server error"}`, http.StatusInternalServerError) + return + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{ + "models": { + "gemini-3.6-flash": { + "displayName": "Gemini 3.6 Flash", + "maxTokens": 1048576, + "maxOutputTokens": 8192 + } + } + }`)) + })) + defer server1.Close() + + server2 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + endpoint2Calls.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{ + "models": { + "gemini-1.5-pro": { + "displayName": "Gemini 1.5 Pro" + } + } + }`)) + })) + defer server2.Close() + + auth := &coreauth.Auth{ + Metadata: map[string]interface{}{ + "access_token": "test-token", + "project_id": "test-project", + }, + } + + // Case 1: First endpoint fails on attempt 1, succeeds on attempt 2. + // It should succeed without calling endpoint 2, and endpoint 1 should have been called 2 times. + models := fetchModelsFromBaseURLs(context.Background(), auth, []string{server1.URL, server2.URL}, server1.Client()) + if len(models) != 1 || models[0].ID != "gemini-3.6-flash" { + t.Fatalf("expected 1 model (gemini-3.6-flash), got: %#v", models) + } + if calls := endpoint1Calls.Load(); calls != 2 { + t.Fatalf("expected endpoint 1 to be called 2 times, got %d", calls) + } + if calls := endpoint2Calls.Load(); calls != 0 { + t.Fatalf("expected endpoint 2 not to be called, got %d", calls) + } +} + +func TestFetchModelsFallbackAfterTwoAttempts(t *testing.T) { + var endpoint1Calls atomic.Int32 + var endpoint2Calls atomic.Int32 + + server1 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + endpoint1Calls.Add(1) + http.Error(w, `{"error":"unavailable"}`, http.StatusServiceUnavailable) + })) + defer server1.Close() + + server2 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + endpoint2Calls.Add(1) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{ + "models": { + "gemini-2.5-flash": { + "displayName": "Gemini 2.5 Flash" + } + } + }`)) + })) + defer server2.Close() + + auth := &coreauth.Auth{ + Metadata: map[string]interface{}{ + "access_token": "test-token", + }, + } + + // Case 2: First endpoint fails all 2 attempts, then falls back to endpoint 2. + models := fetchModelsFromBaseURLs(context.Background(), auth, []string{server1.URL, server2.URL}, server1.Client()) + if len(models) != 1 || models[0].ID != "gemini-2.5-flash" { + t.Fatalf("expected 1 model (gemini-2.5-flash), got: %#v", models) + } + if calls := endpoint1Calls.Load(); calls != 2 { + t.Fatalf("expected endpoint 1 to be called 2 times before fallback, got %d", calls) + } + if calls := endpoint2Calls.Load(); calls != 1 { + t.Fatalf("expected endpoint 2 to be called 1 time, got %d", calls) + } +}