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:
2026-05-05 21:28:37 +05:00
parent 0e3a3b4b0e
commit 731b60863a
12 changed files with 571 additions and 11 deletions
+1
View File
@@ -4,6 +4,7 @@ build:
go build -o bin/stats-gateway ./cmd/server go build -o bin/stats-gateway ./cmd/server
go build -o bin/stats-collector ./cmd/collector go build -o bin/stats-collector ./cmd/collector
go build -o bin/stats-snapshot ./cmd/snapshot go build -o bin/stats-snapshot ./cmd/snapshot
go build -o bin/stats-issue-token ./cmd/issue-token
test: test:
go test ./... go test ./...
+33
View File
@@ -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)
}
+35 -2
View File
@@ -1,20 +1,27 @@
package main package main
import ( import (
"bytes"
"flag" "flag"
"log" "log"
"net" "net"
"os"
"path/filepath"
"google.golang.org/grpc" "google.golang.org/grpc"
"gopkg.in/yaml.v3"
statsv1 "github.com/bergabruh/stats-gateway/gen/stats/v1" 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/service"
"github.com/bergabruh/stats-gateway/internal/storage" "github.com/bergabruh/stats-gateway/internal/storage"
"github.com/bergabruh/stats-gateway/internal/tunnel"
) )
func main() { func main() {
addr := flag.String("addr", "127.0.0.1:50051", "listen address") addr := flag.String("addr", "127.0.0.1:50051", "listen address")
dbPath := flag.String("db", "/var/lib/stats-gateway/stats.db", "sqlite path") 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() flag.Parse()
st, err := storage.Open(*dbPath) st, err := storage.Open(*dbPath)
@@ -23,15 +30,41 @@ func main() {
} }
defer st.Close() 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) lis, err := net.Listen("tcp", *addr)
if err != nil { if err != nil {
log.Fatalf("listen: %v", err) log.Fatalf("listen: %v", err)
} }
srv := grpc.NewServer() 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 { if err := srv.Serve(lis); err != nil {
log.Fatalf("serve: %v", err) log.Fatalf("serve: %v", err)
} }
+6
View File
@@ -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
+2
View File
@@ -4,8 +4,10 @@ go 1.25.0
require ( require (
github.com/mattn/go-sqlite3 v1.14.44 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/grpc v1.81.0
google.golang.org/protobuf v1.36.11 google.golang.org/protobuf v1.36.11
gopkg.in/yaml.v3 v3.0.1
) )
require ( require (
+6
View File
@@ -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/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 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk=
golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA= 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 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= 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= 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/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 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= 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=
+154
View File
@@ -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) }
+75
View File
@@ -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)
}
}
+31 -4
View File
@@ -2,23 +2,30 @@ package service
import ( import (
"context" "context"
"strings"
"time" "time"
statsv1 "github.com/bergabruh/stats-gateway/gen/stats/v1" 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/storage"
"github.com/bergabruh/stats-gateway/internal/tunnel"
"google.golang.org/grpc/codes" "google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/status" "google.golang.org/grpc/status"
) )
// Service implements statsv1.SpeedStatusServer with public RecentResults // Service implements statsv1.SpeedStatusServer.
// and Aggregate methods. Tunnel returns Unimplemented until Task 11 wires
// auth + handler.
type Service struct { type Service struct {
statsv1.UnimplementedSpeedStatusServer 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 { func (s *Service) RecentResults(f *statsv1.Filter, stream statsv1.SpeedStatus_RecentResultsServer) error {
limit := int(f.Limit) limit := int(f.Limit)
@@ -87,3 +94,23 @@ func (s *Service) Aggregate(ctx context.Context, r *statsv1.Range) (*statsv1.His
} }
return h, nil 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)
}
+38 -4
View File
@@ -8,19 +8,27 @@ import (
"time" "time"
"google.golang.org/grpc" "google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/credentials/insecure"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/status"
"google.golang.org/grpc/test/bufconn" "google.golang.org/grpc/test/bufconn"
statsv1 "github.com/bergabruh/stats-gateway/gen/stats/v1" 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/probes"
"github.com/bergabruh/stats-gateway/internal/storage" "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()) { func newTestServer(t *testing.T, st *storage.Store) (statsv1.SpeedStatusClient, func()) {
t.Helper() 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) lis := bufconn.Listen(1024 * 1024)
srv := grpc.NewServer() srv := grpc.NewServer()
statsv1.RegisterSpeedStatusServer(srv, New(st)) statsv1.RegisterSpeedStatusServer(srv, New(st, verifier, handler))
go srv.Serve(lis) go srv.Serve(lis)
conn, err := grpc.NewClient("passthrough:///bufnet", 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() dir := t.TempDir()
st, _ := storage.Open(filepath.Join(dir, "test.db")) st, _ := storage.Open(filepath.Join(dir, "test.db"))
defer st.Close() 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"}) stream.Send(&statsv1.Frame{Kind: statsv1.Frame_HELLO, Target: "1.2.3.4:80"})
_, err = stream.Recv() _, err = stream.Recv()
if err == nil { 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.
} }
+167
View File
@@ -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
}
+22
View File
@@ -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)
}
}
}