package ratelimit

import (
	"context"
	"math/rand"
	"sync"
	"time"
)

// bucket is a single token bucket state (guarded by Limiter.mu).
type bucket struct {
	tokens float64
	last   time.Time
}

// Limiter throttles outgoing requests with two token buckets: one global
// bucket shared by every caller (all services, all proxies) and one bucket per
// key (a proxy URL, so each exit IP has its own budget). A zero or negative
// rate disables that layer; callers then rely on the worker pool size alone.
type Limiter struct {
	globalRate  float64
	globalBurst int
	keyRate     float64
	keyBurst    int

	mu     sync.Mutex
	global bucket
	keys   map[string]*bucket
}

// New builds a limiter with a global request rate plus an optional per-key
// rate. Both are requests per second; 0 or less disables that layer.
func New(globalRatePerSec, keyRatePerSec float64) *Limiter {
	return &Limiter{
		globalRate:  globalRatePerSec,
		globalBurst: burstFor(globalRatePerSec),
		keyRate:     keyRatePerSec,
		keyBurst:    burstFor(keyRatePerSec),
		keys:        make(map[string]*bucket),
	}
}

func burstFor(rate float64) int {
	b := int(rate) + 1
	if b < 1 {
		b = 1
	}
	return b
}

// Wait blocks until both the global bucket and the bucket for key have budget.
// The key is the proxy URL ("direct" when no proxy is configured), so one
// slow or blocked IP cannot stall the whole pool while the global cap still
// bounds total request rate.
func (l *Limiter) Wait(ctx context.Context, key string) error {
	if err := l.waitBucket(ctx, &l.global, l.globalRate, l.globalBurst); err != nil {
		return err
	}
	l.mu.Lock()
	b := l.keys[key]
	if b == nil {
		b = &bucket{}
		l.keys[key] = b
	}
	l.mu.Unlock()
	return l.waitBucket(ctx, b, l.keyRate, l.keyBurst)
}

func (l *Limiter) waitBucket(ctx context.Context, b *bucket, rate float64, burst int) error {
	if rate <= 0 {
		return ctx.Err()
	}
	for {
		l.mu.Lock()
		now := time.Now()
		tokens, last := b.tokens, b.last
		if last.IsZero() {
			tokens = float64(burst)
			last = now
		}
		elapsed := now.Sub(last).Seconds()
		tokens += elapsed * rate
		if tokens > float64(burst) {
			tokens = float64(burst)
		}
		if tokens >= 1 {
			tokens--
			b.tokens, b.last = tokens, now
			l.mu.Unlock()
			return ctx.Err()
		}
		wait := time.Duration((1-tokens)/rate*float64(time.Second)) + time.Duration(rand.Intn(50))*time.Millisecond
		b.tokens, b.last = tokens, now
		l.mu.Unlock()

		select {
		case <-ctx.Done():
			return ctx.Err()
		case <-time.After(wait):
		}
	}
}

// JitterSleep sleeps a random duration between min and max (Python parity for
// anti-bot pacing where no proxy-level rate limit is configured).
func JitterSleep(minSec, maxSec float64) {
	d := minSec + rand.Float64()*(maxSec-minSec)
	time.Sleep(time.Duration(d * float64(time.Second)))
}
