mirror of
https://github.com/karust/openserp.git
synced 2026-08-05 16:53:54 +08:00
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:
2
.gitignore
vendored
2
.gitignore
vendored
@@ -25,3 +25,5 @@ logs.txt
|
||||
.release
|
||||
core/test/
|
||||
.aider*
|
||||
.gocache/
|
||||
openserp
|
||||
|
||||
57
cmd/root.go
57
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")
|
||||
}
|
||||
|
||||
29
cmd/serve.go
29
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 {
|
||||
|
||||
44
config.yaml
44
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
|
||||
|
||||
210
core/circuit_breaker.go
Normal file
210
core/circuit_breaker.go
Normal 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")
|
||||
165
core/circuit_breaker_test.go
Normal file
165
core/circuit_breaker_test.go
Normal 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
144
core/middleware.go
Normal 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
101
core/middleware_test.go
Normal 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
232
core/resilient.go
Normal 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
91
core/retry.go
Normal 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
86
core/retry_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
610
core/server.go
610
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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user