Add retry/circuit breaker, configurable fallback and CORS.

Based in part on work from PR #21 by @Sai-Prashanth123, adapted and integrated with project-specific fixes.
This commit is contained in:
Rustem Kamalov
2026-03-25 23:31:41 +03:00
parent 0fcfc06baa
commit f219c78d84
13 changed files with 1650 additions and 364 deletions

2
.gitignore vendored
View File

@@ -25,3 +25,5 @@ logs.txt
.release
core/test/
.aider*
.gocache/
openserp

View File

@@ -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")
}

View File

@@ -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 {

View File

@@ -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

210
core/circuit_breaker.go Normal file
View File

@@ -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")

View File

@@ -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))
}
}

144
core/middleware.go Normal file
View File

@@ -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"
}
}

101
core/middleware_test.go Normal file
View File

@@ -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)
}
}

232
core/resilient.go Normal file
View File

@@ -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")

91
core/retry.go Normal file
View File

@@ -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)
}

86
core/retry_test.go Normal file
View File

@@ -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)
}
}
}

View File

@@ -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)
}

View File

@@ -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)
}
}