diff --git a/.gitignore b/.gitignore index e96c274..bc833d7 100644 --- a/.gitignore +++ b/.gitignore @@ -25,3 +25,5 @@ logs.txt .release core/test/ .aider* +.gocache/ +openserp diff --git a/cmd/root.go b/cmd/root.go index 0f589ca..2f46bde 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -20,6 +20,9 @@ const ( type Config struct { App AppConfig `mapstructure:"app"` + Resilience ResilienceConfig `mapstructure:"resilience"` + CircuitBreaker CircuitBreakerConfig `mapstructure:"circuit_breaker"` + CORS CORSConfig `mapstructure:"cors"` Config2Capcha Config2Captcha `mapstructure:"2captcha"` GoogleConfig core.SearchEngineOptions `mapstructure:"google"` YandexConfig core.SearchEngineOptions `mapstructure:"yandex"` @@ -49,14 +52,38 @@ type AppConfig struct { IsStealth bool `mapstructure:"stealth"` } +type ResilienceConfig struct { + MaxRetries int `mapstructure:"max_retries"` + AllowEndpointFallback bool `mapstructure:"allow_endpoint_fallback"` +} + +type CircuitBreakerConfig struct { + Failures int `mapstructure:"failures"` + RecoverySeconds int `mapstructure:"recovery_seconds"` + Successes int `mapstructure:"successes"` +} + +type CORSConfig struct { + Enabled bool `mapstructure:"enabled"` + AllowOrigins string `mapstructure:"allow_origins"` + AllowMethods string `mapstructure:"allow_methods"` + AllowHeaders string `mapstructure:"allow_headers"` + MaxAge int `mapstructure:"max_age"` +} + var config = Config{} var flagToConfigKey = map[string]string{ - "config": "app.config_path", - "browser-path": "app.browser_path", - "leave": "app.leave_head", - "raw": "app.raw_requests", - "2captcha_key": "2captcha.apikey", + "config": "app.config_path", + "browser-path": "app.browser_path", + "leave": "app.leave_head", + "raw": "app.raw_requests", + "2captcha_key": "2captcha.apikey", + "max_retries": "resilience.max_retries", + "allow_endpoint_fallback": "resilience.allow_endpoint_fallback", + "cb_failures": "circuit_breaker.failures", + "cb_recovery": "circuit_breaker.recovery_seconds", + "cb_successes": "circuit_breaker.successes", } var RootCmd = &cobra.Command{ @@ -123,6 +150,7 @@ func parseFlagValue(flg *pflag.Flag) (interface{}, error) { // Initialize Viper func initializeConfig(cmd *cobra.Command) error { v := viper.New() + setConfigDefaults(v) // Base name of the config file, without the file extension v.SetConfigName(defaultConfigFilename) @@ -160,6 +188,20 @@ func initializeConfig(cmd *cobra.Command) error { return nil } +func setConfigDefaults(v *viper.Viper) { + // Keep stage2 defaults stable even when config file is absent. + v.SetDefault("resilience.max_retries", 3) + v.SetDefault("resilience.allow_endpoint_fallback", false) + v.SetDefault("circuit_breaker.failures", 5) + v.SetDefault("circuit_breaker.recovery_seconds", 60) + v.SetDefault("circuit_breaker.successes", 2) + v.SetDefault("cors.enabled", true) + v.SetDefault("cors.allow_origins", "*") + v.SetDefault("cors.allow_methods", "GET, POST, OPTIONS") + v.SetDefault("cors.allow_headers", "Origin, Content-Type, Accept, Authorization") + v.SetDefault("cors.max_age", 86400) +} + func init() { RootCmd.PersistentFlags().IntVarP(&config.App.Port, "port", "p", 7070, "Port number to run server") RootCmd.PersistentFlags().StringVarP(&config.App.Host, "host", "a", "127.0.0.1", "Host address to run server") @@ -176,4 +218,9 @@ func init() { RootCmd.PersistentFlags().StringVarP(&config.App.ProxyURL, "proxy", "x", "", "HTTP or Socks5 proxy URL (e.g. http://user:pass@127.0.0.1:8080)") RootCmd.PersistentFlags().BoolVarP(&config.App.IsStealth, "stealth", "s", false, "Use stealth browser plugin") RootCmd.PersistentFlags().BoolVarP(&config.App.Insecure, "insecure", "k", false, "Allow insecure TLS connections") + RootCmd.PersistentFlags().IntVar(&config.Resilience.MaxRetries, "max_retries", 3, "Max retry attempts per search engine (0 to disable)") + RootCmd.PersistentFlags().BoolVar(&config.Resilience.AllowEndpointFallback, "allow_endpoint_fallback", false, "Allow dedicated endpoints to fallback to other engines") + RootCmd.PersistentFlags().IntVar(&config.CircuitBreaker.Failures, "cb_failures", 5, "Consecutive failures before circuit breaker opens") + RootCmd.PersistentFlags().IntVar(&config.CircuitBreaker.RecoverySeconds, "cb_recovery", 60, "Seconds before retrying an engine with open circuit") + RootCmd.PersistentFlags().IntVar(&config.CircuitBreaker.Successes, "cb_successes", 2, "Consecutive successful half-open checks needed to close circuit") } diff --git a/cmd/serve.go b/cmd/serve.go index 30251d1..ad8b6de 100644 --- a/cmd/serve.go +++ b/cmd/serve.go @@ -63,9 +63,34 @@ var serveCMD = &cobra.Command{ } func serve(cmd *cobra.Command, args []string) { + corsCfg := core.DefaultCORSConfig() + corsCfg.AllowOrigins = config.CORS.AllowOrigins + corsCfg.AllowMethods = config.CORS.AllowMethods + corsCfg.AllowHeaders = config.CORS.AllowHeaders + corsCfg.MaxAge = config.CORS.MaxAge + + serverOpts := core.ServerOptions{ + EnableCORS: config.CORS.Enabled, + CORS: corsCfg, + AllowEndpointFallback: config.Resilience.AllowEndpointFallback, + Resilience: core.ResilientConfig{ + Retry: core.RetryConfig{ + MaxRetries: config.Resilience.MaxRetries, + InitialBackoff: 1 * time.Second, + MaxBackoff: 30 * time.Second, + BackoffFactor: 2.0, + }, + CircuitBreaker: core.CircuitBreakerConfig{ + FailureThreshold: config.CircuitBreaker.Failures, + RecoveryTimeout: time.Duration(config.CircuitBreaker.RecoverySeconds) * time.Second, + SuccessThreshold: config.CircuitBreaker.Successes, + }, + }, + } + if config.App.IsRawRequests { logrus.Warn("Browserless results are very inconsistent or may not even work!") - serv := core.NewServer(config.App.Host, config.App.Port, + serv := core.NewServerWithOptions(config.App.Host, config.App.Port, serverOpts, &rawEngine{name: "google"}, &rawEngine{name: "yandex"}, &rawEngine{name: "baidu"}, @@ -102,7 +127,7 @@ func serve(cmd *cobra.Command, args []string) { bing := bing.New(*browser, config.BingConfig) ddg := duckduckgo.New(*browser, config.DuckDuckGoConfig) - serv := core.NewServer(config.App.Host, config.App.Port, gogl, yand, baidu, bing, ddg) + serv := core.NewServerWithOptions(config.App.Host, config.App.Port, serverOpts, gogl, yand, baidu, bing, ddg) err = serv.Listen() if err != nil { diff --git a/config.yaml b/config.yaml index a0596b5..dfa5694 100644 --- a/config.yaml +++ b/config.yaml @@ -1,25 +1,41 @@ app: - host: 0.0.0.0 - port: 7000 - debug: false - verbose: true - timeout: 15 - head: false - leakless: false - leave_head: false - stealth: false - insecure: true - + host: 0.0.0.0 # API host to bind + port: 7000 # API port to bind + debug: false # Enable debug logs and force browser UI mode + verbose: true # Enable info-level request logs + timeout: 15 # Browser/search timeout in seconds + head: false # Show browser UI (headful mode) + leakless: false # Force browser process cleanup after request + leave_head: false # Keep tabs open after request for debugging + stealth: false # Enable stealth browser plugin + insecure: true # Allow insecure TLS connections + proxy: "" # Optional HTTP/SOCKS5 proxy URL # Optional custom browser binary path (chrome/chromium/edge..) #browser_path: "C:/Program Files/BraveSoftware/Brave-Browser/Application/brave.exe" +resilience: + max_retries: 3 # Retry attempts per engine request (0 disables retries) + allow_endpoint_fallback: false # Keep dedicated endpoints engine-pure by default + +circuit_breaker: + failures: 5 # Consecutive failures required to open circuit + recovery_seconds: 60 # Wait time before moving open circuit to half-open + successes: 2 # Consecutive half-open successes required to close circuit + +cors: + enabled: true + allow_origins: "*" + allow_methods: "GET, POST, OPTIONS" + allow_headers: "Origin, Content-Type, Accept, Authorization" + max_age: 86400 # Browser preflight cache in seconds + 2captcha: apikey: "123123123123123" google: - rate_requests: 4 # Number of requests per Minute - rate_burst: 2 # Number of non-ratelimited requests per Minute - captcha: true + rate_requests: 4 # Allowed average requests per minute + rate_burst: 2 # Burst requests before limiter applies + captcha: true # Enable captcha solver path yandex: rate_requests: 4 diff --git a/core/circuit_breaker.go b/core/circuit_breaker.go new file mode 100644 index 0000000..3303bff --- /dev/null +++ b/core/circuit_breaker.go @@ -0,0 +1,210 @@ +package core + +import ( + "fmt" + "sync" + "time" + + "github.com/sirupsen/logrus" +) + +type CircuitState int + +const ( + CircuitClosed CircuitState = iota + CircuitOpen + CircuitHalfOpen +) + +func (s CircuitState) String() string { + switch s { + case CircuitClosed: + return "closed" + case CircuitOpen: + return "open" + case CircuitHalfOpen: + return "half-open" + default: + return "unknown" + } +} + +type CircuitBreakerConfig struct { + FailureThreshold int + RecoveryTimeout time.Duration + SuccessThreshold int +} + +func DefaultCircuitBreakerConfig() CircuitBreakerConfig { + return CircuitBreakerConfig{ + FailureThreshold: 5, + RecoveryTimeout: 60 * time.Second, + SuccessThreshold: 2, + } +} + +// CircuitBreaker tracks failure state for one engine. +type CircuitBreaker struct { + mu sync.RWMutex + name string + state CircuitState + config CircuitBreakerConfig + failureCount int + successCount int + lastFailureTime time.Time + lastStateChange time.Time +} + +func NewCircuitBreaker(name string, cfg CircuitBreakerConfig) *CircuitBreaker { + return &CircuitBreaker{ + name: name, + state: CircuitClosed, + config: cfg, + lastStateChange: time.Now(), + } +} + +func (cb *CircuitBreaker) AllowRequest() bool { + cb.mu.Lock() + defer cb.mu.Unlock() + + switch cb.state { + case CircuitClosed: + return true + case CircuitOpen: + if time.Since(cb.lastFailureTime) >= cb.config.RecoveryTimeout { + cb.setState(CircuitHalfOpen) + logrus.Infof("[CircuitBreaker][%s] Recovery timeout elapsed, moving to half-open", cb.name) + return true + } + return false + case CircuitHalfOpen: + return true + default: + return true + } +} + +func (cb *CircuitBreaker) RecordSuccess() { + cb.mu.Lock() + defer cb.mu.Unlock() + + switch cb.state { + case CircuitHalfOpen: + cb.successCount++ + if cb.successCount >= cb.config.SuccessThreshold { + cb.setState(CircuitClosed) + cb.failureCount = 0 + cb.successCount = 0 + logrus.Infof("[CircuitBreaker][%s] Recovered, circuit closed", cb.name) + } + case CircuitClosed: + cb.failureCount = 0 + } +} + +func (cb *CircuitBreaker) RecordFailure() { + cb.mu.Lock() + defer cb.mu.Unlock() + + cb.lastFailureTime = time.Now() + + switch cb.state { + case CircuitClosed: + cb.failureCount++ + if cb.failureCount >= cb.config.FailureThreshold { + cb.setState(CircuitOpen) + logrus.Warnf("[CircuitBreaker][%s] Circuit OPENED after %d consecutive failures (will retry in %s)", + cb.name, cb.failureCount, cb.config.RecoveryTimeout) + } + case CircuitHalfOpen: + cb.setState(CircuitOpen) + cb.successCount = 0 + logrus.Warnf("[CircuitBreaker][%s] Failed during half-open, circuit re-opened", cb.name) + } +} + +func (cb *CircuitBreaker) State() CircuitState { + cb.mu.RLock() + defer cb.mu.RUnlock() + return cb.state +} + +func (cb *CircuitBreaker) Stats() map[string]interface{} { + cb.mu.RLock() + defer cb.mu.RUnlock() + + stats := map[string]interface{}{ + "engine": cb.name, + "state": cb.state.String(), + "failure_count": cb.failureCount, + "last_changed": cb.lastStateChange.Format(time.RFC3339), + } + + if cb.state == CircuitOpen { + remaining := cb.config.RecoveryTimeout - time.Since(cb.lastFailureTime) + if remaining < 0 { + remaining = 0 + } + + // Expose retry_in as integer seconds for easier client-side processing. + retryInSeconds := int64(0) + if remaining > 0 { + retryInSeconds = int64((remaining + time.Second - time.Nanosecond) / time.Second) + } + stats["retry_in"] = retryInSeconds + } + + return stats +} + +func (cb *CircuitBreaker) setState(state CircuitState) { + cb.state = state + cb.lastStateChange = time.Now() +} + +type CircuitBreakerManager struct { + mu sync.RWMutex + breakers map[string]*CircuitBreaker + config CircuitBreakerConfig +} + +func NewCircuitBreakerManager(cfg CircuitBreakerConfig) *CircuitBreakerManager { + return &CircuitBreakerManager{ + breakers: make(map[string]*CircuitBreaker), + config: cfg, + } +} + +func (m *CircuitBreakerManager) Get(engineName string) *CircuitBreaker { + m.mu.RLock() + if cb, ok := m.breakers[engineName]; ok { + m.mu.RUnlock() + return cb + } + m.mu.RUnlock() + + m.mu.Lock() + defer m.mu.Unlock() + + if cb, ok := m.breakers[engineName]; ok { + return cb + } + + cb := NewCircuitBreaker(engineName, m.config) + m.breakers[engineName] = cb + return cb +} + +func (m *CircuitBreakerManager) AllStats() []map[string]interface{} { + m.mu.RLock() + defer m.mu.RUnlock() + + stats := make([]map[string]interface{}, 0, len(m.breakers)) + for _, cb := range m.breakers { + stats = append(stats, cb.Stats()) + } + return stats +} + +var ErrCircuitOpen = fmt.Errorf("circuit breaker is open - engine temporarily disabled") diff --git a/core/circuit_breaker_test.go b/core/circuit_breaker_test.go new file mode 100644 index 0000000..91b09bb --- /dev/null +++ b/core/circuit_breaker_test.go @@ -0,0 +1,165 @@ +package core + +import ( + "testing" + "time" +) + +func newTestCircuitBreaker(t *testing.T, cfg CircuitBreakerConfig) *CircuitBreaker { + t.Helper() + return NewCircuitBreaker("test-engine", cfg) +} + +// TestCircuitBreaker_OpensAfterThreshold verifies that consecutive failures in closed state +// move the breaker to open exactly on configured threshold and block new requests. +func TestCircuitBreaker_OpensAfterThreshold(t *testing.T) { + cfg := CircuitBreakerConfig{ + FailureThreshold: 3, + RecoveryTimeout: time.Second, + SuccessThreshold: 1, + } + cb := newTestCircuitBreaker(t, cfg) + + cb.RecordFailure() + cb.RecordFailure() + if cb.State() != CircuitClosed { + t.Fatalf("expected closed after 2 failures, got: %s", cb.State()) + } + + cb.RecordFailure() + if cb.State() != CircuitOpen { + t.Fatalf("expected open after %d failures, got: %s", cfg.FailureThreshold, cb.State()) + } + if cb.AllowRequest() { + t.Error("expected request blocked in open state") + } +} + +// TestCircuitBreaker_RecoveryToHalfOpen verifies timed recovery from open to half-open +// when recovery timeout elapses and a new request is attempted. +func TestCircuitBreaker_RecoveryToHalfOpen(t *testing.T) { + cfg := CircuitBreakerConfig{ + FailureThreshold: 2, + RecoveryTimeout: 50 * time.Millisecond, + SuccessThreshold: 1, + } + cb := newTestCircuitBreaker(t, cfg) + + cb.RecordFailure() + cb.RecordFailure() + if cb.State() != CircuitOpen { + t.Fatal("expected open") + } + + time.Sleep(60 * time.Millisecond) + if !cb.AllowRequest() { + t.Error("should allow request after recovery timeout") + } + if cb.State() != CircuitHalfOpen { + t.Errorf("expected half-open, got: %s", cb.State()) + } +} + +// TestCircuitBreaker_HalfOpenSuccessClosesCircuit verifies that half-open state closes +// only after configured number of successful probes. +func TestCircuitBreaker_HalfOpenSuccessClosesCircuit(t *testing.T) { + cfg := CircuitBreakerConfig{ + FailureThreshold: 1, + RecoveryTimeout: 20 * time.Millisecond, + SuccessThreshold: 2, + } + cb := newTestCircuitBreaker(t, cfg) + + cb.RecordFailure() + if cb.State() != CircuitOpen { + t.Fatalf("expected open, got: %s", cb.State()) + } + + time.Sleep(30 * time.Millisecond) + if !cb.AllowRequest() { + t.Fatal("expected request to pass in recovery window") + } + if cb.State() != CircuitHalfOpen { + t.Fatalf("expected half-open after recovery timeout, got: %s", cb.State()) + } + + cb.RecordSuccess() + if cb.State() != CircuitHalfOpen { + t.Fatalf("expected to stay half-open until success threshold reached, got: %s", cb.State()) + } + + cb.RecordSuccess() + if cb.State() != CircuitClosed { + t.Fatalf("expected closed after success threshold reached, got: %s", cb.State()) + } +} + +// TestCircuitBreaker_HalfOpenFailureReopens verifies that a failed probe in half-open +// immediately re-opens the circuit. +func TestCircuitBreaker_HalfOpenFailureReopens(t *testing.T) { + cfg := CircuitBreakerConfig{ + FailureThreshold: 1, + RecoveryTimeout: 20 * time.Millisecond, + SuccessThreshold: 1, + } + cb := newTestCircuitBreaker(t, cfg) + + cb.RecordFailure() + time.Sleep(30 * time.Millisecond) + if !cb.AllowRequest() { + t.Fatal("expected probe request in half-open") + } + if cb.State() != CircuitHalfOpen { + t.Fatalf("expected half-open, got: %s", cb.State()) + } + + cb.RecordFailure() + if cb.State() != CircuitOpen { + t.Fatalf("expected open after failed half-open probe, got: %s", cb.State()) + } +} + +// TestCircuitBreaker_Stats verifies stats payload fields and that retry_in is exposed +// only when breaker is open. +func TestCircuitBreaker_Stats(t *testing.T) { + cb := NewCircuitBreaker("test-engine", DefaultCircuitBreakerConfig()) + cb.RecordFailure() + + stats := cb.Stats() + if stats["engine"] != "test-engine" { + t.Fatalf("expected engine=test-engine, got: %v", stats["engine"]) + } + if stats["state"] != "closed" { + t.Fatalf("expected state=closed, got: %v", stats["state"]) + } + if stats["failure_count"].(int) != 1 { + t.Fatalf("expected failure_count=1, got: %v", stats["failure_count"]) + } + if _, ok := stats["retry_in"]; ok { + t.Fatalf("did not expect retry_in in closed state, got: %v", stats["retry_in"]) + } + + openCfg := CircuitBreakerConfig{FailureThreshold: 1, RecoveryTimeout: time.Second, SuccessThreshold: 1} + openCB := NewCircuitBreaker("open-engine", openCfg) + openCB.RecordFailure() + openStats := openCB.Stats() + retryIn, ok := openStats["retry_in"].(int64) + if !ok { + t.Fatalf("expected retry_in int64 in open state, got: %T", openStats["retry_in"]) + } + if retryIn <= 0 { + t.Fatalf("expected retry_in > 0 in open state, got: %d", retryIn) + } +} + +// TestCircuitBreakerManager_AllStats verifies manager creates and reports per-engine breakers. +func TestCircuitBreakerManager_AllStats(t *testing.T) { + mgr := NewCircuitBreakerManager(DefaultCircuitBreakerConfig()) + mgr.Get("google") + mgr.Get("yandex") + + stats := mgr.AllStats() + if len(stats) != 2 { + t.Errorf("expected 2 entries, got: %d", len(stats)) + } +} diff --git a/core/middleware.go b/core/middleware.go new file mode 100644 index 0000000..c935df8 --- /dev/null +++ b/core/middleware.go @@ -0,0 +1,144 @@ +package core + +import ( + "fmt" + "strings" + "time" + + "github.com/gofiber/fiber/v2" + "github.com/sirupsen/logrus" +) + +type JSONErrorResponse struct { + Error string `json:"error"` + Code int `json:"code"` + Message string `json:"message,omitempty"` +} + +type CORSConfig struct { + AllowOrigins string + AllowMethods string + AllowHeaders string + MaxAge int +} + +func DefaultCORSConfig() CORSConfig { + return CORSConfig{ + AllowOrigins: "*", + AllowMethods: "GET, POST, OPTIONS", + AllowHeaders: "Origin, Content-Type, Accept, Authorization", + MaxAge: 86400, + } +} + +func CORSMiddleware(cfg CORSConfig) fiber.Handler { + cfg = normalizeCORSConfig(cfg) + + return func(c *fiber.Ctx) error { + c.Set("Access-Control-Allow-Origin", cfg.AllowOrigins) + c.Set("Access-Control-Allow-Methods", cfg.AllowMethods) + c.Set("Access-Control-Allow-Headers", cfg.AllowHeaders) + c.Set("Access-Control-Max-Age", fmt.Sprintf("%d", cfg.MaxAge)) + + if c.Method() == "OPTIONS" { + return c.SendStatus(fiber.StatusNoContent) + } + return c.Next() + } +} + +// normalizeCORSConfig keeps CORS behavior predictable when config provides partial values. +func normalizeCORSConfig(cfg CORSConfig) CORSConfig { + defaults := DefaultCORSConfig() + + if strings.TrimSpace(cfg.AllowOrigins) == "" { + cfg.AllowOrigins = defaults.AllowOrigins + } + if strings.TrimSpace(cfg.AllowMethods) == "" { + cfg.AllowMethods = defaults.AllowMethods + } + if strings.TrimSpace(cfg.AllowHeaders) == "" { + cfg.AllowHeaders = defaults.AllowHeaders + } + if cfg.MaxAge <= 0 { + cfg.MaxAge = defaults.MaxAge + } + + return cfg +} + +func RequestLoggerMiddleware() fiber.Handler { + return func(c *fiber.Ctx) error { + start := time.Now() + err := c.Next() + + latency := time.Since(start) + status := c.Response().StatusCode() + if err != nil { + if e, ok := err.(*fiber.Error); ok { + status = e.Code + } else { + status = fiber.StatusInternalServerError + } + } + + logFields := logrus.Fields{ + "method": c.Method(), + "path": c.Path(), + "status": status, + "latency": latency.String(), + "ip": c.IP(), + } + if query := c.Query("text"); query != "" { + logFields["query"] = query + } + + entry := logrus.WithFields(logFields) + if status >= 500 { + entry.Error("Request failed") + } else if status >= 400 { + entry.Warn("Request error") + } else { + entry.Info("Request completed") + } + + return err + } +} + +func JSONErrorMiddleware() fiber.ErrorHandler { + return func(c *fiber.Ctx, err error) error { + code := fiber.StatusInternalServerError + if e, ok := err.(*fiber.Error); ok { + code = e.Code + } + + resp := JSONErrorResponse{ + Error: statusText(code), + Code: code, + Message: err.Error(), + } + + c.Set("Content-Type", "application/json") + return c.Status(code).JSON(resp) + } +} + +func statusText(code int) string { + switch { + case code == 400: + return "bad_request" + case code == 404: + return "not_found" + case code == 429: + return "rate_limited" + case code == 503: + return "service_unavailable" + case code >= 400 && code < 500: + return "client_error" + case code >= 500: + return "server_error" + default: + return "error" + } +} diff --git a/core/middleware_test.go b/core/middleware_test.go new file mode 100644 index 0000000..d4271d0 --- /dev/null +++ b/core/middleware_test.go @@ -0,0 +1,101 @@ +package core + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gofiber/fiber/v2" +) + +func TestStatusText(t *testing.T) { + tests := []struct { + code int + expected string + }{ + {400, "bad_request"}, + {404, "not_found"}, + {429, "rate_limited"}, + {503, "service_unavailable"}, + {401, "client_error"}, + {500, "server_error"}, + {200, "error"}, + } + + for _, tt := range tests { + result := statusText(tt.code) + if result != tt.expected { + t.Errorf("statusText(%d) = %s, want %s", tt.code, result, tt.expected) + } + } +} + +// Middleware unit tests validate CORS behavior itself (header composition and preflight semantics). +func TestCORSMiddleware_UsesConfiguredHeaders(t *testing.T) { + app := fiber.New() + app.Use(CORSMiddleware(CORSConfig{ + AllowOrigins: "https://example.com", + AllowMethods: "GET,OPTIONS", + AllowHeaders: "Authorization,Content-Type", + MaxAge: 1200, + })) + app.Get("/ping", func(c *fiber.Ctx) error { + return c.SendStatus(fiber.StatusOK) + }) + + req := httptest.NewRequest(http.MethodGet, "/ping", nil) + resp, err := app.Test(req, -1) + if err != nil { + t.Fatalf("request failed: %v", err) + } + + if got := resp.Header.Get("Access-Control-Allow-Origin"); got != "https://example.com" { + t.Fatalf("unexpected allow-origin: %q", got) + } + if got := resp.Header.Get("Access-Control-Allow-Methods"); got != "GET,OPTIONS" { + t.Fatalf("unexpected allow-methods: %q", got) + } + if got := resp.Header.Get("Access-Control-Allow-Headers"); got != "Authorization,Content-Type" { + t.Fatalf("unexpected allow-headers: %q", got) + } + if got := resp.Header.Get("Access-Control-Max-Age"); got != "1200" { + t.Fatalf("unexpected max-age: %q", got) + } +} + +func TestCORSMiddleware_OPTIONSReturnsNoContent(t *testing.T) { + app := fiber.New() + app.Use(CORSMiddleware(DefaultCORSConfig())) + app.Get("/ping", func(c *fiber.Ctx) error { + return c.SendStatus(fiber.StatusOK) + }) + + req := httptest.NewRequest(http.MethodOptions, "/ping", nil) + resp, err := app.Test(req, -1) + if err != nil { + t.Fatalf("request failed: %v", err) + } + if resp.StatusCode != fiber.StatusNoContent { + t.Fatalf("expected 204 for OPTIONS, got %d", resp.StatusCode) + } +} + +func TestNormalizeCORSConfig_FillsMissingValues(t *testing.T) { + cfg := normalizeCORSConfig(CORSConfig{ + AllowOrigins: "https://example.com", + }) + def := DefaultCORSConfig() + + if cfg.AllowOrigins != "https://example.com" { + t.Fatalf("expected custom allow_origins preserved, got %q", cfg.AllowOrigins) + } + if cfg.AllowMethods != def.AllowMethods { + t.Fatalf("expected default allow_methods, got %q", cfg.AllowMethods) + } + if cfg.AllowHeaders != def.AllowHeaders { + t.Fatalf("expected default allow_headers, got %q", cfg.AllowHeaders) + } + if cfg.MaxAge != def.MaxAge { + t.Fatalf("expected default max_age, got %d", cfg.MaxAge) + } +} diff --git a/core/resilient.go b/core/resilient.go new file mode 100644 index 0000000..01bb22e --- /dev/null +++ b/core/resilient.go @@ -0,0 +1,232 @@ +package core + +import ( + "context" + "fmt" + "sync" + + "github.com/sirupsen/logrus" +) + +// ResilientSearcher wraps engines with retry and circuit breaker protection. +type ResilientSearcher struct { + engines []SearchEngine + cbManager *CircuitBreakerManager + retryCfg RetryConfig +} + +type ResilientConfig struct { + Retry RetryConfig + CircuitBreaker CircuitBreakerConfig +} + +func DefaultResilientConfig() ResilientConfig { + return ResilientConfig{ + Retry: DefaultRetryConfig(), + CircuitBreaker: DefaultCircuitBreakerConfig(), + } +} + +func NewResilientSearcher(engines []SearchEngine, cfg ResilientConfig) *ResilientSearcher { + return &ResilientSearcher{ + engines: engines, + cbManager: NewCircuitBreakerManager(cfg.CircuitBreaker), + retryCfg: cfg.Retry, + } +} + +// SearchPrimary keeps dedicated endpoints engine-pure (no fallback). +func (rs *ResilientSearcher) SearchPrimary(primaryEngine SearchEngine, q Query) ([]SearchResult, string, error) { + results, err := rs.searchWithProtection(primaryEngine, q) + if err != nil { + return nil, primaryEngine.Name(), err + } + return results, primaryEngine.Name(), nil +} + +// SearchWithFallback retries primary and then tries other initialized engines. +func (rs *ResilientSearcher) SearchWithFallback(primaryEngine SearchEngine, q Query) ([]SearchResult, string, error) { + results, err := rs.searchWithProtection(primaryEngine, q) + if err == nil { + return results, primaryEngine.Name(), nil + } + + logrus.Warnf("[Resilient] Primary engine %s failed: %s. Trying fallback engines...", primaryEngine.Name(), err) + for _, fallbackEngine := range rs.engines { + if fallbackEngine.Name() == primaryEngine.Name() || !fallbackEngine.IsInitialized() { + continue + } + + results, err := rs.searchWithProtection(fallbackEngine, q) + if err == nil { + logrus.Infof("[Resilient] Fallback to %s succeeded with %d results", fallbackEngine.Name(), len(results)) + return results, fallbackEngine.Name(), nil + } + logrus.Warnf("[Resilient] Fallback engine %s also failed: %s", fallbackEngine.Name(), err) + } + + return nil, primaryEngine.Name(), ErrAllEnginesFailed +} + +func (rs *ResilientSearcher) SearchImagePrimary(primaryEngine SearchEngine, q Query) ([]SearchResult, string, error) { + results, err := rs.searchImageWithProtection(primaryEngine, q) + if err != nil { + return nil, primaryEngine.Name(), err + } + return results, primaryEngine.Name(), nil +} + +func (rs *ResilientSearcher) SearchImageWithFallback(primaryEngine SearchEngine, q Query) ([]SearchResult, string, error) { + results, err := rs.searchImageWithProtection(primaryEngine, q) + if err == nil { + return results, primaryEngine.Name(), nil + } + + logrus.Warnf("[Resilient] Primary engine %s image search failed: %s. Trying fallback engines...", primaryEngine.Name(), err) + for _, fallbackEngine := range rs.engines { + if fallbackEngine.Name() == primaryEngine.Name() || !fallbackEngine.IsInitialized() { + continue + } + + results, err := rs.searchImageWithProtection(fallbackEngine, q) + if err == nil { + logrus.Infof("[Resilient] Image fallback to %s succeeded with %d results", fallbackEngine.Name(), len(results)) + return results, fallbackEngine.Name(), nil + } + } + + return nil, primaryEngine.Name(), ErrAllEnginesFailed +} + +func (rs *ResilientSearcher) searchWithProtection(engine SearchEngine, q Query) ([]SearchResult, error) { + cb := rs.cbManager.Get(engine.Name()) + if !cb.AllowRequest() { + return nil, ErrCircuitOpen + } + + result := RetryableSearch(rs.retryCfg, engine.Name(), func() ([]SearchResult, error) { + limiter := engine.GetRateLimiter() + if limiter != nil { + if err := limiter.Wait(context.Background()); err != nil { + return nil, err + } + } + return engine.Search(q) + }) + + if result.Err != nil { + cb.RecordFailure() + return nil, result.Err + } + + cb.RecordSuccess() + return result.Results, nil +} + +func (rs *ResilientSearcher) searchImageWithProtection(engine SearchEngine, q Query) ([]SearchResult, error) { + cb := rs.cbManager.Get(engine.Name()) + if !cb.AllowRequest() { + return nil, ErrCircuitOpen + } + + result := RetryableSearch(rs.retryCfg, engine.Name(), func() ([]SearchResult, error) { + limiter := engine.GetRateLimiter() + if limiter != nil { + if err := limiter.Wait(context.Background()); err != nil { + return nil, err + } + } + return engine.SearchImage(q) + }) + + if result.Err != nil { + cb.RecordFailure() + return nil, result.Err + } + + cb.RecordSuccess() + return result.Results, nil +} + +// SearchAllParallel applies retry/circuit protections per engine for mega search. +func (rs *ResilientSearcher) SearchAllParallel(q Query, engines []SearchEngine) []MegaSearchResult { + var wg sync.WaitGroup + var mu sync.Mutex + var allResults []MegaSearchResult + + for _, engine := range engines { + if !engine.IsInitialized() { + continue + } + if !rs.cbManager.Get(engine.Name()).AllowRequest() { + logrus.Infof("[Resilient] Skipping %s in megasearch (circuit open)", engine.Name()) + continue + } + + wg.Add(1) + go func(eng SearchEngine) { + defer wg.Done() + + results, err := rs.searchWithProtection(eng, q) + if err != nil { + return + } + + mu.Lock() + for _, r := range results { + allResults = append(allResults, MegaSearchResult{ + SearchResult: r, + Engine: eng.Name(), + }) + } + mu.Unlock() + }(engine) + } + + wg.Wait() + return allResults +} + +func (rs *ResilientSearcher) SearchAllImageParallel(q Query, engines []SearchEngine) []MegaSearchResult { + var wg sync.WaitGroup + var mu sync.Mutex + var allResults []MegaSearchResult + + for _, engine := range engines { + if !engine.IsInitialized() { + continue + } + if !rs.cbManager.Get(engine.Name()).AllowRequest() { + logrus.Infof("[Resilient] Skipping %s in megaimage (circuit open)", engine.Name()) + continue + } + + wg.Add(1) + go func(eng SearchEngine) { + defer wg.Done() + + results, err := rs.searchImageWithProtection(eng, q) + if err != nil { + return + } + + mu.Lock() + for _, r := range results { + allResults = append(allResults, MegaSearchResult{ + SearchResult: r, + Engine: eng.Name(), + }) + } + mu.Unlock() + }(engine) + } + + wg.Wait() + return allResults +} + +func (rs *ResilientSearcher) GetCircuitBreakerStats() []map[string]interface{} { + return rs.cbManager.AllStats() +} + +var ErrAllEnginesFailed = fmt.Errorf("all search engines failed") diff --git a/core/retry.go b/core/retry.go new file mode 100644 index 0000000..76109f6 --- /dev/null +++ b/core/retry.go @@ -0,0 +1,91 @@ +package core + +import ( + "fmt" + "math" + "time" + + "github.com/sirupsen/logrus" +) + +// RetryConfig controls retry behavior. +type RetryConfig struct { + MaxRetries int + InitialBackoff time.Duration + MaxBackoff time.Duration + BackoffFactor float64 +} + +func DefaultRetryConfig() RetryConfig { + return RetryConfig{ + MaxRetries: 3, + InitialBackoff: time.Second, + MaxBackoff: 30 * time.Second, + BackoffFactor: 2.0, + } +} + +type RetryResult struct { + Results []SearchResult + Err error + Attempts int + Engine string +} + +// RetryableSearch executes searchFn with exponential backoff retries. +// CAPTCHA errors are not retried. +func RetryableSearch(cfg RetryConfig, engineName string, searchFn func() ([]SearchResult, error)) RetryResult { + if cfg.BackoffFactor <= 0 { + cfg.BackoffFactor = 2.0 + } + + var lastErr error + for attempt := 0; attempt <= cfg.MaxRetries; attempt++ { + if attempt > 0 { + backoff := calculateBackoff(cfg, attempt) + logrus.Warnf("[%s] Retry attempt %d/%d after %s", engineName, attempt, cfg.MaxRetries, backoff) + time.Sleep(backoff) + } + + results, err := searchFn() + if err == nil { + if attempt > 0 { + logrus.Infof("[%s] Succeeded on retry attempt %d", engineName, attempt) + } + return RetryResult{ + Results: results, + Attempts: attempt + 1, + Engine: engineName, + } + } + + lastErr = err + if err == ErrCaptcha { + logrus.Warnf("[%s] CAPTCHA detected, skipping retries", engineName) + return RetryResult{ + Err: err, + Attempts: attempt + 1, + Engine: engineName, + } + } + + logrus.Warnf("[%s] Attempt %d failed: %s", engineName, attempt+1, err) + } + + return RetryResult{ + Err: fmt.Errorf("all %d attempts failed for %s: %w", cfg.MaxRetries+1, engineName, lastErr), + Attempts: cfg.MaxRetries + 1, + Engine: engineName, + } +} + +func calculateBackoff(cfg RetryConfig, attempt int) time.Duration { + backoff := float64(cfg.InitialBackoff) * math.Pow(cfg.BackoffFactor, float64(attempt-1)) + if backoff > float64(cfg.MaxBackoff) { + backoff = float64(cfg.MaxBackoff) + } + if backoff < 0 { + backoff = 0 + } + return time.Duration(backoff) +} diff --git a/core/retry_test.go b/core/retry_test.go new file mode 100644 index 0000000..ca3faf6 --- /dev/null +++ b/core/retry_test.go @@ -0,0 +1,86 @@ +package core + +import ( + "errors" + "testing" + "time" +) + +func TestRetryableSearch_SuccessOnFirstAttempt(t *testing.T) { + cfg := RetryConfig{MaxRetries: 3, InitialBackoff: 10 * time.Millisecond, MaxBackoff: 100 * time.Millisecond, BackoffFactor: 2.0} + calls := 0 + + result := RetryableSearch(cfg, "test", func() ([]SearchResult, error) { + calls++ + return []SearchResult{{Title: "result1"}}, nil + }) + + if result.Err != nil { + t.Fatalf("expected no error, got: %v", result.Err) + } + if result.Attempts != 1 { + t.Errorf("expected 1 attempt, got: %d", result.Attempts) + } + if calls != 1 { + t.Errorf("expected 1 call, got: %d", calls) + } +} + +func TestRetryableSearch_AllAttemptsFail(t *testing.T) { + cfg := RetryConfig{MaxRetries: 2, InitialBackoff: 10 * time.Millisecond, MaxBackoff: 50 * time.Millisecond, BackoffFactor: 2.0} + calls := 0 + + result := RetryableSearch(cfg, "test", func() ([]SearchResult, error) { + calls++ + return nil, errors.New("persistent failure") + }) + + if result.Err == nil { + t.Fatal("expected error, got nil") + } + if calls != 3 { + t.Errorf("expected 3 calls (1 + 2 retries), got: %d", calls) + } + if result.Attempts != 3 { + t.Errorf("expected 3 attempts, got: %d", result.Attempts) + } +} + +func TestRetryableSearch_CaptchaNotRetried(t *testing.T) { + cfg := RetryConfig{MaxRetries: 3, InitialBackoff: 10 * time.Millisecond, MaxBackoff: 100 * time.Millisecond, BackoffFactor: 2.0} + calls := 0 + + result := RetryableSearch(cfg, "test", func() ([]SearchResult, error) { + calls++ + return nil, ErrCaptcha + }) + + if !errors.Is(result.Err, ErrCaptcha) { + t.Fatalf("expected ErrCaptcha, got: %v", result.Err) + } + if calls != 1 { + t.Errorf("expected 1 call, got: %d", calls) + } +} + +func TestCalculateBackoff(t *testing.T) { + cfg := RetryConfig{InitialBackoff: 1 * time.Second, MaxBackoff: 10 * time.Second, BackoffFactor: 2.0} + + tests := []struct { + attempt int + expected time.Duration + }{ + {1, 1 * time.Second}, + {2, 2 * time.Second}, + {3, 4 * time.Second}, + {4, 8 * time.Second}, + {5, 10 * time.Second}, + } + + for _, tt := range tests { + got := calculateBackoff(cfg, tt.attempt) + if got != tt.expected { + t.Errorf("attempt %d: expected %s, got %s", tt.attempt, tt.expected, got) + } + } +} diff --git a/core/server.go b/core/server.go index 853ec92..587dbdc 100644 --- a/core/server.go +++ b/core/server.go @@ -26,355 +26,133 @@ type Server struct { app *fiber.App addr string searchEngines []SearchEngine + resilient *ResilientSearcher startTime time.Time + opts ServerOptions +} + +type ServerOptions struct { + EnableCORS bool + CORS CORSConfig + AllowEndpointFallback bool + Resilience ResilientConfig +} + +func DefaultServerOptions() ServerOptions { + return ServerOptions{ + EnableCORS: true, + CORS: DefaultCORSConfig(), + AllowEndpointFallback: false, + Resilience: DefaultResilientConfig(), + } } func NewServer(host string, port int, searchEngines ...SearchEngine) *Server { + return NewServerWithOptions(host, port, DefaultServerOptions(), searchEngines...) +} + +func NewServerWithOptions(host string, port int, opts ServerOptions, searchEngines ...SearchEngine) *Server { addr := fmt.Sprintf("%s:%d", host, port) + app := fiber.New(fiber.Config{ + ErrorHandler: JSONErrorMiddleware(), + }) + serv := Server{ - app: fiber.New(), + app: app, addr: addr, searchEngines: searchEngines, + resilient: NewResilientSearcher(searchEngines, opts.Resilience), startTime: time.Now(), + opts: opts, + } + logrus.Info("Resilient search enabled: retry + circuit breaker") + if opts.AllowEndpointFallback { + logrus.Warn("Dedicated endpoint fallback is enabled") } - // Health endpoint used by orchestration probes. - serv.app.Get("/health", serv.handleHealthCheck) + if opts.EnableCORS { + app.Use(CORSMiddleware(opts.CORS)) + } + app.Use(RequestLoggerMiddleware()) + + app.Get("/health", serv.handleHealthCheck) + app.Get("/resilience/stats", serv.handleResilienceStats) for _, engine := range searchEngines { locEngine := engine - limiter := engine.GetRateLimiter() - // Custom endpoint mapping for DuckDuckGo endpointName := strings.ToLower(locEngine.Name()) if endpointName == "duckduckgo" { endpointName = "duck" } serv.app.Get(fmt.Sprintf("/%s/search", endpointName), func(c *fiber.Ctx) error { - q := Query{} - err := q.InitFromContext(c) - if err != nil { - logrus.Errorf("Error while setting %s query: %s", locEngine.Name(), err) - return err - } - - logrus.Infof("Starting SERP search request using %s engine for query: %s", locEngine.Name(), q.Text) - - err = limiter.Wait(context.Background()) - if err != nil { - logrus.Errorf("Ratelimiter error during %s query: %s", locEngine.Name(), err) - } - - res, err := locEngine.Search(q) - if err != nil { - switch err { - case ErrCaptcha: - err = fmt.Errorf("captcha found, please stop sending requests for a while\n%s", err) - case ErrSearchTimeout: - err = fmt.Errorf("%s", err) - } - - logrus.Errorf("Error during %s search: %s", locEngine.Name(), err) - return fiber.NewError(fiber.StatusServiceUnavailable, err.Error()) - } - - logrus.Infof("Successfully completed SERP search using %s engine, returned %d results", locEngine.Name(), len(res)) - return c.JSON(res) + return serv.handleDedicatedEndpoint(c, locEngine, false) }) serv.app.Get(fmt.Sprintf("/%s/image", endpointName), func(c *fiber.Ctx) error { - q := Query{} - err := q.InitFromContext(c) - if err != nil { - logrus.Errorf("Error while setting %s query: %s", locEngine.Name(), err) - return err - } - - logrus.Infof("Starting SERP image search request using %s engine for query: %s", locEngine.Name(), q.Text) - - err = limiter.Wait(context.Background()) - if err != nil { - logrus.Errorf("Ratelimiter error during %s query: %s", locEngine.Name(), err) - } - - res, err := locEngine.SearchImage(q) - - if err != nil && len(res) > 0 { - logrus.Warnf("Partial results returned from %s image search despite error: %s", locEngine.Name(), err) - c.Status(503) - return c.JSON(res) - } - - if err != nil { - switch err { - case ErrCaptcha: - err = fmt.Errorf("captcha found, please stop sending requests for a while: %s", err) - case ErrSearchTimeout: - err = fmt.Errorf("%s", err) - } - - logrus.Errorf("Error during %s image search: %s", locEngine.Name(), err) - return fiber.NewError(fiber.StatusServiceUnavailable, err.Error()) - } - - logrus.Infof("Successfully completed SERP image search using [%s], returned %d results", locEngine.Name(), len(res)) - return c.JSON(res) + return serv.handleDedicatedEndpoint(c, locEngine, true) }) } - // Add megasearch endpoint serv.app.Get("/mega/search", serv.handleMegaSearch) - - // Add megasearch image endpoint serv.app.Get("/mega/image", serv.handleMegaImage) - - // Add endpoint to list available engines serv.app.Get("/mega/engines", serv.handleListEngines) return &serv } -// MegaSearchResult represents a search result with engine information -type MegaSearchResult struct { - SearchResult - Engine string `json:"engine"` -} - -// handleMegaSearch handles the /megasearch endpoint -func (s *Server) handleMegaSearch(c *fiber.Ctx) error { +func (s *Server) handleDedicatedEndpoint(c *fiber.Ctx, engine SearchEngine, isImage bool) error { q := Query{} - err := q.InitFromContext(c) - if err != nil { - logrus.Errorf("Error while setting megasearch query: %s", err) + if err := q.InitFromContext(c); err != nil { + logrus.Errorf("Error while setting %s query: %s", engine.Name(), err) return err } - // Get engines parameter to filter which engines to use - enginesParam := c.Query("engines", "") - var enginesToUse []SearchEngine + action := "search" + if isImage { + action = "image" + } + logrus.Infof("Starting SERP %s request using %s engine for query: %s", action, engine.Name(), q.Text) - if enginesParam != "" { - // Parse comma-separated list of engines - engineNames := strings.Split(enginesParam, ",") - for _, engineName := range engineNames { - engineName = strings.TrimSpace(strings.ToLower(engineName)) - for _, engine := range s.searchEngines { - if strings.ToLower(engine.Name()) == engineName { - enginesToUse = append(enginesToUse, engine) - break - } - } + var ( + res []SearchResult + usedEngine string + searchErr error + ) + + if isImage { + if s.opts.AllowEndpointFallback { + res, usedEngine, searchErr = s.resilient.SearchImageWithFallback(engine, q) + } else { + res, usedEngine, searchErr = s.resilient.SearchImagePrimary(engine, q) } } else { - // Use all engines if no specific engines specified - enginesToUse = s.searchEngines - } - - if len(enginesToUse) == 0 { - return fiber.NewError(fiber.StatusBadRequest, "No valid search engines specified") - } - - // Log which engines will be used - engineNames := make([]string, len(enginesToUse)) - for i, engine := range enginesToUse { - engineNames[i] = engine.Name() - } - logrus.Infof("Starting SERP megasearch request using engines: %s for query: %s", strings.Join(engineNames, ", "), q.Text) - - // Execute searches in parallel across selected engines - results := s.searchSelectedEngines(q, enginesToUse) - - // Deduplicate results while preserving engine information - dedupedResults := s.deduplicateMegaResults(results) - - logrus.Infof("Successfully completed SERP megasearch using %d engines, returned %d deduplicated results", len(enginesToUse), len(dedupedResults)) - return c.JSON(dedupedResults) -} - -// handleMegaImage handles the /mega/image endpoint -func (s *Server) handleMegaImage(c *fiber.Ctx) error { - q := Query{} - err := q.InitFromContext(c) - if err != nil { - logrus.Errorf("Error while setting megasearch image query: %s", err) - return err - } - - // Get engines parameter to filter which engines to use - enginesParam := c.Query("engines", "") - var enginesToUse []SearchEngine - - if enginesParam != "" { - // Parse comma-separated list of engines - engineNames := strings.Split(enginesParam, ",") - for _, engineName := range engineNames { - engineName = strings.TrimSpace(strings.ToLower(engineName)) - for _, engine := range s.searchEngines { - if strings.ToLower(engine.Name()) == engineName { - enginesToUse = append(enginesToUse, engine) - break - } - } - } - } else { - // Use all engines if no specific engines specified - enginesToUse = s.searchEngines - } - - if len(enginesToUse) == 0 { - return fiber.NewError(fiber.StatusBadRequest, "No valid search engines specified") - } - - // Log which engines will be used - engineNames := make([]string, len(enginesToUse)) - for i, engine := range enginesToUse { - engineNames[i] = engine.Name() - } - logrus.Infof("Starting SERP megasearch image request using engines: %s for query: %s", strings.Join(engineNames, ", "), q.Text) - - // Execute image searches in parallel across selected engines - results := s.searchSelectedEnginesImage(q, enginesToUse) - - // Deduplicate results while preserving engine information - dedupedResults := s.deduplicateMegaResults(results) - - logrus.Infof("Successfully completed SERP megasearch image using %d engines, returned %d deduplicated results", len(enginesToUse), len(dedupedResults)) - return c.JSON(dedupedResults) -} - -// handleListEngines lists all available search engines -func (s *Server) handleListEngines(c *fiber.Ctx) error { - var engines []map[string]interface{} - - for _, engine := range s.searchEngines { - engines = append(engines, map[string]interface{}{ - "name": engine.Name(), - "initialized": engine.IsInitialized(), - }) - } - - return c.JSON(map[string]interface{}{ - "engines": engines, - "total": len(engines), - }) -} - -// searchSelectedEngines performs parallel searches across selected engines -func (s *Server) searchSelectedEngines(q Query, engines []SearchEngine) []MegaSearchResult { - var wg sync.WaitGroup - var mu sync.Mutex - var allResults []MegaSearchResult - - for _, engine := range engines { - wg.Add(1) - go func(eng SearchEngine) { - defer wg.Done() - - // Apply rate limiting - limiter := eng.GetRateLimiter() - if limiter != nil { - err := limiter.Wait(context.Background()) - if err != nil { - logrus.Errorf("Ratelimiter error during %s megasearch: %s", eng.Name(), err) - } - } - - // Perform search - results, err := eng.Search(q) - if err != nil { - logrus.Errorf("Error during %s megasearch: %s", eng.Name(), err) - return - } - - // Convert to MegaSearchResult with engine info - mu.Lock() - for _, result := range results { - megaResult := MegaSearchResult{ - SearchResult: result, - Engine: eng.Name(), - } - allResults = append(allResults, megaResult) - } - mu.Unlock() - }(engine) - } - - wg.Wait() - return allResults -} - -// searchSelectedEnginesImage performs parallel image searches across selected engines -func (s *Server) searchSelectedEnginesImage(q Query, engines []SearchEngine) []MegaSearchResult { - var wg sync.WaitGroup - var mu sync.Mutex - var allResults []MegaSearchResult - - for _, engine := range engines { - wg.Add(1) - go func(eng SearchEngine) { - defer wg.Done() - - // Apply rate limiting - limiter := eng.GetRateLimiter() - if limiter != nil { - err := limiter.Wait(context.Background()) - if err != nil { - logrus.Errorf("Ratelimiter error during %s megasearch image: %s", eng.Name(), err) - } - } - - // Perform image search - results, err := eng.SearchImage(q) - if err != nil { - logrus.Errorf("Error during %s megasearch image: %s", eng.Name(), err) - return - } - - // Convert to MegaSearchResult with engine info - mu.Lock() - for _, result := range results { - megaResult := MegaSearchResult{ - SearchResult: result, - Engine: eng.Name(), - } - allResults = append(allResults, megaResult) - } - mu.Unlock() - }(engine) - } - - wg.Wait() - return allResults -} - -// deduplicateMegaResults deduplicates results while preserving engine information -func (s *Server) deduplicateMegaResults(results []MegaSearchResult) []MegaSearchResult { - urlMap := make(map[string]MegaSearchResult) - - // Process results and keep the first occurrence of each URL - for _, result := range results { - if result.URL == "" { - continue - } - - if _, exists := urlMap[result.URL]; !exists { - urlMap[result.URL] = result + if s.opts.AllowEndpointFallback { + res, usedEngine, searchErr = s.resilient.SearchWithFallback(engine, q) + } else { + res, usedEngine, searchErr = s.resilient.SearchPrimary(engine, q) } } - // Convert map back to slice and sort by rank - var deduped []MegaSearchResult - for _, result := range urlMap { - deduped = append(deduped, result) + if searchErr != nil { + errToReturn := searchErr + switch searchErr { + case ErrCaptcha: + errToReturn = fmt.Errorf("captcha found, please stop sending requests for a while: %w", searchErr) + case ErrSearchTimeout: + errToReturn = fmt.Errorf("%s", searchErr) + } + logrus.Errorf("Error during resilient %s %s: %s", engine.Name(), action, searchErr) + return fiber.NewError(fiber.StatusServiceUnavailable, errToReturn.Error()) } - // Sort by rank - sort.Slice(deduped, func(i, j int) bool { - return deduped[i].Rank < deduped[j].Rank - }) + if usedEngine != "" && usedEngine != engine.Name() { + c.Set("X-Fallback-Engine", usedEngine) + } - return deduped + logrus.Infof("Successfully completed SERP %s using %s, returned %d results", action, usedEngine, len(res)) + return c.JSON(res) } type HealthStatus struct { @@ -398,9 +176,26 @@ func (s *Server) handleHealthCheck(c *fiber.Ctx) error { for _, engine := range s.searchEngines { status := "ready" + isAvailable := true if !engine.IsInitialized() { status = "not_initialized" - } else { + isAvailable = false + } + + for _, cbStat := range s.resilient.GetCircuitBreakerStats() { + engineName, _ := cbStat["engine"].(string) + if engineName != engine.Name() { + continue + } + circuitState, _ := cbStat["state"].(string) + if circuitState == "open" { + status = "circuit_open" + isAvailable = false + } + break + } + + if isAvailable { availableEngines++ } @@ -440,6 +235,215 @@ func (s *Server) handleHealthCheck(c *fiber.Ctx) error { return c.JSON(health) } +func (s *Server) handleResilienceStats(c *fiber.Ctx) error { + return c.JSON(map[string]interface{}{ + "circuit_breakers": s.resilient.GetCircuitBreakerStats(), + }) +} + +type MegaSearchResult struct { + SearchResult + Engine string `json:"engine"` +} + +func (s *Server) handleMegaSearch(c *fiber.Ctx) error { + q := Query{} + if err := q.InitFromContext(c); err != nil { + logrus.Errorf("Error while setting megasearch query: %s", err) + return err + } + + enginesToUse := s.resolveEngines(c.Query("engines", "")) + if len(enginesToUse) == 0 { + return fiber.NewError(fiber.StatusBadRequest, "No valid search engines specified") + } + + engineNames := make([]string, len(enginesToUse)) + for i, engine := range enginesToUse { + engineNames[i] = engine.Name() + } + logrus.Infof("Starting SERP megasearch request using engines: %s for query: %s", strings.Join(engineNames, ", "), q.Text) + + results := s.resilient.SearchAllParallel(q, enginesToUse) + dedupedResults := s.deduplicateMegaResults(results) + + logrus.Infof("Successfully completed SERP megasearch using %d engines, returned %d deduplicated results", len(enginesToUse), len(dedupedResults)) + return c.JSON(dedupedResults) +} + +func (s *Server) handleMegaImage(c *fiber.Ctx) error { + q := Query{} + if err := q.InitFromContext(c); err != nil { + logrus.Errorf("Error while setting megasearch image query: %s", err) + return err + } + + enginesToUse := s.resolveEngines(c.Query("engines", "")) + if len(enginesToUse) == 0 { + return fiber.NewError(fiber.StatusBadRequest, "No valid search engines specified") + } + + engineNames := make([]string, len(enginesToUse)) + for i, engine := range enginesToUse { + engineNames[i] = engine.Name() + } + logrus.Infof("Starting SERP megasearch image request using engines: %s for query: %s", strings.Join(engineNames, ", "), q.Text) + + results := s.resilient.SearchAllImageParallel(q, enginesToUse) + dedupedResults := s.deduplicateMegaResults(results) + + logrus.Infof("Successfully completed SERP megasearch image using %d engines, returned %d deduplicated results", len(enginesToUse), len(dedupedResults)) + return c.JSON(dedupedResults) +} + +func (s *Server) handleListEngines(c *fiber.Ctx) error { + var engines []map[string]interface{} + + for _, engine := range s.searchEngines { + engineInfo := map[string]interface{}{ + "name": engine.Name(), + "initialized": engine.IsInitialized(), + } + + for _, cbStat := range s.resilient.GetCircuitBreakerStats() { + engineName, _ := cbStat["engine"].(string) + if engineName == engine.Name() { + engineInfo["circuit_state"] = cbStat["state"] + break + } + } + + engines = append(engines, engineInfo) + } + + return c.JSON(map[string]interface{}{ + "engines": engines, + "total": len(engines), + }) +} + +func (s *Server) resolveEngines(enginesParam string) []SearchEngine { + if enginesParam == "" { + return s.searchEngines + } + + var enginesToUse []SearchEngine + engineNames := strings.Split(enginesParam, ",") + for _, engineName := range engineNames { + engineName = strings.TrimSpace(strings.ToLower(engineName)) + for _, engine := range s.searchEngines { + if strings.ToLower(engine.Name()) == engineName { + enginesToUse = append(enginesToUse, engine) + break + } + } + } + return enginesToUse +} + +func (s *Server) searchSelectedEngines(q Query, engines []SearchEngine) []MegaSearchResult { + var wg sync.WaitGroup + var mu sync.Mutex + var allResults []MegaSearchResult + + for _, engine := range engines { + wg.Add(1) + go func(eng SearchEngine) { + defer wg.Done() + + limiter := eng.GetRateLimiter() + if limiter != nil { + err := limiter.Wait(context.Background()) + if err != nil { + logrus.Errorf("Ratelimiter error during %s megasearch: %s", eng.Name(), err) + } + } + + results, err := eng.Search(q) + if err != nil { + logrus.Errorf("Error during %s megasearch: %s", eng.Name(), err) + return + } + + mu.Lock() + for _, result := range results { + megaResult := MegaSearchResult{ + SearchResult: result, + Engine: eng.Name(), + } + allResults = append(allResults, megaResult) + } + mu.Unlock() + }(engine) + } + + wg.Wait() + return allResults +} + +func (s *Server) searchSelectedEnginesImage(q Query, engines []SearchEngine) []MegaSearchResult { + var wg sync.WaitGroup + var mu sync.Mutex + var allResults []MegaSearchResult + + for _, engine := range engines { + wg.Add(1) + go func(eng SearchEngine) { + defer wg.Done() + + limiter := eng.GetRateLimiter() + if limiter != nil { + err := limiter.Wait(context.Background()) + if err != nil { + logrus.Errorf("Ratelimiter error during %s megasearch image: %s", eng.Name(), err) + } + } + + results, err := eng.SearchImage(q) + if err != nil { + logrus.Errorf("Error during %s megasearch image: %s", eng.Name(), err) + return + } + + mu.Lock() + for _, result := range results { + megaResult := MegaSearchResult{ + SearchResult: result, + Engine: eng.Name(), + } + allResults = append(allResults, megaResult) + } + mu.Unlock() + }(engine) + } + + wg.Wait() + return allResults +} + +func (s *Server) deduplicateMegaResults(results []MegaSearchResult) []MegaSearchResult { + urlMap := make(map[string]MegaSearchResult) + + for _, result := range results { + if result.URL == "" { + continue + } + if _, exists := urlMap[result.URL]; !exists { + urlMap[result.URL] = result + } + } + + var deduped []MegaSearchResult + for _, result := range urlMap { + deduped = append(deduped, result) + } + + sort.Slice(deduped, func(i, j int) bool { + return deduped[i].Rank < deduped[j].Rank + }) + return deduped +} + func (s *Server) Listen() error { return s.app.Listen(s.addr) } diff --git a/core/server_test.go b/core/server_test.go index a431265..8f0471e 100644 --- a/core/server_test.go +++ b/core/server_test.go @@ -1,62 +1,225 @@ package core import ( - "fmt" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "sync" "testing" "time" "golang.org/x/time/rate" ) -var ( - servHost = "127.0.0.1" - servPort = 7070 - servAddr = fmt.Sprintf("%s:%d", servHost, servPort) -) +type engineMock struct { + name string + initialized bool + limiter *rate.Limiter + searchFn func(Query) ([]SearchResult, error) + imageFn func(Query) ([]SearchResult, error) -type SeMock struct { - EngineName string + mu sync.Mutex + searchCalls int } -func (s SeMock) Name() string { - return s.EngineName +func (e *engineMock) Name() string { return e.name } +func (e *engineMock) IsInitialized() bool { + return e.initialized } -func (SeMock) IsInitialized() bool { - return true -} -func (s SeMock) Search(q Query) (res []SearchResult, err error) { - return []SearchResult{{Title: s.EngineName}}, nil -} -func (s SeMock) SearchImage(q Query) (res []SearchResult, err error) { - return []SearchResult{{Title: s.EngineName}}, nil -} -func (s SeMock) GetRateLimiter() *rate.Limiter { - return nil +func (e *engineMock) GetRateLimiter() *rate.Limiter { return e.limiter } + +func (e *engineMock) Search(q Query) ([]SearchResult, error) { + e.mu.Lock() + e.searchCalls++ + e.mu.Unlock() + if e.searchFn != nil { + return e.searchFn(q) + } + return []SearchResult{{Rank: 1, URL: "https://example.com/" + e.name, Title: e.name}}, nil } -func TestCreateServer(t *testing.T) { - se1 := SeMock{"mock_engine_1"} - se2 := SeMock{"mock_engine_2"} +func (e *engineMock) SearchImage(q Query) ([]SearchResult, error) { + if e.imageFn != nil { + return e.imageFn(q) + } + return []SearchResult{{Rank: 1, URL: "https://img.example.com/" + e.name, Title: e.name}}, nil +} - server := NewServer(servHost, servPort, se1, se2) - - go func() { - time.Sleep(1 * time.Second) - server.Shutdown() - }() - - err := server.Listen() +func request(t *testing.T, s *Server, path string) *http.Response { + t.Helper() + req := httptest.NewRequest(http.MethodGet, path, nil) + resp, err := s.app.Test(req, -1) if err != nil { - t.Fatalf("Error failed initializing browser: %s", err) + t.Fatalf("request failed for %s: %v", path, err) + } + return resp +} + +func TestHealthEndpointStatusSemantics(t *testing.T) { + ready := &engineMock{name: "google", initialized: true, limiter: rate.NewLimiter(rate.Every(time.Second), 1)} + notReady := &engineMock{name: "yandex", initialized: false, limiter: rate.NewLimiter(rate.Every(time.Second), 1)} + + srv := NewServerWithOptions("127.0.0.1", 7070, DefaultServerOptions(), ready, notReady) + resp := request(t, srv, "/health") + if resp.StatusCode != http.StatusOK { + t.Fatalf("expected degraded health to return 200, got %d", resp.StatusCode) + } + + var health HealthStatus + if err := json.NewDecoder(resp.Body).Decode(&health); err != nil { + t.Fatalf("decode health response: %v", err) + } + if health.Status != "degraded" { + t.Fatalf("expected degraded status, got %s", health.Status) + } + + down := &engineMock{name: "google", initialized: false, limiter: rate.NewLimiter(rate.Every(time.Second), 1)} + srvUnhealthy := NewServerWithOptions("127.0.0.1", 7071, DefaultServerOptions(), down) + resp = request(t, srvUnhealthy, "/health") + if resp.StatusCode != http.StatusServiceUnavailable { + t.Fatalf("expected unhealthy health to return 503, got %d", resp.StatusCode) } } -// func TestServerSearch(t *testing.T) { -// se := SeMock{"mock_engine"} -// server := NewServer("127.0.0.1", 7070, se) +func TestDedicatedEndpointNoFallbackByDefault(t *testing.T) { + primary := &engineMock{ + name: "google", + initialized: true, + searchFn: func(q Query) ([]SearchResult, error) { + return nil, errors.New("primary failed") + }, + } + fallback := &engineMock{name: "yandex", initialized: true} -// go func() { -// server.Listen() -// }() + opts := DefaultServerOptions() + opts.AllowEndpointFallback = false + opts.Resilience.Retry.MaxRetries = 0 + srv := NewServerWithOptions("127.0.0.1", 7072, opts, primary, fallback) -// } + resp := request(t, srv, "/google/search?text=golang") + if resp.StatusCode != http.StatusServiceUnavailable { + t.Fatalf("expected 503 when primary fails and fallback disabled, got %d", resp.StatusCode) + } + if got := resp.Header.Get("X-Fallback-Engine"); got != "" { + t.Fatalf("unexpected fallback header: %s", got) + } + if fallback.searchCalls != 0 { + t.Fatalf("fallback engine should not be called, got %d calls", fallback.searchCalls) + } +} + +func TestResilienceStatsContainsRetryInWhenCircuitOpen(t *testing.T) { + primary := &engineMock{ + name: "google", + initialized: true, + searchFn: func(q Query) ([]SearchResult, error) { + return nil, errors.New("forced failure") + }, + } + + opts := DefaultServerOptions() + opts.Resilience.Retry.MaxRetries = 0 + opts.Resilience.CircuitBreaker.FailureThreshold = 1 + opts.Resilience.CircuitBreaker.RecoveryTimeout = 5 * time.Minute + srv := NewServerWithOptions("127.0.0.1", 7074, opts, primary) + + _ = request(t, srv, "/google/search?text=golang") + statsResp := request(t, srv, "/resilience/stats") + if statsResp.StatusCode != http.StatusOK { + t.Fatalf("expected stats endpoint to return 200, got %d", statsResp.StatusCode) + } + + var stats map[string]interface{} + if err := json.NewDecoder(statsResp.Body).Decode(&stats); err != nil { + t.Fatalf("decode stats response: %v", err) + } + + breakers, ok := stats["circuit_breakers"].([]interface{}) + if !ok || len(breakers) == 0 { + t.Fatalf("expected at least one circuit breaker entry, got %#v", stats["circuit_breakers"]) + } + + first := breakers[0].(map[string]interface{}) + if state, _ := first["state"].(string); state != "open" { + t.Fatalf("expected circuit to be open, got %q", state) + } + retryIn, ok := first["retry_in"].(float64) + if !ok { + t.Fatalf("expected retry_in number in JSON response, got %T", first["retry_in"]) + } + if retryIn <= 0 { + t.Fatalf("expected retry_in to be present when circuit is open, got %v", first["retry_in"]) + } +} + +func TestRetryAppliesRateLimiterOnEachAttempt(t *testing.T) { + engine := &engineMock{ + name: "google", + initialized: true, + limiter: rate.NewLimiter(rate.Every(120*time.Millisecond), 1), + searchFn: func(q Query) ([]SearchResult, error) { + return nil, errors.New("always fail") + }, + } + + opts := DefaultServerOptions() + opts.Resilience.Retry.MaxRetries = 2 + opts.Resilience.Retry.InitialBackoff = 0 + opts.Resilience.Retry.MaxBackoff = 0 + opts.Resilience.Retry.BackoffFactor = 1 + srv := NewServerWithOptions("127.0.0.1", 7075, opts, engine) + + start := time.Now() + resp := request(t, srv, "/google/search?text=golang") + elapsed := time.Since(start) + if resp.StatusCode != http.StatusServiceUnavailable { + t.Fatalf("expected failure response, got %d", resp.StatusCode) + } + + if engine.searchCalls != 3 { + t.Fatalf("expected 3 attempts (1 + 2 retries), got %d", engine.searchCalls) + } + if elapsed < 200*time.Millisecond { + t.Fatalf("expected limiter to delay retries, elapsed only %s", elapsed) + } +} + +// Server-level CORS tests are intentionally smoke-level: +// they verify middleware registration and option wiring, not header semantics. +func TestServerOptions_WiresCustomCORSConfig(t *testing.T) { + engine := &engineMock{name: "google", initialized: true} + + opts := DefaultServerOptions() + opts.CORS = CORSConfig{ + AllowOrigins: "https://client.local", + AllowMethods: "GET,OPTIONS", + AllowHeaders: "Authorization", + MaxAge: 600, + } + + srv := NewServerWithOptions("127.0.0.1", 7076, opts, engine) + resp := request(t, srv, "/health") + if resp.StatusCode != http.StatusOK { + t.Fatalf("expected 200, got %d", resp.StatusCode) + } + if got := resp.Header.Get("Access-Control-Allow-Origin"); got != "https://client.local" { + t.Fatalf("unexpected allow-origin: %q", got) + } +} + +func TestServerOptions_DisableCORSMiddleware(t *testing.T) { + engine := &engineMock{name: "google", initialized: true} + + opts := DefaultServerOptions() + opts.EnableCORS = false + + srv := NewServerWithOptions("127.0.0.1", 7077, opts, engine) + resp := request(t, srv, "/health") + if resp.StatusCode != http.StatusOK { + t.Fatalf("expected 200, got %d", resp.StatusCode) + } + if got := resp.Header.Get("Access-Control-Allow-Origin"); got != "" { + t.Fatalf("expected CORS headers to be absent when disabled, got allow-origin=%q", got) + } +}