feat: tunnel JWT auth
Tunnel method now requires Bearer JWT (HMAC-HS256) in gRPC metadata.
- internal/auth: pure-stdlib verifier + issuer + revoked-jti list
- internal/tunnel: per-target rate limiter + egress allowlist with
wildcard support + crypto/rand frame padding to {1KB, 4KB, 16KB}
- cmd/server reads /etc/stats-gateway/{jwt.key,tunnel.yaml,revoked.txt}
via --config-dir flag
- cmd/issue-token Go binary wraps auth.Issue
- configs/tunnel.yaml example with projectshitpost wildcard
Tests: 13 PASS across 5 packages.
This commit is contained in:
@@ -0,0 +1,167 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user