cache: Per-client fair-share bandwidth on table uplink
CI / vulncheck (pull_request) Successful in 20s
CI / check-and-test (pull_request) Failing after 21s

This commit is contained in:
2026-09-09 20:22:23 +00:00
parent 50ca0a071c
commit b7710de0ca
13 changed files with 439 additions and 18 deletions
+9
View File
@@ -236,6 +236,10 @@ While most configuration is done via the YAML file, some runtime options are sti
./steamcache2 --max-concurrent-requests 8 ./steamcache2 --max-concurrent-requests 8
./steamcache2 --max-requests-per-client 4 ./steamcache2 --max-requests-per-client 4
# Table-tier uplink shaping (empty/0 = use config / disabled)
./steamcache2 --uplink-bandwidth 10MB
./steamcache2 --max-bytes-per-client-per-sec 2500000
# Show help # Show help
./steamcache2 --help ./steamcache2 --help
``` ```
@@ -252,6 +256,11 @@ listen_address: :80
max_object_size: "0" # 0=unlimited; set e.g. "256MB" for response size DoS protection max_object_size: "0" # 0=unlimited; set e.g. "256MB" for response size DoS protection
trusted_proxies: [] # empty = safe (ignore XFF for rate limit); set CIDRs for trusted proxies trusted_proxies: [] # empty = safe (ignore XFF for rate limit); set CIDRs for trusted proxies
# Table-tier uplink bandwidth shaping (bytes/sec). Empty/0 = disabled (unlimited).
# Distinct from max_requests_per_client (concurrency). See "Table-tier uplink fair-share".
uplink_bandwidth: "" # e.g. "10MB" = 10e6 bytes/sec shared fairly across active clients
max_bytes_per_client_per_sec: 0 # optional absolute per-client cap; 0 = no absolute cap
# Cache configuration # Cache configuration
cache: cache:
# Memory cache settings # Memory cache settings
+12
View File
@@ -22,6 +22,8 @@ var (
maxConcurrentRequests int64 maxConcurrentRequests int64
maxRequestsPerClient int64 maxRequestsPerClient int64
uplinkBandwidth string
maxBytesPerClientPerSec int64
) )
var rootCmd = &cobra.Command{ var rootCmd = &cobra.Command{
@@ -107,6 +109,12 @@ var rootCmd = &cobra.Command{
if maxRequestsPerClient > 0 { if maxRequestsPerClient > 0 {
finalMaxRequestsPerClient = maxRequestsPerClient finalMaxRequestsPerClient = maxRequestsPerClient
} }
if uplinkBandwidth != "" {
cfg.UplinkBandwidth = uplinkBandwidth
}
if maxBytesPerClientPerSec > 0 {
cfg.MaxBytesPerClientPerSec = maxBytesPerClientPerSec
}
// Validate after loading and applying CLI overrides (fail fast, do not create default on validate error) // Validate after loading and applying CLI overrides (fail fast, do not create default on validate error)
if err := cfg.Validate(); err != nil { if err := cfg.Validate(); err != nil {
@@ -130,6 +138,8 @@ var rootCmd = &cobra.Command{
cfg.MaxObjectSize, cfg.MaxObjectSize,
cfg.TrustedProxies, cfg.TrustedProxies,
cfg.Cache.NegativeTTL, cfg.Cache.NegativeTTL,
cfg.UplinkBandwidth,
cfg.MaxBytesPerClientPerSec,
) )
if err != nil { if err != nil {
logger.Logger.Error(). logger.Logger.Error().
@@ -170,4 +180,6 @@ func init() {
rootCmd.Flags().Int64Var(&maxConcurrentRequests, "max-concurrent-requests", 0, "Maximum concurrent requests (0 = use config file value)") rootCmd.Flags().Int64Var(&maxConcurrentRequests, "max-concurrent-requests", 0, "Maximum concurrent requests (0 = use config file value)")
rootCmd.Flags().Int64Var(&maxRequestsPerClient, "max-requests-per-client", 0, "Maximum concurrent requests per client IP (0 = use config file value)") rootCmd.Flags().Int64Var(&maxRequestsPerClient, "max-requests-per-client", 0, "Maximum concurrent requests per client IP (0 = use config file value)")
rootCmd.Flags().StringVar(&uplinkBandwidth, "uplink-bandwidth", "", "Table uplink bandwidth bytes/sec human size e.g. 10MB (empty = use config; 0 disables)")
rootCmd.Flags().Int64Var(&maxBytesPerClientPerSec, "max-bytes-per-client-per-sec", 0, "Absolute per-client bytes/sec cap (0 = use config file value)")
} }
+13
View File
@@ -19,6 +19,11 @@ type Config struct {
MaxConcurrentRequests int64 `yaml:"max_concurrent_requests" default:"200"` MaxConcurrentRequests int64 `yaml:"max_concurrent_requests" default:"200"`
MaxRequestsPerClient int64 `yaml:"max_requests_per_client" default:"5"` MaxRequestsPerClient int64 `yaml:"max_requests_per_client" default:"5"`
// Table-tier uplink bandwidth shaping (bytes/sec). Distinct from MaxRequestsPerClient.
// Empty/"0" uplink and 0 max_bytes_per_client_per_sec = disabled (current unlimited behavior).
UplinkBandwidth string `yaml:"uplink_bandwidth"` // e.g. "10MB" via go-units = bytes/sec
MaxBytesPerClientPerSec int64 `yaml:"max_bytes_per_client_per_sec"` // absolute per-client cap; 0 = none
// Hardening limits (security/correctness) // Hardening limits (security/correctness)
MaxObjectSize string `yaml:"max_object_size" default:"0"` // 0=unlimited; e.g. "256MB" protects against OOM from huge/malicious upstream responses MaxObjectSize string `yaml:"max_object_size" default:"0"` // 0=unlimited; e.g. "256MB" protects against OOM from huge/malicious upstream responses
TrustedProxies []string `yaml:"trusted_proxies"` // CIDR list; empty=never trust X-Forwarded-For (safe default). See README security notes. TrustedProxies []string `yaml:"trusted_proxies"` // CIDR list; empty=never trust X-Forwarded-For (safe default). See README security notes.
@@ -186,6 +191,14 @@ func (c Config) Validate() error {
if c.MaxRequestsPerClient < 0 { if c.MaxRequestsPerClient < 0 {
return fmt.Errorf("negative per-client limit not allowed") return fmt.Errorf("negative per-client limit not allowed")
} }
if c.MaxBytesPerClientPerSec < 0 {
return fmt.Errorf("negative max_bytes_per_client_per_sec not allowed")
}
if c.UplinkBandwidth != "" && c.UplinkBandwidth != "0" {
if _, err := units.FromHumanSize(c.UplinkBandwidth); err != nil {
return fmt.Errorf("invalid uplink_bandwidth: %w", err)
}
}
if c.Cache.Memory.GCAlgorithm != "" { if c.Cache.Memory.GCAlgorithm != "" {
switch c.Cache.Memory.GCAlgorithm { switch c.Cache.Memory.GCAlgorithm {
+69
View File
@@ -211,3 +211,72 @@ func TestValidate(t *testing.T) {
}) })
} }
} }
func TestValidateUplinkBandwidth(t *testing.T) {
cases := []struct {
name string
mutate func(*Config)
wantErr bool
errSub string
}{
{
name: "empty uplink ok",
mutate: func(c *Config) {
c.UplinkBandwidth = ""
c.MaxBytesPerClientPerSec = 0
},
},
{
name: "zero uplink ok",
mutate: func(c *Config) {
c.UplinkBandwidth = "0"
},
},
{
name: "valid human size",
mutate: func(c *Config) {
c.UplinkBandwidth = "10MB"
},
},
{
name: "invalid uplink",
mutate: func(c *Config) {
c.UplinkBandwidth = "not-a-size"
},
wantErr: true,
errSub: "uplink_bandwidth",
},
{
name: "negative max bytes",
mutate: func(c *Config) {
c.MaxBytesPerClientPerSec = -1
},
wantErr: true,
errSub: "max_bytes_per_client_per_sec",
},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
c := GetDefaultConfig()
tt.mutate(&c)
err := c.Validate()
if tt.wantErr {
if err == nil {
t.Fatalf("Validate() error = nil, wantErr")
}
if tt.errSub != "" && !contains(err.Error(), tt.errSub) {
t.Fatalf("Validate() error %q does not contain %q", err.Error(), tt.errSub)
}
return
}
if err != nil {
t.Fatalf("Validate() unexpected error: %v", err)
}
})
}
}
func contains(s, sub string) bool {
return strings.Contains(s, sub)
}
+1
View File
@@ -38,6 +38,7 @@
listen_address: :80 listen_address: :80
max_concurrent_requests: 1000 max_concurrent_requests: 1000
# uplink_bandwidth / max_bytes_per_client_per_sec default off (unlimited)
max_requests_per_client: 10 max_requests_per_client: 10
max_object_size: "0" # unlimited for validation (real Steam files can be large) max_object_size: "0" # unlimited for validation (real Steam files can be large)
+1
View File
@@ -9,6 +9,7 @@ require (
github.com/spf13/cobra v1.8.1 github.com/spf13/cobra v1.8.1
golang.org/x/sync v0.16.0 golang.org/x/sync v0.16.0
golang.org/x/sys v0.12.0 golang.org/x/sys v0.12.0
golang.org/x/time v0.16.0
gopkg.in/yaml.v3 v3.0.1 gopkg.in/yaml.v3 v3.0.1
) )
+2
View File
@@ -27,6 +27,8 @@ golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBc
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.12.0 h1:CM0HF96J0hcLAwsHPJZjfdNzs0gftsLfgKt57wWHJ0o= golang.org/x/sys v0.12.0 h1:CM0HF96J0hcLAwsHPJZjfdNzs0gftsLfgKt57wWHJ0o=
golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.12.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/time v0.16.0 h1:vMb6ptszcQMkcwiRTAuNNU50gom6++Q/6gY2hDM6VDE=
golang.org/x/time v0.16.0/go.mod h1:rVKOqvZeKvrDKTQiAHJ7wmwP0RzleSphoEA9RcdLA0s=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= 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/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 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
+158
View File
@@ -0,0 +1,158 @@
// steamcache/bandwidth.go
// Per-client fair-share / absolute bandwidth shaping for table-tier uplink.
// Distinct from max_requests_per_client concurrency (semaphores in ratelimit.go).
package steamcache
import (
"context"
"net/http"
"sync"
"golang.org/x/time/rate"
)
const bandwidthWriteChunk = 32 * 1024
// clientBandwidthLimiter fair-shares uplinkBytesPerSec among active clients and/or
// applies an absolute per-client bytes/sec cap. Both 0 disables shaping.
type clientBandwidthLimiter struct {
uplinkBytesPerSec int64
absoluteCap int64
mu sync.Mutex
active map[string]int // refcount of in-flight shaped responses per client IP
limiters map[string]*rate.Limiter
}
func newClientBandwidthLimiter(uplinkBytesPerSec, absoluteCap int64) *clientBandwidthLimiter {
if uplinkBytesPerSec < 0 {
uplinkBytesPerSec = 0
}
if absoluteCap < 0 {
absoluteCap = 0
}
return &clientBandwidthLimiter{
uplinkBytesPerSec: uplinkBytesPerSec,
absoluteCap: absoluteCap,
active: make(map[string]int),
limiters: make(map[string]*rate.Limiter),
}
}
func (b *clientBandwidthLimiter) enabled() bool {
return b != nil && (b.uplinkBytesPerSec > 0 || b.absoluteCap > 0)
}
// acquire registers clientIP as actively downloading and returns its limiter
// (nil if shaping disabled) plus a release func that must be deferred.
func (b *clientBandwidthLimiter) acquire(clientIP string) (*rate.Limiter, func()) {
if !b.enabled() {
return nil, func() {}
}
b.mu.Lock()
b.active[clientIP]++
lim := b.ensureLimiterLocked(clientIP)
b.recomputeRatesLocked()
b.mu.Unlock()
var once sync.Once
release := func() {
once.Do(func() {
b.mu.Lock()
defer b.mu.Unlock()
if n := b.active[clientIP]; n <= 1 {
delete(b.active, clientIP)
} else {
b.active[clientIP] = n - 1
}
b.recomputeRatesLocked()
})
}
return lim, release
}
func (b *clientBandwidthLimiter) ensureLimiterLocked(clientIP string) *rate.Limiter {
if lim, ok := b.limiters[clientIP]; ok {
return lim
}
// Start with a placeholder; recomputeRatesLocked sets the real rate.
lim := rate.NewLimiter(rate.Limit(1), 1)
b.limiters[clientIP] = lim
return lim
}
func (b *clientBandwidthLimiter) recomputeRatesLocked() {
n := len(b.active)
if n == 0 {
return
}
var fair int64
if b.uplinkBytesPerSec > 0 {
fair = b.uplinkBytesPerSec / int64(n)
if fair < 1 {
fair = 1
}
}
for ip := range b.active {
r := fair
if b.absoluteCap > 0 {
if r == 0 || b.absoluteCap < r {
r = b.absoluteCap
}
}
if r < 1 {
r = 1
}
lim := b.ensureLimiterLocked(ip)
burst := int(r)
if burst < bandwidthWriteChunk {
burst = bandwidthWriteChunk
}
// Cap burst to avoid huge memory spikes on huge uplinks.
if burst > 4*bandwidthWriteChunk {
burst = 4 * bandwidthWriteChunk
}
lim.SetLimit(rate.Limit(r))
lim.SetBurst(burst)
}
}
// limitedResponseWriter rate-limits response body Write calls. Headers/WriteHeader
// are unlimited. Implements http.ResponseWriter (+ optional Flusher/Hijacker passthrough
// is intentionally omitted — SteamCache body path only needs Write).
type limitedResponseWriter struct {
http.ResponseWriter
lim *rate.Limiter
ctx context.Context
}
func (w *limitedResponseWriter) Write(p []byte) (int, error) {
if w.lim == nil || len(p) == 0 {
return w.ResponseWriter.Write(p)
}
ctx := w.ctx
if ctx == nil {
ctx = context.Background()
}
total := 0
for total < len(p) {
chunk := p[total:]
if len(chunk) > bandwidthWriteChunk {
chunk = chunk[:bandwidthWriteChunk]
}
if err := w.lim.WaitN(ctx, len(chunk)); err != nil {
return total, err
}
n, err := w.ResponseWriter.Write(chunk)
total += n
if err != nil {
return total, err
}
}
return total, nil
}
// Unwrap exposes the underlying ResponseWriter for http.ResponseController etc.
func (w *limitedResponseWriter) Unwrap() http.ResponseWriter {
return w.ResponseWriter
}
+128
View File
@@ -0,0 +1,128 @@
package steamcache
import (
"bytes"
"context"
"net/http"
"net/http/httptest"
"testing"
"time"
)
func TestBandwidthFairShareRates(t *testing.T) {
b := newClientBandwidthLimiter(1000, 0)
lim1, rel1 := b.acquire("1.1.1.1")
defer rel1()
lim2, rel2 := b.acquire("2.2.2.2")
defer rel2()
if lim1 == nil || lim2 == nil {
t.Fatal("expected limiters")
}
// With 2 active clients, each should get ~500 bytes/sec.
got1 := float64(lim1.Limit())
got2 := float64(lim2.Limit())
if got1 < 400 || got1 > 600 || got2 < 400 || got2 > 600 {
t.Fatalf("fair-share rates = %v,%v want ~500", got1, got2)
}
rel2()
// After release, sole client should get full uplink.
lim1b, rel1b := b.acquire("1.1.1.1")
defer rel1b()
if float64(lim1b.Limit()) < 900 {
t.Fatalf("after release limit=%v want ~1000", lim1b.Limit())
}
}
func TestBandwidthAbsoluteCap(t *testing.T) {
b := newClientBandwidthLimiter(0, 250)
lim, rel := b.acquire("9.9.9.9")
defer rel()
if lim == nil {
t.Fatal("expected limiter")
}
if float64(lim.Limit()) != 250 {
t.Fatalf("limit=%v want 250", lim.Limit())
}
}
func TestBandwidthDisabled(t *testing.T) {
b := newClientBandwidthLimiter(0, 0)
lim, rel := b.acquire("9.9.9.9")
defer rel()
if lim != nil {
t.Fatal("expected nil limiter when disabled")
}
}
func TestLimitedResponseWriterShapes(t *testing.T) {
var buf bytes.Buffer
rec := httptest.NewRecorder()
// Use a custom writer sink via ResponseRecorder is fine; WaitN will delay.
lim := newClientBandwidthLimiter(0, 2000) // 2KB/s
l, rel := lim.acquire("127.0.0.1")
defer rel()
w := &limitedResponseWriter{ResponseWriter: rec, lim: l, ctx: context.Background()}
payload := bytes.Repeat([]byte("x"), 4000)
start := time.Now()
n, err := w.Write(payload)
elapsed := time.Since(start)
if err != nil {
t.Fatal(err)
}
if n != len(payload) {
t.Fatalf("wrote %d want %d", n, len(payload))
}
_ = buf
// 4000 bytes at 2000 B/s should take ~2s (allow slack for CI).
if elapsed < 1500*time.Millisecond {
t.Fatalf("elapsed %v too fast for 2KB/s shaping of 4KB", elapsed)
}
if elapsed > 8*time.Second {
t.Fatalf("elapsed %v unexpectedly slow", elapsed)
}
}
func TestServeHTTPBandwidthCap(t *testing.T) {
body := bytes.Repeat([]byte("a"), 3000)
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/octet-stream")
w.Header().Set("Content-Length", "3000")
w.WriteHeader(200)
_, _ = w.Write(body)
}))
t.Cleanup(upstream.Close)
sc, err := NewWithOptions(Options{
Address: "127.0.0.1:0",
MemorySize: "1MB",
DiskSize: "0",
Upstream: upstream.URL,
MemoryGC: "lru",
DiskGC: "lru",
MaxConcurrentRequests: 20,
MaxRequestsPerClient: 10,
MaxObjectSize: "0",
MaxBytesPerClientPerSec: 1500, // 1.5KB/s
})
if err != nil {
t.Fatalf("NewWithOptions: %v", err)
}
t.Cleanup(func() { sc.Shutdown() })
req := httptest.NewRequest(http.MethodGet, "/depot/bw/chunk", nil)
req.Header.Set("User-Agent", "Valve/Steam HTTP Client 1.0")
rr := httptest.NewRecorder()
start := time.Now()
sc.ServeHTTP(rr, req)
elapsed := time.Since(start)
if rr.Code != 200 {
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
}
got := rr.Body.Bytes()
if len(got) != len(body) {
t.Fatalf("body len=%d want %d", len(got), len(body))
}
if elapsed < time.Second {
t.Fatalf("elapsed %v too fast for shaping", elapsed)
}
}
+14 -1
View File
@@ -38,10 +38,14 @@ type Options struct {
// NegativeTTL is a Go duration string for 404/410 negative cache entries. // NegativeTTL is a Go duration string for 404/410 negative cache entries.
// Empty defaults to 5m. "0" / "0s" disables storing negatives. // Empty defaults to 5m. "0" / "0s" disables storing negatives.
NegativeTTL string NegativeTTL string
// Table-tier uplink bandwidth shaping (bytes/sec). Empty/0 = disabled.
UplinkBandwidth string
MaxBytesPerClientPerSec int64
} }
func NewWithOptions(o Options) (*SteamCache, error) { func NewWithOptions(o Options) (*SteamCache, error) {
return New(o.Address, o.MemorySize, o.DiskSize, o.DiskPath, o.Upstream, o.MemoryGC, o.DiskGC, o.MaxConcurrentRequests, o.MaxRequestsPerClient, o.MaxObjectSize, o.TrustedProxies, o.NegativeTTL) return New(o.Address, o.MemorySize, o.DiskSize, o.DiskPath, o.Upstream, o.MemoryGC, o.DiskGC, o.MaxConcurrentRequests, o.MaxRequestsPerClient, o.MaxObjectSize, o.TrustedProxies, o.NegativeTTL, o.UplinkBandwidth, o.MaxBytesPerClientPerSec)
} }
// handleSpecialEndpoints handles non-content paths (health, heartbeat, metrics) and // handleSpecialEndpoints handles non-content paths (health, heartbeat, metrics) and
@@ -413,6 +417,15 @@ func (sc *SteamCache) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return return
} }
// Per-client uplink bandwidth shaping (table-tier). Distinct from concurrency limits above.
if sc.bandwidth != nil && sc.bandwidth.enabled() {
lim, release := sc.bandwidth.acquire(clientIP)
defer release()
if lim != nil {
w = &limitedResponseWriter{ResponseWriter: w, lim: lim, ctx: r.Context()}
}
}
// Check if this is a request from a supported service // Check if this is a request from a supported service
if service, isSupported := sc.detectService(r); isSupported { if service, isSupported := sc.detectService(r); isSupported {
// Cache key is the path only, never the Host: Steam rotates CDN hostnames // Cache key is the path only, never the Host: Steam rotates CDN hostnames
+1 -1
View File
@@ -170,7 +170,7 @@ func TestSerializeNegativeHeader(t *testing.T) {
} }
func TestNewInvalidNegativeTTL(t *testing.T) { func TestNewInvalidNegativeTTL(t *testing.T) {
sc, err := New("127.0.0.1:0", "1MB", "0", t.TempDir(), "", "lru", "lru", 10, 5, "0", nil, "not-a-duration") sc, err := New("127.0.0.1:0", "1MB", "0", t.TempDir(), "", "lru", "lru", 10, 5, "0", nil, "not-a-duration", "", 0)
if err == nil { if err == nil {
if sc != nil { if sc != nil {
sc.Shutdown() sc.Shutdown()
+16 -1
View File
@@ -55,6 +55,9 @@ type SteamCache struct {
clientRateLimiter *clientRateLimiter clientRateLimiter *clientRateLimiter
maxRequestsPerClient int64 maxRequestsPerClient int64
// Per-client uplink bandwidth shaping (see bandwidth.go); nil/disabled = unlimited
bandwidth *clientBandwidthLimiter
// Hardening config fields (plumbed) // Hardening config fields (plumbed)
maxObjectSize int64 maxObjectSize int64
trustedProxies []string trustedProxies []string
@@ -85,7 +88,7 @@ const DefaultNegativeTTL = 5 * time.Minute
// negativeTTL is a Go duration string for 404/410 negative cache entries; empty means 5m. // negativeTTL is a Go duration string for 404/410 negative cache entries; empty means 5m.
// Callers must check the returned error. // Callers must check the returned error.
// Prefer NewWithOptions (or config file) for forward compatibility. See README migration notes. // Prefer NewWithOptions (or config file) for forward compatibility. See README migration notes.
func New(address string, memorySize string, diskSize string, diskPath, upstream, memoryGC, diskGC string, maxConcurrentRequests int64, maxRequestsPerClient int64, maxObjectSize string, trustedProxies []string, negativeTTL string) (*SteamCache, error) { func New(address string, memorySize string, diskSize string, diskPath, upstream, memoryGC, diskGC string, maxConcurrentRequests int64, maxRequestsPerClient int64, maxObjectSize string, trustedProxies []string, negativeTTL string, uplinkBandwidth string, maxBytesPerClientPerSec int64) (*SteamCache, error) {
memorysize, err := units.FromHumanSize(memorySize) memorysize, err := units.FromHumanSize(memorySize)
if err != nil { if err != nil {
return nil, fmt.Errorf("invalid memory size: %w", err) return nil, fmt.Errorf("invalid memory size: %w", err)
@@ -114,6 +117,17 @@ func New(address string, memorySize string, diskSize string, diskPath, upstream,
return nil, err return nil, err
} }
var uplinkBytes int64
if uplinkBandwidth != "" && uplinkBandwidth != "0" {
uplinkBytes, err = units.FromHumanSize(uplinkBandwidth)
if err != nil {
return nil, fmt.Errorf("invalid uplink bandwidth: %w", err)
}
}
if maxBytesPerClientPerSec < 0 {
return nil, fmt.Errorf("negative max_bytes_per_client_per_sec not allowed")
}
c := cache.New() c := cache.New()
var m *memory.MemoryFS var m *memory.MemoryFS
@@ -178,6 +192,7 @@ func New(address string, memorySize string, diskSize string, diskPath, upstream,
requestSemaphore: semaphore.NewWeighted(maxConcurrentRequests), requestSemaphore: semaphore.NewWeighted(maxConcurrentRequests),
clientRateLimiter: newClientRateLimiter(maxRequestsPerClient), clientRateLimiter: newClientRateLimiter(maxRequestsPerClient),
maxRequestsPerClient: maxRequestsPerClient, maxRequestsPerClient: maxRequestsPerClient,
bandwidth: newClientBandwidthLimiter(uplinkBytes, maxBytesPerClientPerSec),
shutdownCh: make(chan struct{}), shutdownCh: make(chan struct{}),
// Hardening config plumbed // Hardening config plumbed
+15 -15
View File
@@ -28,7 +28,7 @@ import (
func TestCaching(t *testing.T) { func TestCaching(t *testing.T) {
td := t.TempDir() td := t.TempDir()
sc, err := New("localhost:8080", "1G", "1G", td, "", "lru", "lru", 200, 5, "0", nil, "") sc, err := New("localhost:8080", "1G", "1G", td, "", "lru", "lru", 200, 5, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("failed to create SteamCache: %v", err) t.Fatalf("failed to create SteamCache: %v", err)
} }
@@ -133,7 +133,7 @@ func TestCaching(t *testing.T) {
} }
func TestCacheMissAndHit(t *testing.T) { func TestCacheMissAndHit(t *testing.T) {
sc, err := New("localhost:8080", "1MB", "1G", t.TempDir(), "", "lru", "lru", 200, 5, "0", nil, "") sc, err := New("localhost:8080", "1MB", "1G", t.TempDir(), "", "lru", "lru", 200, 5, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("failed to create SteamCache: %v", err) t.Fatalf("failed to create SteamCache: %v", err)
} }
@@ -376,7 +376,7 @@ func TestServiceManagerExpandability(t *testing.T) {
// Removed hash calculation tests since we switched to lightweight validation // Removed hash calculation tests since we switched to lightweight validation
func TestSteamKeySharding(t *testing.T) { func TestSteamKeySharding(t *testing.T) {
sc, err := New("localhost:8080", "1MB", "1G", t.TempDir(), "", "lru", "lru", 200, 5, "0", nil, "") sc, err := New("localhost:8080", "1MB", "1G", t.TempDir(), "", "lru", "lru", 200, 5, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("failed to create SteamCache: %v", err) t.Fatalf("failed to create SteamCache: %v", err)
} }
@@ -483,7 +483,7 @@ func TestErrorTypes(t *testing.T) {
// TestMetrics tests the metrics functionality // TestMetrics tests the metrics functionality
func TestMetrics(t *testing.T) { func TestMetrics(t *testing.T) {
td := t.TempDir() td := t.TempDir()
sc, err := New("localhost:8080", "1G", "1G", td, "", "lru", "lru", 200, 5, "0", nil, "") sc, err := New("localhost:8080", "1G", "1G", td, "", "lru", "lru", 200, 5, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("failed to create SteamCache: %v", err) t.Fatalf("failed to create SteamCache: %v", err)
} }
@@ -587,7 +587,7 @@ func newTestCacheWithFakeUpstream(t *testing.T, h http.HandlerFunc, mem, disk st
s := httptest.NewServer(h) s := httptest.NewServer(h)
t.Cleanup(s.Close) t.Cleanup(s.Close)
d := t.TempDir() d := t.TempDir()
sc, err := New("127.0.0.1:0", mem, disk, d, s.URL, "lru", "lru", 200, 10, "0", nil, "") sc, err := New("127.0.0.1:0", mem, disk, d, s.URL, "lru", "lru", 200, 10, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("failed to create SteamCache: %v", err) t.Fatalf("failed to create SteamCache: %v", err)
} }
@@ -749,7 +749,7 @@ func TestErrorMetrics(t *testing.T) {
// Cover 503 capacity path + accounting skew: force Acquire err via canceled ctx. // Cover 503 capacity path + accounting skew: force Acquire err via canceled ctx.
// Asserts Errors+RateLimited inc, Total unchanged (per documented design in code comment). // Asserts Errors+RateLimited inc, Total unchanged (per documented design in code comment).
tdCap := t.TempDir() tdCap := t.TempDir()
scCap, err := New("127.0.0.1:0", "1MB", "0", tdCap, "", "lru", "lru", 200, 5, "0", nil, "") scCap, err := New("127.0.0.1:0", "1MB", "0", tdCap, "", "lru", "lru", 200, 5, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("cap sc: %v", err) t.Fatalf("cap sc: %v", err)
} }
@@ -813,7 +813,7 @@ func TestErrorMetrics(t *testing.T) {
func TestExpandedErrorMetrics(t *testing.T) { func TestExpandedErrorMetrics(t *testing.T) {
t.Parallel() t.Parallel()
td := t.TempDir() td := t.TempDir()
sc, err := New("localhost:0", "1MB", "0", td, "", "lru", "lru", 10, 5, "0", nil, "") sc, err := New("localhost:0", "1MB", "0", td, "", "lru", "lru", 10, 5, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("create: %v", err) t.Fatalf("create: %v", err)
} }
@@ -903,7 +903,7 @@ func TestNewInvalidSizes(t *testing.T) {
} }
for _, c := range cases { for _, c := range cases {
t.Run(c.mem+"_"+c.disk, func(t *testing.T) { t.Run(c.mem+"_"+c.disk, func(t *testing.T) {
sc, err := New("127.0.0.1:0", c.mem, c.disk, t.TempDir(), "", "lru", "lru", 10, 5, c.maxobj, nil, "") sc, err := New("127.0.0.1:0", c.mem, c.disk, t.TempDir(), "", "lru", "lru", 10, 5, c.maxobj, nil, "", "", 0)
if err == nil { if err == nil {
t.Fatal("expected error for bad size, got nil") t.Fatal("expected error for bad size, got nil")
} }
@@ -924,7 +924,7 @@ func TestNewRunShutdownHygiene(t *testing.T) {
t.Skip("skips Run hygiene in -short per existing pattern") t.Skip("skips Run hygiene in -short per existing pattern")
} }
d := t.TempDir() d := t.TempDir()
sc, err := New("127.0.0.1:0", "1MB", "0", d, "", "lru", "lru", 10, 5, "0", nil, "") sc, err := New("127.0.0.1:0", "1MB", "0", d, "", "lru", "lru", 10, 5, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("new: %v", err) t.Fatalf("new: %v", err)
} }
@@ -1073,7 +1073,7 @@ func TestDiskOnlyDelayedAttach(t *testing.T) {
}) })
// mem=0, disk>0 -> pure disk delayed path (go func) // mem=0, disk>0 -> pure disk delayed path (go func)
sc, err := New("localhost:0", "0", "10MB", diskPath, "", "lru", "lru", 10, 1, "0", nil, "") sc, err := New("localhost:0", "0", "10MB", diskPath, "", "lru", "lru", 10, 1, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("New disk-only: %v", err) t.Fatalf("New disk-only: %v", err)
} }
@@ -1145,7 +1145,7 @@ func TestDiskOnlyDelayedAttach(t *testing.T) {
// TestDiskTierSignalMemoryOnly covers memory-only mode: DiskTierReady=1 (N/A, not // TestDiskTierSignalMemoryOnly covers memory-only mode: DiskTierReady=1 (N/A, not
// waiting on disk attach) and heartbeat header X-SteamCache-Disk-Tier: disabled. // waiting on disk attach) and heartbeat header X-SteamCache-Disk-Tier: disabled.
func TestDiskTierSignalMemoryOnly(t *testing.T) { func TestDiskTierSignalMemoryOnly(t *testing.T) {
sc, err := New("127.0.0.1:0", "1MB", "0", t.TempDir(), "", "lru", "lru", 10, 5, "0", nil, "") sc, err := New("127.0.0.1:0", "1MB", "0", t.TempDir(), "", "lru", "lru", 10, 5, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("New memory-only: %v", err) t.Fatalf("New memory-only: %v", err)
} }
@@ -1189,7 +1189,7 @@ func TestDiskTierSignalMixedPendingReady(t *testing.T) {
disk.ClearInitHold(diskPath) disk.ClearInitHold(diskPath)
}) })
sc, err := New("127.0.0.1:0", "1MB", "10MB", diskPath, "", "lru", "lru", 10, 1, "0", nil, "") sc, err := New("127.0.0.1:0", "1MB", "10MB", diskPath, "", "lru", "lru", 10, 1, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("New mixed: %v", err) t.Fatalf("New mixed: %v", err)
} }
@@ -1341,7 +1341,7 @@ func TestHostAllowedForDirectFetch(t *testing.T) {
func TestDirectFetchRejectsNonSteamHost(t *testing.T) { func TestDirectFetchRejectsNonSteamHost(t *testing.T) {
td := t.TempDir() td := t.TempDir()
sc, err := New("127.0.0.1:0", "1MB", "0", td, "", "lru", "lru", 200, 5, "0", nil, "") sc, err := New("127.0.0.1:0", "1MB", "0", td, "", "lru", "lru", 200, 5, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("New: %v", err) t.Fatalf("New: %v", err)
} }
@@ -1505,7 +1505,7 @@ func TestCacheKeySharedAcrossCDNHostAliases(t *testing.T) {
// still report the configured disk capacity. // still report the configured disk capacity.
func TestGetMetricsCapacityGauges(t *testing.T) { func TestGetMetricsCapacityGauges(t *testing.T) {
t.Run("memory-only", func(t *testing.T) { t.Run("memory-only", func(t *testing.T) {
sc, err := New("127.0.0.1:0", "1MB", "0", t.TempDir(), "", "lru", "lru", 10, 5, "0", nil, "") sc, err := New("127.0.0.1:0", "1MB", "0", t.TempDir(), "", "lru", "lru", 10, 5, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("New memory-only: %v", err) t.Fatalf("New memory-only: %v", err)
} }
@@ -1551,7 +1551,7 @@ func TestGetMetricsCapacityGauges(t *testing.T) {
disk.ClearInitHold(diskPath) disk.ClearInitHold(diskPath)
}) })
sc, err := New("127.0.0.1:0", "1MB", "10MB", diskPath, "", "lru", "lru", 10, 1, "0", nil, "") sc, err := New("127.0.0.1:0", "1MB", "10MB", diskPath, "", "lru", "lru", 10, 1, "0", nil, "", "", 0)
if err != nil { if err != nil {
t.Fatalf("New mixed: %v", err) t.Fatalf("New mixed: %v", err)
} }