From fe91bb286489b8fe724268a07f4b70a3267c8315 Mon Sep 17 00:00:00 2001 From: Josef Strzibny Date: Thu, 8 Oct 2026 13:00:03 +0200 Subject: [PATCH] Add -timeout flag to increase client timeout --- README.md | 8 ++++ pkg/api/client.go | 55 ++++++++++++++++++++++-- pkg/api/client_test.go | 97 ++++++++++++++++++++++++++++++++++++++++++ pkg/cmd/account.go | 8 ++-- pkg/cmd/archive.go | 7 ++- pkg/cmd/locations.go | 7 ++- pkg/cmd/login.go | 11 +++-- pkg/cmd/root.go | 46 ++++++++++++++++++-- pkg/cmd/root_test.go | 53 +++++++++++++++++++++++ pkg/cmd/search.go | 8 ++-- 10 files changed, 280 insertions(+), 20 deletions(-) create mode 100644 pkg/api/client_test.go create mode 100644 pkg/cmd/root_test.go diff --git a/README.md b/README.md index 53590cf..5826397 100644 --- a/README.md +++ b/README.md @@ -119,6 +119,14 @@ serpapi login ``` Quote the expression with single quotes in bash/zsh to avoid shell interpretation of `$`, `|`, and `"`. - `--api-key ` — Override API key (takes priority over environment and config file) +- `--timeout ` — HTTP request timeout (default `60`, env: `SERPAPI_TIMEOUT`). Use `0` to wait indefinitely. Slow engines such as `google_ai_mode`, or searches with `no_cache=true`, can take longer than the default; raise this if you see a `network_error` mentioning `Client.Timeout exceeded`. Note that a timed-out search may still complete on SerpApi's side and be billed and retrievable with `serpapi archive `. + ```bash + # Give an AI Mode query up to two minutes + serpapi search --timeout 120 engine=google_ai_mode q="..." + + # Or set it once for the shell session + export SERPAPI_TIMEOUT=120 + ``` ## Configuration diff --git a/pkg/api/client.go b/pkg/api/client.go index 1c9654d..9a6d7ab 100644 --- a/pkg/api/client.go +++ b/pkg/api/client.go @@ -4,8 +4,10 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "io" + "net" "net/http" "net/url" "time" @@ -15,24 +17,39 @@ import ( ) const ( - defaultTimeout = 30 * time.Second - maxResponseBytes = 100 << 20 // 100 MB + // DefaultTimeout bounds a single HTTP request. Slow engines such as + // google_ai_mode can take well over 30s with no_cache=true, so this + // matches the 60s default used by the official serpapi-golang library. + DefaultTimeout = 60 * time.Second + maxResponseBytes = 100 << 20 // 100 MB ) // Client is an HTTP client for the SerpApi service. type Client struct { apiKey string baseURL string + timeout time.Duration http *http.Client } -// New creates a new SerpApi client. apiKey may be empty for unauthenticated requests. +// New creates a new SerpApi client with DefaultTimeout. apiKey may be empty +// for unauthenticated requests. func New(apiKey string) *Client { + return NewWithTimeout(apiKey, DefaultTimeout) +} + +// NewWithTimeout creates a new SerpApi client whose requests are bounded by +// timeout. A zero or negative timeout disables the limit entirely. +func NewWithTimeout(apiKey string, timeout time.Duration) *Client { + if timeout < 0 { + timeout = 0 + } return &Client{ apiKey: apiKey, baseURL: "https://serpapi.com", + timeout: timeout, http: &http.Client{ - Timeout: defaultTimeout, + Timeout: timeout, CheckRedirect: func(req *http.Request, via []*http.Request) error { return http.ErrUseLastResponse }, @@ -64,6 +81,9 @@ func (c *Client) doGet(ctx context.Context, endpoint string, params map[string]s resp, err := c.http.Do(req) if err != nil { + if c.isTimeout(req, err) { + return nil, &clierrors.NetworkError{Message: c.timeoutMessage(), Cause: err} + } return nil, &clierrors.NetworkError{Message: err.Error(), Cause: err} } defer resp.Body.Close() @@ -104,6 +124,33 @@ func (c *Client) doGet(ctx context.Context, endpoint string, params map[string]s return body, nil } +// isTimeout reports whether err was caused by the client's own timeout rather +// than by the caller cancelling the request context. +func (c *Client) isTimeout(req *http.Request, err error) bool { + if c.timeout <= 0 { + return false + } + if ctxErr := req.Context().Err(); ctxErr != nil { + // The caller's context ended first (Ctrl-C, parent deadline, ...). + return false + } + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + return true + } + return errors.Is(err, context.DeadlineExceeded) +} + +func (c *Client) timeoutMessage() string { + return fmt.Sprintf( + "Request timed out after %gs waiting for SerpApi to respond. "+ + "Slow engines (e.g. google_ai_mode) or no_cache=true can take longer; "+ + "retry with --timeout (0 to wait indefinitely). "+ + "The search may still have completed server-side and be available in your SerpApi search archive.", + c.timeout.Seconds(), + ) +} + // checkAPIError returns an APIError if the JSON body contains a top-level "error" key. func checkAPIError(body []byte) error { var envelope struct { diff --git a/pkg/api/client_test.go b/pkg/api/client_test.go new file mode 100644 index 0000000..646c9dc --- /dev/null +++ b/pkg/api/client_test.go @@ -0,0 +1,97 @@ +package api + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + clierrors "github.com/serpapi/serpapi-cli/pkg/errors" +) + +// newSlowServer returns a client pointed at a server that waits for `delay` +// (or until the request is cancelled) before answering with a tiny JSON body. +func newSlowServer(t *testing.T, delay time.Duration, timeout time.Duration) *Client { + t.Helper() + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + select { + case <-time.After(delay): + case <-r.Context().Done(): + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"ok":true}`)) + })) + t.Cleanup(srv.Close) + + client := NewWithTimeout("secret_key", timeout) + client.baseURL = srv.URL + return client +} + +func TestNewUsesDefaultTimeout(t *testing.T) { + if got := New("k").http.Timeout; got != DefaultTimeout { + t.Fatalf("expected default timeout %s, got %s", DefaultTimeout, got) + } + if DefaultTimeout < 60*time.Second { + t.Fatalf("default timeout %s is too short for slow engines like google_ai_mode", DefaultTimeout) + } +} + +func TestTimeoutProducesActionableNetworkError(t *testing.T) { + client := newSlowServer(t, 5*time.Second, 50*time.Millisecond) + + _, err := client.Search(context.Background(), map[string]string{"q": "x"}) + var netErr *clierrors.NetworkError + if !errors.As(err, &netErr) { + t.Fatalf("expected NetworkError, got %T: %v", err, err) + } + for _, want := range []string{"timed out", "--timeout", "archive"} { + if !strings.Contains(netErr.Message, want) { + t.Errorf("message %q should mention %q", netErr.Message, want) + } + } + if strings.Contains(netErr.Message, "api_key=") { + t.Errorf("timeout message should not echo the request URL: %q", netErr.Message) + } +} + +func TestZeroTimeoutDisablesLimit(t *testing.T) { + client := newSlowServer(t, 200*time.Millisecond, 0) + if client.http.Timeout != 0 { + t.Fatalf("expected no http timeout, got %s", client.http.Timeout) + } + + if _, err := client.Search(context.Background(), map[string]string{"q": "x"}); err != nil { + t.Fatalf("expected slow request to succeed with timeout disabled, got %v", err) + } +} + +func TestLongerTimeoutAllowsSlowResponse(t *testing.T) { + client := newSlowServer(t, 200*time.Millisecond, 5*time.Second) + if _, err := client.Search(context.Background(), map[string]string{"q": "x"}); err != nil { + t.Fatalf("expected slow request to succeed within timeout, got %v", err) + } +} + +func TestCallerCancellationIsNotReportedAsTimeout(t *testing.T) { + client := newSlowServer(t, 5*time.Second, 10*time.Second) + + ctx, cancel := context.WithCancel(context.Background()) + go func() { + time.Sleep(50 * time.Millisecond) + cancel() + }() + + _, err := client.Search(ctx, map[string]string{"q": "x"}) + var netErr *clierrors.NetworkError + if !errors.As(err, &netErr) { + t.Fatalf("expected NetworkError, got %T: %v", err, err) + } + if strings.Contains(netErr.Message, "timed out") { + t.Errorf("caller cancellation should not be reported as a timeout: %q", netErr.Message) + } +} diff --git a/pkg/cmd/account.go b/pkg/cmd/account.go index c55db92..9aa4843 100644 --- a/pkg/cmd/account.go +++ b/pkg/cmd/account.go @@ -5,8 +5,6 @@ import ( "regexp" "github.com/spf13/cobra" - - "github.com/serpapi/serpapi-cli/pkg/api" ) var accountCmd = &cobra.Command{ @@ -26,10 +24,14 @@ func runAccount(cmd *cobra.Command, args []string) error { return err } + client, err := newClient(apiKey) + if err != nil { + return err + } + sp := newSpinner("Fetching account...") sp.Start() defer sp.Stop() - client := api.New(apiKey) result, err := client.Account(cmd.Context()) if err != nil { return err diff --git a/pkg/cmd/archive.go b/pkg/cmd/archive.go index dae295a..01497d4 100644 --- a/pkg/cmd/archive.go +++ b/pkg/cmd/archive.go @@ -5,7 +5,6 @@ import ( "github.com/spf13/cobra" - "github.com/serpapi/serpapi-cli/pkg/api" clierrors "github.com/serpapi/serpapi-cli/pkg/errors" ) @@ -33,10 +32,14 @@ func runArchive(cmd *cobra.Command, args []string) error { return err } + client, err := newClient(apiKey) + if err != nil { + return err + } + sp := newSpinner("Fetching archive...") sp.Start() defer sp.Stop() - client := api.New(apiKey) result, err := client.Archive(cmd.Context(), id) if err != nil { return err diff --git a/pkg/cmd/locations.go b/pkg/cmd/locations.go index 95b2aa4..741e554 100644 --- a/pkg/cmd/locations.go +++ b/pkg/cmd/locations.go @@ -3,7 +3,6 @@ package cmd import ( "github.com/spf13/cobra" - "github.com/serpapi/serpapi-cli/pkg/api" "github.com/serpapi/serpapi-cli/pkg/params" ) @@ -25,10 +24,14 @@ func runLocations(cmd *cobra.Command, args []string) error { } paramsMap := params.ParamsToMap(parsed) + client, err := newClient("") + if err != nil { + return err + } + sp := newSpinner("Fetching locations...") sp.Start() defer sp.Stop() - client := api.New("") result, err := client.Locations(cmd.Context(), paramsMap) if err != nil { return err diff --git a/pkg/cmd/login.go b/pkg/cmd/login.go index b311097..39d6e3f 100644 --- a/pkg/cmd/login.go +++ b/pkg/cmd/login.go @@ -11,7 +11,6 @@ import ( "github.com/spf13/cobra" "golang.org/x/term" - "github.com/serpapi/serpapi-cli/pkg/api" "github.com/serpapi/serpapi-cli/pkg/config" clierrors "github.com/serpapi/serpapi-cli/pkg/errors" ) @@ -30,7 +29,10 @@ func init() { func runLogin(cmd *cobra.Command, args []string) error { // Check if already authenticated if existingKey, ok := config.LoadAPIKey(); ok { - client := api.New(existingKey) + client, err := newClient(existingKey) + if err != nil { + return err + } raw, err := client.Account(cmd.Context()) if err == nil { var account struct { @@ -80,7 +82,10 @@ func runLogin(cmd *cobra.Command, args []string) error { return &clierrors.UsageError{Message: "API key cannot be empty."} } - client := api.New(apiKey) + client, err := newClient(apiKey) + if err != nil { + return err + } raw, err := client.Account(cmd.Context()) if err != nil { return err diff --git a/pkg/cmd/root.go b/pkg/cmd/root.go index 31b3e0a..41eb197 100644 --- a/pkg/cmd/root.go +++ b/pkg/cmd/root.go @@ -5,9 +5,12 @@ import ( "encoding/json" "fmt" "os" + "strconv" + "time" "github.com/spf13/cobra" + "github.com/serpapi/serpapi-cli/pkg/api" "github.com/serpapi/serpapi-cli/pkg/config" clierrors "github.com/serpapi/serpapi-cli/pkg/errors" "github.com/serpapi/serpapi-cli/pkg/jq" @@ -17,9 +20,10 @@ import ( ) var ( - apiKeyFlag string - fieldsFlag string - jqFlag string + apiKeyFlag string + fieldsFlag string + jqFlag string + timeoutFlag string ) var rootCmd = &cobra.Command{ @@ -49,6 +53,8 @@ func init() { rootCmd.PersistentFlags().StringVar(&apiKeyFlag, "api-key", "", "SerpApi API key (env: SERPAPI_KEY)") rootCmd.PersistentFlags().StringVar(&fieldsFlag, "fields", "", "Restrict JSON fields (server-side json_restrictor)") rootCmd.PersistentFlags().StringVar(&jqFlag, "jq", "", "Apply jq filter to output") + rootCmd.PersistentFlags().StringVar(&timeoutFlag, "timeout", "", + fmt.Sprintf("HTTP request timeout in seconds, 0 to disable (default %d; env: SERPAPI_TIMEOUT)", int(api.DefaultTimeout/time.Second))) // Wrap cobra flag-parsing errors as UsageError so they get exit code 2. rootCmd.SetFlagErrorFunc(func(_ *cobra.Command, err error) error { @@ -80,6 +86,40 @@ func resolveAPIKeyOptional() string { return "" } +// resolveTimeout merges --timeout with SERPAPI_TIMEOUT, falling back to +// api.DefaultTimeout. Values are whole or fractional seconds; 0 disables. +func resolveTimeout() (time.Duration, error) { + raw := timeoutFlag + source := "--timeout" + if raw == "" { + raw = os.Getenv("SERPAPI_TIMEOUT") + source = "SERPAPI_TIMEOUT" + } + if raw == "" { + return api.DefaultTimeout, nil + } + return parseTimeout(raw, source) +} + +func parseTimeout(raw, source string) (time.Duration, error) { + secs, err := strconv.ParseFloat(raw, 64) + if err != nil || secs < 0 || secs != secs { // secs != secs rejects NaN + return 0, &clierrors.UsageError{ + Message: fmt.Sprintf("Invalid %s value %q: expected a non-negative number of seconds", source, raw), + } + } + return time.Duration(secs * float64(time.Second)), nil +} + +// newClient builds an API client honoring the configured request timeout. +func newClient(apiKey string) (*api.Client, error) { + timeout, err := resolveTimeout() + if err != nil { + return nil, err + } + return api.NewWithTimeout(apiKey, timeout), nil +} + // handleOutput applies --jq filtering and prints result. func handleOutput(raw json.RawMessage) error { if jqFlag != "" { diff --git a/pkg/cmd/root_test.go b/pkg/cmd/root_test.go new file mode 100644 index 0000000..471af0a --- /dev/null +++ b/pkg/cmd/root_test.go @@ -0,0 +1,53 @@ +package cmd + +import ( + "errors" + "testing" + "time" + + "github.com/serpapi/serpapi-cli/pkg/api" + clierrors "github.com/serpapi/serpapi-cli/pkg/errors" +) + +func TestResolveTimeout(t *testing.T) { + tests := []struct { + name string + flag string + env string + want time.Duration + wantErr bool + }{ + {name: "default", want: api.DefaultTimeout}, + {name: "flag seconds", flag: "120", want: 120 * time.Second}, + {name: "flag fractional", flag: "1.5", want: 1500 * time.Millisecond}, + {name: "flag zero disables", flag: "0", want: 0}, + {name: "env fallback", env: "90", want: 90 * time.Second}, + {name: "flag overrides env", flag: "10", env: "90", want: 10 * time.Second}, + {name: "negative rejected", flag: "-5", wantErr: true}, + {name: "garbage rejected", flag: "soon", wantErr: true}, + {name: "env garbage rejected", env: "abc", wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + timeoutFlag = tt.flag + t.Cleanup(func() { timeoutFlag = "" }) + t.Setenv("SERPAPI_TIMEOUT", tt.env) + + got, err := resolveTimeout() + if tt.wantErr { + var ue *clierrors.UsageError + if !errors.As(err, &ue) { + t.Fatalf("expected UsageError, got %T: %v", err, err) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != tt.want { + t.Errorf("expected %s, got %s", tt.want, got) + } + }) + } +} diff --git a/pkg/cmd/search.go b/pkg/cmd/search.go index b42216f..2470a3f 100644 --- a/pkg/cmd/search.go +++ b/pkg/cmd/search.go @@ -10,7 +10,6 @@ import ( "github.com/spf13/cobra" - "github.com/serpapi/serpapi-cli/pkg/api" clierrors "github.com/serpapi/serpapi-cli/pkg/errors" "github.com/serpapi/serpapi-cli/pkg/params" ) @@ -53,6 +52,11 @@ func runSearch(cmd *cobra.Command, args []string) error { paramsMap := params.ParamsToMap(parsed) params.ApplyFields(paramsMap, fieldsFlag) + client, err := newClient(apiKey) + if err != nil { + return err + } + hasMaxPages := cmd.Flags().Changed("max-pages") if hasMaxPages && !allPagesFlag { @@ -63,7 +67,6 @@ func runSearch(cmd *cobra.Command, args []string) error { sp := newSpinner("Searching...") sp.Start() defer sp.Stop() - client := api.New(apiKey) raw, err := client.Search(cmd.Context(), paramsMap) if err != nil { return err @@ -82,7 +85,6 @@ func runSearch(cmd *cobra.Command, args []string) error { paramsMap["json_restrictor"] = paginationRestrictor } - client := api.New(apiKey) currentParams := paramsMap var accumulated map[string]any visitedPages := make(map[string]bool)