diff --git a/Makefile b/Makefile index 6931cfd..ec025a7 100644 --- a/Makefile +++ b/Makefile @@ -4,6 +4,7 @@ build: go build -o bin/stats-gateway ./cmd/server go build -o bin/stats-collector ./cmd/collector go build -o bin/stats-snapshot ./cmd/snapshot + go build -o bin/stats-issue-token ./cmd/issue-token test: go test ./... diff --git a/cmd/issue-token/main.go b/cmd/issue-token/main.go new file mode 100644 index 0000000..a44efd2 --- /dev/null +++ b/cmd/issue-token/main.go @@ -0,0 +1,33 @@ +package main + +import ( + "bytes" + "flag" + "fmt" + "log" + "os" + "time" + + "github.com/bergabruh/stats-gateway/internal/auth" +) + +func main() { + keyPath := flag.String("key", "/etc/stats-gateway/jwt.key", "path to HMAC secret key file") + subject := flag.String("sub", "", "client-id (required)") + days := flag.Int("days", 90, "token validity in days") + flag.Parse() + + if *subject == "" { + log.Fatal("--sub required") + } + + secret, err := os.ReadFile(*keyPath) + if err != nil { + log.Fatalf("read key: %v", err) + } + tok, err := auth.Issue(bytes.TrimSpace(secret), *subject, "tunnel", time.Duration(*days)*24*time.Hour) + if err != nil { + log.Fatal(err) + } + fmt.Println(tok) +} diff --git a/cmd/server/main.go b/cmd/server/main.go index 0a5d8eb..6fb51de 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -1,20 +1,27 @@ package main import ( + "bytes" "flag" "log" "net" + "os" + "path/filepath" "google.golang.org/grpc" + "gopkg.in/yaml.v3" statsv1 "github.com/bergabruh/stats-gateway/gen/stats/v1" + "github.com/bergabruh/stats-gateway/internal/auth" "github.com/bergabruh/stats-gateway/internal/service" "github.com/bergabruh/stats-gateway/internal/storage" + "github.com/bergabruh/stats-gateway/internal/tunnel" ) func main() { addr := flag.String("addr", "127.0.0.1:50051", "listen address") dbPath := flag.String("db", "/var/lib/stats-gateway/stats.db", "sqlite path") + cfgDir := flag.String("config-dir", "/etc/stats-gateway", "config dir (jwt.key, tunnel.yaml, revoked.txt)") flag.Parse() st, err := storage.Open(*dbPath) @@ -23,15 +30,41 @@ func main() { } defer st.Close() + secret, err := os.ReadFile(filepath.Join(*cfgDir, "jwt.key")) + if err != nil { + log.Fatalf("read jwt.key: %v", err) + } + secret = bytes.TrimSpace(secret) + + cfgBytes, err := os.ReadFile(filepath.Join(*cfgDir, "tunnel.yaml")) + if err != nil { + log.Fatalf("read tunnel.yaml: %v", err) + } + var tunnelCfg struct { + Egress tunnel.EgressConfig `yaml:"egress"` + } + if err := yaml.Unmarshal(cfgBytes, &tunnelCfg); err != nil { + log.Fatalf("parse tunnel.yaml: %v", err) + } + + revoked, err := auth.LoadRevokedFile(filepath.Join(*cfgDir, "revoked.txt")) + if err != nil { + log.Fatalf("load revoked: %v", err) + } + + verifier := auth.NewVerifier(secret, revoked) + handler := tunnel.NewHandler(tunnelCfg.Egress) + lis, err := net.Listen("tcp", *addr) if err != nil { log.Fatalf("listen: %v", err) } srv := grpc.NewServer() - statsv1.RegisterSpeedStatusServer(srv, service.New(st)) + statsv1.RegisterSpeedStatusServer(srv, service.New(st, verifier, handler)) - log.Printf("stats-gateway listening on %s", *addr) + log.Printf("stats-gateway listening on %s (egress=%d targets, rate=%.0fmbps)", + *addr, len(tunnelCfg.Egress.Allowed), tunnelCfg.Egress.RateMbps) if err := srv.Serve(lis); err != nil { log.Fatalf("serve: %v", err) } diff --git a/configs/tunnel.yaml b/configs/tunnel.yaml new file mode 100644 index 0000000..12af34e --- /dev/null +++ b/configs/tunnel.yaml @@ -0,0 +1,6 @@ +egress: + allowed: + - "192.0.2.10:8443" # FL xray inbound (production) + - "*.projectshitpost.fun:443" # any subdomain HTTPS + # TODO: add PL1 IP:port once confirmed (controller to fill in during deploy) + rate_mbps: 50 diff --git a/go.mod b/go.mod index f23b1d8..3765fd2 100644 --- a/go.mod +++ b/go.mod @@ -4,8 +4,10 @@ go 1.25.0 require ( github.com/mattn/go-sqlite3 v1.14.44 + golang.org/x/time v0.15.0 google.golang.org/grpc v1.81.0 google.golang.org/protobuf v1.36.11 + gopkg.in/yaml.v3 v3.0.1 ) require ( diff --git a/go.sum b/go.sum index 1b580a2..7673f52 100644 --- a/go.sum +++ b/go.sum @@ -30,6 +30,8 @@ golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo= golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk= golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA= +golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U= +golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= google.golang.org/genproto/googleapis/rpc v0.0.0-20260226221140-a57be14db171 h1:ggcbiqK8WWh6l1dnltU4BgWGIGo+EVYxCaAPih/zQXQ= @@ -38,3 +40,7 @@ google.golang.org/grpc v1.81.0 h1:W3G9N3KQf3BU+YuCtGKJk0CmxQNbAISICD/9AORxLIw= google.golang.org/grpc v1.81.0/go.mod h1:xGH9GfzOyMTGIOXBJmXt+BX/V0kcdQbdcuwQ/zNw42I= google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/auth/jwt.go b/internal/auth/jwt.go new file mode 100644 index 0000000..949c74c --- /dev/null +++ b/internal/auth/jwt.go @@ -0,0 +1,154 @@ +package auth + +import ( + "bufio" + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "os" + "strings" + "time" +) + +// Claims holds the JWT payload fields used by this service. +type Claims struct { + Subject string `json:"sub"` + Scope string `json:"scope"` + IssuedAt int64 `json:"iat"` + Expires int64 `json:"exp"` + JTI string `json:"jti"` +} + +// Verifier validates HMAC-HS256 tokens against a secret and optional revocation list. +type Verifier struct { + secret []byte + revoked map[string]bool +} + +// NewVerifier creates a Verifier. If revoked is nil an empty map is used. +func NewVerifier(secret []byte, revoked map[string]bool) *Verifier { + if revoked == nil { + revoked = map[string]bool{} + } + return &Verifier{secret: secret, revoked: revoked} +} + +// Verify checks signature, expiry, scope=="tunnel", and revocation list. +// Returns parsed Claims on success. +func (v *Verifier) Verify(token string) (*Claims, error) { + parts := strings.Split(token, ".") + if len(parts) != 3 { + return nil, errors.New("malformed token") + } + + signingInput := parts[0] + "." + parts[1] + expectedSig, err := base64URLDecode(parts[2]) + if err != nil { + return nil, fmt.Errorf("decode sig: %w", err) + } + mac := hmac.New(sha256.New, v.secret) + mac.Write([]byte(signingInput)) + if !hmac.Equal(mac.Sum(nil), expectedSig) { + return nil, errors.New("bad signature") + } + + headerJSON, err := base64URLDecode(parts[0]) + if err != nil { + return nil, fmt.Errorf("decode header: %w", err) + } + var header struct { + Alg string `json:"alg"` + Typ string `json:"typ"` + } + if err := json.Unmarshal(headerJSON, &header); err != nil { + return nil, fmt.Errorf("header: %w", err) + } + if header.Alg != "HS256" { + return nil, fmt.Errorf("alg %q not allowed", header.Alg) + } + + payload, err := base64URLDecode(parts[1]) + if err != nil { + return nil, fmt.Errorf("decode payload: %w", err) + } + var c Claims + if err := json.Unmarshal(payload, &c); err != nil { + return nil, fmt.Errorf("payload: %w", err) + } + + now := time.Now().Unix() + if c.Expires < now { + return nil, errors.New("token expired") + } + if c.Scope != "tunnel" { + return nil, fmt.Errorf("scope %q not allowed", c.Scope) + } + if v.revoked[c.JTI] { + return nil, errors.New("token revoked") + } + return &c, nil +} + +// Issue creates a new HS256 token. ttl <= 0 produces an already-expired token (useful for tests). +func Issue(secret []byte, subject, scope string, ttl time.Duration) (string, error) { + jtiBytes := make([]byte, 8) + if _, err := rand.Read(jtiBytes); err != nil { + return "", err + } + return IssueWithJTI(secret, subject, scope, ttl, fmt.Sprintf("%x", jtiBytes)) +} + +// IssueWithJTI creates a token with a caller-supplied jti (useful for revocation tests). +func IssueWithJTI(secret []byte, subject, scope string, ttl time.Duration, jti string) (string, error) { + now := time.Now().Unix() + c := Claims{ + Subject: subject, + Scope: scope, + IssuedAt: now, + Expires: now + int64(ttl.Seconds()), + JTI: jti, + } + header := `{"alg":"HS256","typ":"JWT"}` + payload, err := json.Marshal(c) + if err != nil { + return "", err + } + + h := base64URLEncode([]byte(header)) + p := base64URLEncode(payload) + signingInput := h + "." + p + mac := hmac.New(sha256.New, secret) + mac.Write([]byte(signingInput)) + sig := base64URLEncode(mac.Sum(nil)) + return signingInput + "." + sig, nil +} + +// LoadRevokedFile reads one jti per line from path. Lines starting with '#' +// or empty are ignored. Missing file → empty map (no error). +func LoadRevokedFile(path string) (map[string]bool, error) { + f, err := os.Open(path) + if err != nil { + if os.IsNotExist(err) { + return map[string]bool{}, nil + } + return nil, err + } + defer f.Close() + out := map[string]bool{} + sc := bufio.NewScanner(f) + for sc.Scan() { + line := strings.TrimSpace(sc.Text()) + if line == "" || strings.HasPrefix(line, "#") { + continue + } + out[line] = true + } + return out, sc.Err() +} + +func base64URLEncode(b []byte) string { return base64.RawURLEncoding.EncodeToString(b) } +func base64URLDecode(s string) ([]byte, error) { return base64.RawURLEncoding.DecodeString(s) } diff --git a/internal/auth/jwt_test.go b/internal/auth/jwt_test.go new file mode 100644 index 0000000..77dd3e9 --- /dev/null +++ b/internal/auth/jwt_test.go @@ -0,0 +1,75 @@ +package auth + +import ( + "os" + "path/filepath" + "testing" + "time" +) + +func TestVerify_ValidToken(t *testing.T) { + v := NewVerifier([]byte("test-secret"), nil) + tok, _ := Issue([]byte("test-secret"), "alice", "tunnel", time.Hour) + claims, err := v.Verify(tok) + if err != nil { + t.Fatalf("expected ok, got %v", err) + } + if claims.Subject != "alice" || claims.Scope != "tunnel" { + t.Errorf("claims wrong: %+v", claims) + } +} + +func TestVerify_WrongSecret(t *testing.T) { + v := NewVerifier([]byte("server-secret"), nil) + tok, _ := Issue([]byte("attacker-secret"), "eve", "tunnel", time.Hour) + if _, err := v.Verify(tok); err == nil { + t.Fatal("expected verification fail") + } +} + +func TestVerify_Expired(t *testing.T) { + v := NewVerifier([]byte("test-secret"), nil) + tok, _ := Issue([]byte("test-secret"), "alice", "tunnel", -time.Minute) + if _, err := v.Verify(tok); err == nil { + t.Fatal("expected expired") + } +} + +func TestVerify_WrongScope(t *testing.T) { + v := NewVerifier([]byte("test-secret"), nil) + tok, _ := Issue([]byte("test-secret"), "alice", "read-only", time.Hour) + if _, err := v.Verify(tok); err == nil { + t.Fatal("expected scope mismatch") + } +} + +func TestVerify_RevokedJTI(t *testing.T) { + revoked := map[string]bool{"abc123": true} + v := NewVerifier([]byte("test-secret"), revoked) + tok, _ := IssueWithJTI([]byte("test-secret"), "alice", "tunnel", time.Hour, "abc123") + if _, err := v.Verify(tok); err == nil { + t.Fatal("expected revoked") + } +} + +func TestLoadRevokedFile_Missing(t *testing.T) { + m, err := LoadRevokedFile(filepath.Join(t.TempDir(), "nonexistent")) + if err != nil { + t.Fatal(err) + } + if len(m) != 0 { + t.Errorf("expected empty, got %v", m) + } +} + +func TestLoadRevokedFile_ParsesEntries(t *testing.T) { + tmp := filepath.Join(t.TempDir(), "r.txt") + os.WriteFile(tmp, []byte("# comment\nabc\n\nxyz\n"), 0644) + m, err := LoadRevokedFile(tmp) + if err != nil { + t.Fatal(err) + } + if !m["abc"] || !m["xyz"] || len(m) != 2 { + t.Errorf("got %v", m) + } +} diff --git a/internal/service/service.go b/internal/service/service.go index 23ab978..7e9137b 100644 --- a/internal/service/service.go +++ b/internal/service/service.go @@ -2,23 +2,30 @@ package service import ( "context" + "strings" "time" statsv1 "github.com/bergabruh/stats-gateway/gen/stats/v1" + "github.com/bergabruh/stats-gateway/internal/auth" "github.com/bergabruh/stats-gateway/internal/storage" + "github.com/bergabruh/stats-gateway/internal/tunnel" "google.golang.org/grpc/codes" + "google.golang.org/grpc/metadata" "google.golang.org/grpc/status" ) -// Service implements statsv1.SpeedStatusServer with public RecentResults -// and Aggregate methods. Tunnel returns Unimplemented until Task 11 wires -// auth + handler. +// Service implements statsv1.SpeedStatusServer. type Service struct { statsv1.UnimplementedSpeedStatusServer - store *storage.Store + store *storage.Store + verifier *auth.Verifier + tunnel *tunnel.Handler } -func New(st *storage.Store) *Service { return &Service{store: st} } +// New creates a Service with the given storage, JWT verifier, and tunnel handler. +func New(st *storage.Store, v *auth.Verifier, t *tunnel.Handler) *Service { + return &Service{store: st, verifier: v, tunnel: t} +} func (s *Service) RecentResults(f *statsv1.Filter, stream statsv1.SpeedStatus_RecentResultsServer) error { limit := int(f.Limit) @@ -87,3 +94,23 @@ func (s *Service) Aggregate(ctx context.Context, r *statsv1.Range) (*statsv1.His } return h, nil } + +// Tunnel authenticates the caller via Bearer JWT then delegates to the tunnel handler. +func (s *Service) Tunnel(stream statsv1.SpeedStatus_TunnelServer) error { + md, ok := metadata.FromIncomingContext(stream.Context()) + if !ok { + return status.Error(codes.Unauthenticated, "no metadata") + } + authH := md.Get("authorization") + if len(authH) == 0 { + return status.Error(codes.Unauthenticated, "no auth") + } + tok := strings.TrimPrefix(authH[0], "Bearer ") + if tok == authH[0] { + return status.Error(codes.Unauthenticated, "expected Bearer") + } + if _, err := s.verifier.Verify(tok); err != nil { + return status.Errorf(codes.Unauthenticated, "auth: %v", err) + } + return s.tunnel.Serve(stream) +} diff --git a/internal/service/service_test.go b/internal/service/service_test.go index 45208bf..3bed79f 100644 --- a/internal/service/service_test.go +++ b/internal/service/service_test.go @@ -8,19 +8,27 @@ import ( "time" "google.golang.org/grpc" + "google.golang.org/grpc/codes" "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/status" "google.golang.org/grpc/test/bufconn" statsv1 "github.com/bergabruh/stats-gateway/gen/stats/v1" + "github.com/bergabruh/stats-gateway/internal/auth" "github.com/bergabruh/stats-gateway/internal/probes" "github.com/bergabruh/stats-gateway/internal/storage" + "github.com/bergabruh/stats-gateway/internal/tunnel" ) func newTestServer(t *testing.T, st *storage.Store) (statsv1.SpeedStatusClient, func()) { t.Helper() + verifier := auth.NewVerifier([]byte("test-secret"), nil) + handler := tunnel.NewHandler(tunnel.EgressConfig{Allowed: []string{"127.0.0.1:0"}, RateMbps: 50}) + lis := bufconn.Listen(1024 * 1024) srv := grpc.NewServer() - statsv1.RegisterSpeedStatusServer(srv, New(st)) + statsv1.RegisterSpeedStatusServer(srv, New(st, verifier, handler)) go srv.Serve(lis) conn, err := grpc.NewClient("passthrough:///bufnet", @@ -104,7 +112,7 @@ func TestAggregate_BucketCounts(t *testing.T) { } } -func TestTunnel_Unimplemented(t *testing.T) { +func TestTunnel_RejectsMissingAuth(t *testing.T) { dir := t.TempDir() st, _ := storage.Open(filepath.Join(dir, "test.db")) defer st.Close() @@ -118,7 +126,33 @@ func TestTunnel_Unimplemented(t *testing.T) { stream.Send(&statsv1.Frame{Kind: statsv1.Frame_HELLO, Target: "1.2.3.4:80"}) _, err = stream.Recv() if err == nil { - t.Fatal("expected Unimplemented error") + t.Fatal("expected unauthenticated") + } + st_, ok := status.FromError(err) + if !ok || st_.Code() != codes.Unauthenticated { + t.Errorf("expected Unauthenticated, got %v", err) + } +} + +func TestTunnel_RejectsBadToken(t *testing.T) { + dir := t.TempDir() + st, _ := storage.Open(filepath.Join(dir, "test.db")) + defer st.Close() + client, cleanup := newTestServer(t, st) + defer cleanup() + + ctx := metadata.AppendToOutgoingContext(context.Background(), "authorization", "Bearer not.a.real.token") + stream, err := client.Tunnel(ctx) + if err != nil { + t.Fatal(err) + } + stream.Send(&statsv1.Frame{Kind: statsv1.Frame_HELLO, Target: "1.2.3.4:80"}) + _, err = stream.Recv() + if err == nil { + t.Fatal("expected auth fail") + } + st_, ok := status.FromError(err) + if !ok || st_.Code() != codes.Unauthenticated { + t.Errorf("expected Unauthenticated, got %v", err) } - // Just check it's an error; the exact code/message will be checked indirectly when Task 11 wires real impl. } diff --git a/internal/tunnel/tunnel.go b/internal/tunnel/tunnel.go new file mode 100644 index 0000000..ced2ae0 --- /dev/null +++ b/internal/tunnel/tunnel.go @@ -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 +} diff --git a/internal/tunnel/tunnel_test.go b/internal/tunnel/tunnel_test.go new file mode 100644 index 0000000..4147145 --- /dev/null +++ b/internal/tunnel/tunnel_test.go @@ -0,0 +1,22 @@ +package tunnel + +import "testing" + +func TestAllowEgress(t *testing.T) { + cfg := EgressConfig{Allowed: []string{"192.0.2.10:8443", "*.projectshitpost.fun:443"}} + cases := []struct { + target string + want bool + }{ + {"192.0.2.10:8443", true}, + {"foo.projectshitpost.fun:443", true}, + {"projectshitpost.fun:443", false}, // wildcard "*.x" doesn't match bare "x" + {"8.8.8.8:53", false}, + {"192.0.2.10:22", false}, + } + for _, tc := range cases { + if got := cfg.Allow(tc.target); got != tc.want { + t.Errorf("Allow(%q) = %v, want %v", tc.target, got, tc.want) + } + } +}