package tunnel import ( "crypto/rand" "errors" "fmt" "io" "net" "strings" "sync" "time" "golang.org/x/time/rate" statsv1 "github.com/bergabruh/stats-gateway/gen/stats/v1" ) // EgressConfig defines which targets are allowed and the per-target rate limit. type EgressConfig struct { Allowed []string `yaml:"allowed"` RateMbps float64 `yaml:"rate_mbps"` } // Allow returns true if target ("host:port") matches any entry in Allowed. // Wildcards: "*.suffix:port" matches "anything.suffix:port" but NOT bare "suffix:port". func (e *EgressConfig) Allow(target string) bool { for _, p := range e.Allowed { if p == target { return true } if strings.HasPrefix(p, "*.") { suffix := p[1:] // e.g. ".projectshitpost.fun:443" if strings.HasSuffix(target, suffix) && len(target) > len(suffix) { return true } } } return false } // Handler manages per-target rate limiters and dispatches tunnel sessions. type Handler struct { cfg EgressConfig mu sync.Mutex limiters map[string]*rate.Limiter } // NewHandler creates a Handler with the given egress configuration. func NewHandler(cfg EgressConfig) *Handler { return &Handler{cfg: cfg, limiters: map[string]*rate.Limiter{}} } func (h *Handler) limiter(target string) *rate.Limiter { h.mu.Lock() defer h.mu.Unlock() if l, ok := h.limiters[target]; ok { return l } bps := h.cfg.RateMbps * 1024 * 1024 / 8 if bps <= 0 { bps = 50 * 1024 * 1024 / 8 // default 50 Mbps } l := rate.NewLimiter(rate.Limit(bps), int(bps)) h.limiters[target] = l return l } // Serve runs the bidirectional tunnel after auth has already been validated. // First frame must be HELLO with target "host:port" within Allow(). func (h *Handler) Serve(stream statsv1.SpeedStatus_TunnelServer) error { first, err := stream.Recv() if err != nil { return fmt.Errorf("recv hello: %w", err) } if first.Kind != statsv1.Frame_HELLO { return errors.New("expected HELLO") } if !h.cfg.Allow(first.Target) { return fmt.Errorf("egress denied: %s", first.Target) } conn, err := net.DialTimeout("tcp", first.Target, 10*time.Second) if err != nil { return fmt.Errorf("dial %s: %w", first.Target, err) } defer conn.Close() lim := h.limiter(first.Target) errCh := make(chan error, 2) // stream → conn go func() { for { f, err := stream.Recv() if err != nil { errCh <- err return } if f.Kind == statsv1.Frame_CLOSE { errCh <- nil return } if len(f.Payload) > 0 { if err := lim.WaitN(stream.Context(), len(f.Payload)); err != nil { errCh <- err return } } if _, err := conn.Write(f.Payload); err != nil { errCh <- err return } } }() // conn → stream go func() { buf := make([]byte, 16*1024) for { n, err := conn.Read(buf) if n > 0 { if err := lim.WaitN(stream.Context(), n); err != nil { errCh <- err return } if sErr := stream.Send(&statsv1.Frame{ Kind: statsv1.Frame_DATA, Payload: append([]byte{}, buf[:n]...), Padding: paddingForSize(n), }); sErr != nil { errCh <- sErr return } } if err == io.EOF { errCh <- nil return } if err != nil { errCh <- err return } } }() return <-errCh } // paddingForSize rounds up to the nearest bucket of {1024, 4096, 16384} bytes // and fills the remainder with crypto/rand bytes. This blurs the packet-length // distribution for DPI/ML classifiers. func paddingForSize(n int) []byte { target := 16384 switch { case n <= 1024: target = 1024 case n <= 4096: target = 4096 } pad := target - n if pad <= 0 { return nil } out := make([]byte, pad) _, _ = rand.Read(out) return out }