diff --git a/internal/bootstrap/config.go b/internal/bootstrap/config.go index 8304468080..f634b053fa 100644 --- a/internal/bootstrap/config.go +++ b/internal/bootstrap/config.go @@ -9,11 +9,12 @@ import ( "github.com/OpenListTeam/OpenList/v4/cmd/flags" "github.com/OpenListTeam/OpenList/v4/drivers/base" + "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/mem" "github.com/OpenListTeam/OpenList/v4/internal/net" "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/caarlos0/env/v9" - "github.com/shirou/gopsutil/v4/mem" log "github.com/sirupsen/logrus" ) @@ -96,6 +97,13 @@ func InitConfig() { if !conf.Conf.Force { confFromEnv() } + cachePolicy, policyErr := cache.ResolvePolicy(conf.Conf.CachePolicy, cache.PolicyAuto) + if policyErr != nil { + log.Fatalf("resolve cache policy error: %+v", policyErr) + } + conf.Conf.CachePolicy = cachePolicy + conf.CachePolicy = cachePolicy + log.Infof("cache policy: %s", conf.CachePolicy) if conf.Conf.MaxConcurrency > math.MaxInt32 { net.DefaultConcurrencyLimit = &net.ConcurrencyLimit{Limit: math.MaxInt32} @@ -103,15 +111,20 @@ func InitConfig() { net.DefaultConcurrencyLimit = &net.ConcurrencyLimit{Limit: uint32(conf.Conf.MaxConcurrency)} } - memStat, _ := mem.VirtualMemory() - if memStat != nil { - log.Infof("total memory: %dMB, available: %dMB", memStat.Total>>20, memStat.Available>>20) + memStat, memErr := mem.GetMemorySnapshot() + if memErr != nil { + log.Warnf("memory detection warning: %v", memErr) + } + memoryCacheEnabled := false + if memStat.Limit > 0 { + log.Infof("effective memory: limit=%dMB, used=%dMB, available=%dMB, source=%s", memStat.Limit>>20, memStat.Used>>20, memStat.Available>>20, memStat.Source) if conf.Conf.MinFreeMemory < 0 { conf.MinFreeMemory = 0 log.Info("disable memory cache") } else { + memoryCacheEnabled = true if conf.Conf.MinFreeMemory < 16 { - t := (memStat.Total >> 20) / 10 + t := (memStat.Limit >> 20) / 10 conf.MinFreeMemory = max(16, min(t, 1024)) << 20 } else { conf.MinFreeMemory = uint64(conf.Conf.MinFreeMemory) << 20 @@ -120,7 +133,7 @@ func InitConfig() { } if conf.Conf.MaxBlockLimit < 4 { - t := (memStat.Total >> 20) * 3 / 100 + t := (memStat.Limit >> 20) * 3 / 100 conf.MaxBlockLimit = max(4, min(uint64(t), 64)) << 20 } else { conf.MaxBlockLimit = uint64(conf.Conf.MaxBlockLimit) << 20 @@ -130,6 +143,9 @@ func InitConfig() { conf.MinFreeMemory = 0 log.Warn("failed to get memory info, disable memory cache") } + budgetCapacity := cacheMemoryBudgetCapacity(memStat.Available, conf.MinFreeMemory, memoryCacheEnabled) + mem.CacheMemoryBudget.SetCapacity(budgetCapacity) + log.Infof("cache memory budget: %dMB", budgetCapacity>>20) if conf.Conf.AutoMemoryLimit > 0 { conf.AutoMemoryLimit = uint64(conf.Conf.AutoMemoryLimit) << 20 @@ -169,6 +185,13 @@ func InitConfig() { initURL() } +func cacheMemoryBudgetCapacity(available, minFree uint64, enabled bool) uint64 { + if !enabled || available <= minFree { + return 0 + } + return available - minFree +} + func confFromEnv() { prefix := "OPENLIST_" if flags.NoPrefix { diff --git a/internal/bootstrap/config_test.go b/internal/bootstrap/config_test.go new file mode 100644 index 0000000000..84b1ae01d8 --- /dev/null +++ b/internal/bootstrap/config_test.go @@ -0,0 +1,25 @@ +package bootstrap + +import "testing" + +func TestCacheMemoryBudgetCapacity(t *testing.T) { + tests := []struct { + name string + available uint64 + minFree uint64 + enabled bool + want uint64 + }{ + {name: "enabled", available: 10, minFree: 3, enabled: true, want: 7}, + {name: "disabled", available: 10, minFree: 3, want: 0}, + {name: "reserve equals available", available: 10, minFree: 10, enabled: true, want: 0}, + {name: "reserve exceeds available", available: 10, minFree: 11, enabled: true, want: 0}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := cacheMemoryBudgetCapacity(tt.available, tt.minFree, tt.enabled); got != tt.want { + t.Fatalf("cacheMemoryBudgetCapacity() = %d, want %d", got, tt.want) + } + }) + } +} diff --git a/internal/cache/policy.go b/internal/cache/policy.go new file mode 100644 index 0000000000..4d0dc22c54 --- /dev/null +++ b/internal/cache/policy.go @@ -0,0 +1,70 @@ +package cache + +import ( + "fmt" + "strings" +) + +type Policy string + +const ( + // PolicyInherit is only used by per-instance overrides. It is not a + // user-selectable cache policy. + PolicyInherit Policy = "" + PolicyAuto Policy = "auto" + PolicyMemory Policy = "memory" + PolicyDisk Policy = "disk" +) + +// ParsePolicy parses a user-facing global policy, where an empty value means +// auto. Text unmarshalling keeps an empty per-instance value as inherit. +func ParsePolicy(value string) (Policy, error) { + policy := Policy(strings.ToLower(strings.TrimSpace(value))) + if policy == PolicyInherit { + return PolicyAuto, nil + } + if !policy.IsConcrete() { + return PolicyInherit, fmt.Errorf("invalid cache policy %q: expected auto, memory, or disk", value) + } + return policy, nil +} + +func ResolvePolicy(override, fallback Policy) (Policy, error) { + policy := override + if policy == PolicyInherit { + policy = fallback + } + if !policy.IsConcrete() { + return PolicyInherit, fmt.Errorf("invalid cache policy %q: expected auto, memory, or disk", policy) + } + return policy, nil +} + +func (p Policy) IsConcrete() bool { + return p == PolicyAuto || p == PolicyMemory || p == PolicyDisk +} + +func (p Policy) MarshalText() ([]byte, error) { + if p == PolicyInherit { + return []byte{}, nil + } + if !p.IsConcrete() { + return nil, fmt.Errorf("invalid cache policy %q", p) + } + return []byte(p), nil +} + +func (p *Policy) UnmarshalText(text []byte) error { + if strings.TrimSpace(string(text)) == "" { + // Preserve inherit for per-instance configs. Global configuration + // resolves an empty policy to auto during bootstrap. + *p = PolicyInherit + return nil + } + policy, err := ParsePolicy(string(text)) + if err != nil { + return err + } + *p = policy + return nil +} diff --git a/internal/cache/policy_test.go b/internal/cache/policy_test.go new file mode 100644 index 0000000000..5996e27e52 --- /dev/null +++ b/internal/cache/policy_test.go @@ -0,0 +1,95 @@ +package cache + +import ( + "encoding/json" + "testing" + + "github.com/caarlos0/env/v9" +) + +func TestParsePolicy(t *testing.T) { + tests := []struct { + input string + want Policy + }{ + {"", PolicyAuto}, + {"auto", PolicyAuto}, + {" MEMORY ", PolicyMemory}, + {"Disk", PolicyDisk}, + } + for _, tt := range tests { + got, err := ParsePolicy(tt.input) + if err != nil { + t.Fatalf("ParsePolicy(%q) error = %v", tt.input, err) + } + if got != tt.want { + t.Errorf("ParsePolicy(%q) = %q, want %q", tt.input, got, tt.want) + } + } + if _, err := ParsePolicy("hybrid"); err == nil { + t.Fatal("ParsePolicy() expected an error for an invalid policy") + } +} + +func TestPolicyEnvironment(t *testing.T) { + t.Setenv("OPENLIST_TEST_CACHE_POLICY", " Disk ") + var cfg struct { + Policy Policy `env:"CACHE_POLICY"` + } + if err := env.ParseWithOptions(&cfg, env.Options{Prefix: "OPENLIST_TEST_"}); err != nil { + t.Fatalf("env.ParseWithOptions() error = %v", err) + } + if cfg.Policy != PolicyDisk { + t.Fatalf("environment policy = %q, want disk", cfg.Policy) + } +} + +func TestResolvePolicy(t *testing.T) { + got, err := ResolvePolicy(PolicyInherit, PolicyDisk) + if err != nil || got != PolicyDisk { + t.Fatalf("ResolvePolicy(inherit, disk) = %q, %v", got, err) + } + got, err = ResolvePolicy(PolicyMemory, PolicyDisk) + if err != nil || got != PolicyMemory { + t.Fatalf("ResolvePolicy(memory, disk) = %q, %v", got, err) + } + if _, err := ResolvePolicy(PolicyInherit, Policy("invalid")); err == nil { + t.Fatal("ResolvePolicy() expected an error for an invalid fallback") + } +} + +func TestPolicyJSON(t *testing.T) { + type config struct { + Policy Policy `json:"cache_policy"` + } + var cfg config + if err := json.Unmarshal([]byte(`{"cache_policy":" MEMORY "}`), &cfg); err != nil { + t.Fatalf("json.Unmarshal() error = %v", err) + } + if cfg.Policy != PolicyMemory { + t.Fatalf("json.Unmarshal() policy = %q, want memory", cfg.Policy) + } + b, err := json.Marshal(cfg) + if err != nil { + t.Fatalf("json.Marshal() error = %v", err) + } + if string(b) != `{"cache_policy":"memory"}` { + t.Fatalf("json.Marshal() = %s", b) + } + if err := json.Unmarshal([]byte(`{"cache_policy":"invalid"}`), &cfg); err == nil { + t.Fatal("json.Unmarshal() expected an error for an invalid policy") + } + + cfg.Policy = PolicyInherit + b, err = json.Marshal(cfg) + if err != nil { + t.Fatalf("json.Marshal(inherit) error = %v", err) + } + var roundTrip config + if err := json.Unmarshal(b, &roundTrip); err != nil { + t.Fatalf("json.Unmarshal(inherit) error = %v", err) + } + if roundTrip.Policy != PolicyInherit { + t.Fatalf("inherit round trip policy = %q, want inherit", roundTrip.Policy) + } +} diff --git a/internal/conf/config.go b/internal/conf/config.go index f8423d5908..5e72ce46d8 100644 --- a/internal/conf/config.go +++ b/internal/conf/config.go @@ -3,6 +3,7 @@ package conf import ( "path/filepath" + "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/pkg/utils/random" ) @@ -112,33 +113,34 @@ type MCP struct { } type Config struct { - Force bool `json:"force" env:"FORCE"` - SiteURL string `json:"site_url" env:"SITE_URL"` - Cdn string `json:"cdn" env:"CDN"` - JwtSecret string `json:"jwt_secret" env:"JWT_SECRET"` - TokenExpiresIn int `json:"token_expires_in" env:"TOKEN_EXPIRES_IN"` - Database Database `json:"database" envPrefix:"DB_"` - Meilisearch Meilisearch `json:"meilisearch" envPrefix:"MEILISEARCH_"` - Scheme Scheme `json:"scheme"` - TempDir string `json:"temp_dir" env:"TEMP_DIR"` - BleveDir string `json:"bleve_dir" env:"BLEVE_DIR"` - DistDir string `json:"dist_dir"` - Log LogConfig `json:"log" envPrefix:"LOG_"` - DelayedStart int `json:"delayed_start" env:"DELAYED_START"` - AutoMemoryLimit int `json:"auto_memory_limit" env:"AUTO_MEMORY_LIMIT"` - MinFreeMemory int `json:"min_free_memory" env:"MIN_FREE_MEMORY"` - MaxBlockLimit int `json:"max_block_limit" env:"MAX_BLOCK_LIMIT"` - MaxConnections int `json:"max_connections" env:"MAX_CONNECTIONS"` - MaxConcurrency int `json:"max_concurrency" env:"MAX_CONCURRENCY"` - TlsInsecureSkipVerify bool `json:"tls_insecure_skip_verify" env:"TLS_INSECURE_SKIP_VERIFY"` - Tasks TasksConfig `json:"tasks" envPrefix:"TASKS_"` - Cors Cors `json:"cors" envPrefix:"CORS_"` - S3 S3 `json:"s3" envPrefix:"S3_"` - FTP FTP `json:"ftp" envPrefix:"FTP_"` - SFTP SFTP `json:"sftp" envPrefix:"SFTP_"` - MCP MCP `json:"mcp" envPrefix:"MCP_"` - LastLaunchedVersion string `json:"last_launched_version"` - ProxyAddress string `json:"proxy_address" env:"PROXY_ADDRESS"` + Force bool `json:"force" env:"FORCE"` + SiteURL string `json:"site_url" env:"SITE_URL"` + Cdn string `json:"cdn" env:"CDN"` + JwtSecret string `json:"jwt_secret" env:"JWT_SECRET"` + TokenExpiresIn int `json:"token_expires_in" env:"TOKEN_EXPIRES_IN"` + Database Database `json:"database" envPrefix:"DB_"` + Meilisearch Meilisearch `json:"meilisearch" envPrefix:"MEILISEARCH_"` + Scheme Scheme `json:"scheme"` + TempDir string `json:"temp_dir" env:"TEMP_DIR"` + BleveDir string `json:"bleve_dir" env:"BLEVE_DIR"` + DistDir string `json:"dist_dir"` + Log LogConfig `json:"log" envPrefix:"LOG_"` + DelayedStart int `json:"delayed_start" env:"DELAYED_START"` + CachePolicy cache.Policy `json:"cache_policy" env:"CACHE_POLICY"` + AutoMemoryLimit int `json:"auto_memory_limit" env:"AUTO_MEMORY_LIMIT"` + MinFreeMemory int `json:"min_free_memory" env:"MIN_FREE_MEMORY"` + MaxBlockLimit int `json:"max_block_limit" env:"MAX_BLOCK_LIMIT"` + MaxConnections int `json:"max_connections" env:"MAX_CONNECTIONS"` + MaxConcurrency int `json:"max_concurrency" env:"MAX_CONCURRENCY"` + TlsInsecureSkipVerify bool `json:"tls_insecure_skip_verify" env:"TLS_INSECURE_SKIP_VERIFY"` + Tasks TasksConfig `json:"tasks" envPrefix:"TASKS_"` + Cors Cors `json:"cors" envPrefix:"CORS_"` + S3 S3 `json:"s3" envPrefix:"S3_"` + FTP FTP `json:"ftp" envPrefix:"FTP_"` + SFTP SFTP `json:"sftp" envPrefix:"SFTP_"` + MCP MCP `json:"mcp" envPrefix:"MCP_"` + LastLaunchedVersion string `json:"last_launched_version"` + ProxyAddress string `json:"proxy_address" env:"PROXY_ADDRESS"` } func DefaultConfig(dataDir string) *Config { @@ -185,6 +187,7 @@ func DefaultConfig(dataDir string) *Config { }, }, }, + CachePolicy: cache.PolicyAuto, AutoMemoryLimit: 4, MaxConnections: 0, MaxConcurrency: 64, diff --git a/internal/conf/config_test.go b/internal/conf/config_test.go new file mode 100644 index 0000000000..231e653bd8 --- /dev/null +++ b/internal/conf/config_test.go @@ -0,0 +1,21 @@ +package conf + +import ( + "encoding/json" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/cache" +) + +func TestDefaultCachePolicy(t *testing.T) { + cfg := DefaultConfig(t.TempDir()) + if got := cfg.CachePolicy; got != cache.PolicyAuto { + t.Fatalf("default cache policy = %q, want auto", got) + } + if err := json.Unmarshal([]byte(`{}`), cfg); err != nil { + t.Fatalf("json.Unmarshal() error = %v", err) + } + if got := cfg.CachePolicy; got != cache.PolicyAuto { + t.Fatalf("cache policy after loading an old config = %q, want auto", got) + } +} diff --git a/internal/conf/var.go b/internal/conf/var.go index 6b25bcfb5b..57da166213 100644 --- a/internal/conf/var.go +++ b/internal/conf/var.go @@ -4,6 +4,8 @@ import ( "net/url" "regexp" "sync" + + "github.com/OpenListTeam/OpenList/v4/internal/cache" ) var ( @@ -25,10 +27,11 @@ var FilenameCharMap = make(map[string]string) var PrivacyReg []*regexp.Regexp var ( + CachePolicy cache.Policy = cache.PolicyAuto // 在HybridCache中使用[]byte缓存数据流的限制,内存为Go自动管理,直到GC AutoMemoryLimit uint64 = 4 * 1024 * 1024 - // 最小空闲内存,当内存不足时,HybridCache会回退到文件缓存。 - // 如果为0,HybridCache会使用文件缓存,不占用内存。 + // 最小空闲内存,当内存不足时,auto策略会回退到文件缓存。 + // 如果为0,auto策略会使用文件缓存,不占用内存。 MinFreeMemory uint64 = 16 * 1024 * 1024 // 限制HybridCache手动管理内存单次的扩容大小,超过该阈值将分多次扩容。 // MinFreeMemory大于0时,也限制 Downloader 的PartSize diff --git a/internal/hybrid_cache/buffer.go b/internal/hybrid_cache/buffer.go index 0023996e63..1bd93c775f 100644 --- a/internal/hybrid_cache/buffer.go +++ b/internal/hybrid_cache/buffer.go @@ -3,24 +3,36 @@ package hybrid_cache import ( "fmt" "io" + "sync" ) +// BufferStore is a growable in-memory backing store split into stable blocks. type BufferStore struct { - blocks [][]byte - size int64 + // layoutMu protects blocks and size. ReadAt and WriteAt intentionally hold + // a shared lock so independent block ranges can be streamed concurrently; + // callers must synchronize overlapping byte ranges. + layoutMu sync.RWMutex + blocks [][]byte + size int64 } func (m *BufferStore) Size() int64 { + m.layoutMu.RLock() + defer m.layoutMu.RUnlock() return m.size } -// 用于存储不复用的[]byte +// Append adds a caller-owned block without copying it. func (m *BufferStore) Append(buf []byte) { + m.layoutMu.Lock() + defer m.layoutMu.Unlock() m.size += int64(len(buf)) m.blocks = append(m.blocks, buf) } func (m *BufferStore) Close() error { + m.layoutMu.Lock() + defer m.layoutMu.Unlock() if len(m.blocks) > 0 { clear(m.blocks) m.blocks = m.blocks[:0] @@ -30,6 +42,8 @@ func (m *BufferStore) Close() error { } func (m *BufferStore) ReadAt(p []byte, off int64) (int, error) { + m.layoutMu.RLock() + defer m.layoutMu.RUnlock() if len(p) == 0 { return 0, nil } @@ -55,6 +69,8 @@ func (m *BufferStore) ReadAt(p []byte, off int64) (int, error) { } func (m *BufferStore) WriteAt(p []byte, off int64) (int, error) { + m.layoutMu.RLock() + defer m.layoutMu.RUnlock() if len(p) == 0 { return 0, nil } @@ -80,6 +96,8 @@ func (m *BufferStore) WriteAt(p []byte, off int64) (int, error) { } func (m *BufferStore) GrowTo(size int64) (err error) { + m.layoutMu.Lock() + defer m.layoutMu.Unlock() if size <= m.size { return nil } diff --git a/internal/hybrid_cache/hybrid_cache.go b/internal/hybrid_cache/hybrid_cache.go index c69147937e..645de9b6d0 100644 --- a/internal/hybrid_cache/hybrid_cache.go +++ b/internal/hybrid_cache/hybrid_cache.go @@ -2,76 +2,122 @@ package hybrid_cache import ( "errors" + "fmt" "io" + "math" "runtime" + "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/mem" "github.com/OpenListTeam/OpenList/v4/pkg/buffer" "github.com/OpenListTeam/OpenList/v4/pkg/utils" ) -// 线程不安全,单线程使用,或者外部加锁保护 +// HybridCache stores an in-memory prefix and may spill subsequent data to +// disk. It is not safe for concurrent use without external synchronization. type HybridCache struct { - blockSize uint64 - memoryStore mem.LinearMemory - memoryOffset uint64 - backingStore BackingStore - backingOffset uint64 - cleanup runtime.Cleanup + blockSize uint64 + memoryStore mem.LinearMemory + memoryOffset uint64 + backingStore BackingStore + backingOffset uint64 + cleanup runtime.Cleanup + spillOnMemoryFailure bool + memoryBacking bool + memoryCeiling int64 + memoryReservation *mem.Reservation } -// HybridCache本身是一个大的Block,支持分块成多个小的Block - -// 分配一个新的Block,支持读写,大小为size +// AllocBlock appends a readable and writable block of size bytes. func (hc *HybridCache) AllocBlock(size uint64) (buffer.Block, error) { retry: if hc.backingStore != nil { - if err := hc.backingStore.GrowTo(int64(hc.backingOffset + size)); err != nil { + base, err := hc.growBackingStore(size) + if err != nil { return nil, err } - base := hc.backingOffset - hc.backingOffset += size fs := buffer.NewBlockAdapter( io.NewOffsetWriter(hc.backingStore, int64(base)), io.NewSectionReader(hc.backingStore, int64(base), int64(size)), ) return fs, nil } - all, err := hc.memoryStore.Reallocate(hc.memoryOffset + size) + var all []byte + var err error + if hc.exceedsMemoryCeiling(hc.memoryOffset, size) { + err = mem.ErrNotEnoughMemory + } else { + all, err = hc.memoryStore.Reallocate(hc.memoryOffset + size) + } if err == nil { start := hc.memoryOffset hc.memoryOffset += size return buffer.NewByteBlock(all[start : start+size]), nil } + if !hc.spillOnMemoryFailure { + return nil, err + } if err2 := hc.initFileCache(); err2 != nil { return nil, errors.Join(err, err2) } + hc.shrinkReservationToMemory() goto retry } func (hc *HybridCache) allocWriteAtSeeker(size uint64) (buffer.WriteAtSeeker, error) { retry: if hc.backingStore != nil { - if err := hc.backingStore.GrowTo(int64(hc.backingOffset + size)); err != nil { + base, err := hc.growBackingStore(size) + if err != nil { return nil, err } - base := hc.backingOffset - hc.backingOffset += size return io.NewOffsetWriter(hc.backingStore, int64(base)), nil } - all, err := hc.memoryStore.Reallocate(hc.memoryOffset + size) + var all []byte + var err error + if hc.exceedsMemoryCeiling(hc.memoryOffset, size) { + err = mem.ErrNotEnoughMemory + } else { + all, err = hc.memoryStore.Reallocate(hc.memoryOffset + size) + } if err == nil { start := hc.memoryOffset hc.memoryOffset += size return io.NewOffsetWriter(buffer.NewByteBlock(all[start:start+size]), 0), nil } + if !hc.spillOnMemoryFailure { + return nil, err + } if err2 := hc.initFileCache(); err2 != nil { return nil, errors.Join(err, err2) } + hc.shrinkReservationToMemory() goto retry } +func (hc *HybridCache) growBackingStore(size uint64) (uint64, error) { + if hc.memoryBacking && hc.exceedsMemoryCeiling(hc.backingOffset, size) { + return 0, mem.ErrNotEnoughMemory + } + if hc.backingOffset > math.MaxInt64 || size > math.MaxInt64-hc.backingOffset { + return 0, mem.ErrNotEnoughMemory + } + base := hc.backingOffset + target := base + size + if err := hc.backingStore.GrowTo(int64(target)); err != nil { + return 0, err + } + hc.backingOffset = target + return base, nil +} + +func (hc *HybridCache) shrinkReservationToMemory() { + if hc.memoryReservation != nil { + hc.memoryReservation.Resize(hc.memoryOffset) + } +} + func (hc *HybridCache) NextBlock() (buffer.Block, error) { return hc.AllocBlock(hc.blockSize) } @@ -95,8 +141,20 @@ func (hc *HybridCache) RewindOneBlock() { hc.RewindBySize(hc.blockSize) } +func (hc *HybridCache) exceedsMemoryCeiling(offset, size uint64) bool { + if hc.memoryCeiling < 0 { + return false + } + ceiling := uint64(hc.memoryCeiling) + return offset > ceiling || size > ceiling-offset +} + func (hc *HybridCache) initFileCache() error { - file, err := NewFileStore(int64(hc.blockSize)) + initialSize := hc.blockSize + if hc.memoryCeiling < 0 { + initialSize = 0 + } + file, err := NewFileStore(int64(initialSize)) if err != nil { return err } @@ -120,6 +178,10 @@ func (hc *HybridCache) Close() error { hc.backingStore = nil hc.backingOffset = 0 } + if hc.memoryReservation != nil { + hc.memoryReservation.Release() + hc.memoryReservation = nil + } return err } @@ -187,6 +249,33 @@ func (hc *HybridCache) WriteAt(p []byte, off int64) (n int, err error) { return n + nn, err } +// Write appends p to the cache. It is not safe for concurrent use. +func (hc *HybridCache) Write(p []byte) (n int, err error) { + for len(p) > 0 { + chunkSize := len(p) + if hc.blockSize > 0 && uint64(chunkSize) > hc.blockSize { + chunkSize = int(hc.blockSize) + } + w, allocErr := hc.allocWriteAtSeeker(uint64(chunkSize)) + if allocErr != nil { + return n, allocErr + } + nn, writeErr := w.Write(p[:chunkSize]) + n += nn + if nn < chunkSize { + hc.RewindBySize(uint64(chunkSize - nn)) + if writeErr == nil { + writeErr = io.ErrShortWrite + } + } + if writeErr != nil { + return n, writeErr + } + p = p[chunkSize:] + } + return n, nil +} + func (hc *HybridCache) CopyFromN(src io.Reader, n int64) (written int64, err error) { limit := n for limit > 0 { @@ -208,32 +297,82 @@ func (hc *HybridCache) CopyFromN(src io.Reader, n int64) (written int64, err err return written, nil } -// HybridCache 线程不安全,单线程使用,或者外部加锁保护 -func NewHybridCache(blockSize, maxMemorySize uint64) (hc *HybridCache, err error) { - if conf.MinFreeMemory > 0 { - // 策略1: Go自动内存管理 - if maxMemorySize <= conf.AutoMemoryLimit { - return &HybridCache{backingStore: &BufferStore{}, blockSize: blockSize}, nil +func selectPolicy(requested cache.Policy, memoryCeiling int64) (cache.Policy, error) { + switch requested { + case cache.PolicyMemory, cache.PolicyDisk: + return requested, nil + case cache.PolicyAuto: + if memoryCeiling < 0 { + return cache.PolicyDisk, nil } + return cache.PolicyMemory, nil + default: + return cache.PolicyInherit, fmt.Errorf("invalid cache policy %q", requested) + } +} - // 策略2: 手动内存管理 - if maxMemorySize >= blockSize { - var m mem.LinearMemory - // 手动管理内存,Uinx Mmap 或者 Windows VirtualAlloc - if m, err = mem.NewGuardedMemory(blockSize, maxMemorySize); err == nil { - hc = &HybridCache{memoryStore: m, blockSize: blockSize} - } +// NewHybridCache creates a non-thread-safe cache using the requested policy. +func NewHybridCache(blockSize uint64, memoryCeiling int64, requested cache.Policy) (hc *HybridCache, err error) { + return newHybridCache(blockSize, memoryCeiling, requested, mem.CacheMemoryBudget) +} + +func newHybridCache(blockSize uint64, memoryCeiling int64, requested cache.Policy, budget *mem.Budget) (hc *HybridCache, err error) { + if memoryCeiling < 0 && blockSize == 0 { + return nil, fmt.Errorf("block size must be positive when memory ceiling is unknown") + } + if memoryCeiling > 0 { + blockSize = min(blockSize, uint64(memoryCeiling)) + if blockSize == 0 { + return nil, fmt.Errorf("block size must be positive for a non-empty cache") + } + } + + selected, err := selectPolicy(requested, memoryCeiling) + if err != nil { + return nil, err + } + hc = &HybridCache{blockSize: blockSize, memoryCeiling: memoryCeiling} + if requested == cache.PolicyAuto && selected == cache.PolicyMemory && memoryCeiling > 0 { + if budget == nil { + return nil, errors.New("cache memory budget is unavailable") + } + var ok bool + hc.memoryReservation, ok = budget.Reserve(uint64(memoryCeiling)) + if !ok { + selected = cache.PolicyDisk } } - // 策略3: 文件后备 - if hc == nil { - hc = &HybridCache{blockSize: blockSize} - // 文件 - if err2 := hc.initFileCache(); err2 != nil { - return nil, errors.Join(err, err2) + if selected == cache.PolicyDisk { + if err := hc.initFileCache(); err != nil { + return nil, err } + return hc, nil + } + + if memoryCeiling < 0 || uint64(memoryCeiling) <= conf.AutoMemoryLimit { + hc.backingStore = &BufferStore{} + hc.memoryBacking = true + return hc, nil + } + + hc.memoryStore, err = mem.NewManagedMemory(blockSize, uint64(memoryCeiling)) + if err == nil { + hc.spillOnMemoryFailure = requested == cache.PolicyAuto + return hc, nil + } + + if hc.memoryReservation != nil { + hc.memoryReservation.Release() + hc.memoryReservation = nil + } + if requested == cache.PolicyMemory { + return nil, err + } + if fileErr := hc.initFileCache(); fileErr != nil { + return nil, errors.Join(err, fileErr) } return hc, nil } var _ buffer.Block = (*HybridCache)(nil) +var _ io.Writer = (*HybridCache)(nil) diff --git a/internal/hybrid_cache/policy_test.go b/internal/hybrid_cache/policy_test.go new file mode 100644 index 0000000000..82328e2cee --- /dev/null +++ b/internal/hybrid_cache/policy_test.go @@ -0,0 +1,331 @@ +package hybrid_cache + +import ( + "errors" + "io" + "os" + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/cache" + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/mem" +) + +func TestSelectPolicy(t *testing.T) { + tests := []struct { + name string + requested cache.Policy + ceiling int64 + want cache.Policy + }{ + {"explicit memory", cache.PolicyMemory, -1, cache.PolicyMemory}, + {"explicit disk", cache.PolicyDisk, 1024, cache.PolicyDisk}, + {"auto unknown", cache.PolicyAuto, -1, cache.PolicyDisk}, + {"auto empty", cache.PolicyAuto, 0, cache.PolicyMemory}, + {"auto known", cache.PolicyAuto, 1024, cache.PolicyMemory}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := selectPolicy(tt.requested, tt.ceiling) + if err != nil { + t.Fatalf("selectPolicy() error = %v", err) + } + if got != tt.want { + t.Fatalf("selectPolicy() = %q, want %q", got, tt.want) + } + }) + } + if _, err := selectPolicy(cache.PolicyInherit, 1); err == nil { + t.Fatal("selectPolicy() expected an error for inherit") + } +} + +func TestHybridCacheDiskPolicy(t *testing.T) { + withCacheConfig(t, 0) + hc, err := NewHybridCache(4, 8, cache.PolicyDisk) + if err != nil { + t.Fatalf("NewHybridCache() error = %v", err) + } + store, ok := hc.backingStore.(*singleFileStore) + if !ok || hc.memoryStore != nil { + t.Fatalf("disk policy initialized unexpected stores: memory=%T backing=%T", hc.memoryStore, hc.backingStore) + } + name := store.Name() + if err := hc.Close(); err != nil { + t.Fatalf("Close() error = %v", err) + } + if _, err := os.Stat(name); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("cache file still exists after Close(): %v", err) + } +} + +func TestHybridCachePolicyBackingSelection(t *testing.T) { + tests := []struct { + name string + requested cache.Policy + ceiling int64 + budgetCapacity uint64 + wantMemory bool + }{ + {name: "memory ignores budget", requested: cache.PolicyMemory, ceiling: 8, wantMemory: true}, + {name: "disk uses disk", requested: cache.PolicyDisk, ceiling: 8}, + {name: "auto admitted uses memory", requested: cache.PolicyAuto, ceiling: 8, budgetCapacity: 8, wantMemory: true}, + {name: "auto rejected uses disk", requested: cache.PolicyAuto, ceiling: 8, budgetCapacity: 7}, + {name: "auto unknown uses disk", requested: cache.PolicyAuto, ceiling: -1}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + withCacheConfig(t, 8) + hc, err := newHybridCache(4, tt.ceiling, tt.requested, mem.NewBudget(tt.budgetCapacity)) + if err != nil { + t.Fatalf("newHybridCache() error = %v", err) + } + + if hc.memoryBacking != tt.wantMemory { + t.Fatalf("memory backing = %v, want %v", hc.memoryBacking, tt.wantMemory) + } + if _, isBuffer := hc.backingStore.(*BufferStore); isBuffer != tt.wantMemory { + t.Fatalf("buffer backing = %v, want %v", isBuffer, tt.wantMemory) + } + if _, isFile := hc.backingStore.(*singleFileStore); isFile == tt.wantMemory { + t.Fatalf("file backing = %v, want %v", isFile, !tt.wantMemory) + } + + input := []byte("12345678") + if tt.ceiling < 0 { + input = []byte("unknown") + } + if _, err := hc.Write(input); err != nil { + t.Fatalf("Write() error = %v", err) + } + if hc.memoryOffset != 0 { + t.Fatalf("linear memory offset = %d, want 0", hc.memoryOffset) + } + if hc.backingOffset != uint64(len(input)) { + t.Fatalf("backing offset = %d, want %d", hc.backingOffset, len(input)) + } + got := make([]byte, len(input)) + if _, err := hc.ReadAt(got, 0); err != nil { + t.Fatalf("ReadAt() error = %v", err) + } + if string(got) != string(input) { + t.Fatalf("cache = %q, want %q", got, input) + } + if err := hc.Close(); err != nil { + t.Fatalf("Close() error = %v", err) + } + }) + } +} + +func TestHybridCacheAutoRejectsWholeCeiling(t *testing.T) { + withCacheConfig(t, 1024) + hc, err := newHybridCache(4, 8, cache.PolicyAuto, mem.NewBudget(7)) + if err != nil { + t.Fatalf("NewHybridCache() error = %v", err) + } + t.Cleanup(func() { _ = hc.Close() }) + if hc.memoryStore != nil || hc.memoryBacking { + t.Fatalf("auto policy initialized memory stores: memory=%T backing=%T", hc.memoryStore, hc.backingStore) + } + if _, ok := hc.backingStore.(*singleFileStore); !ok { + t.Fatalf("auto policy backing = %T, want singleFileStore", hc.backingStore) + } +} + +func TestHybridCacheZeroCeilingRejectsUnexpectedWrites(t *testing.T) { + withCacheConfig(t, 8) + hc, err := newHybridCache(4, 0, cache.PolicyAuto, mem.NewBudget(0)) + if err != nil { + t.Fatalf("newHybridCache() error = %v", err) + } + t.Cleanup(func() { _ = hc.Close() }) + if _, err := hc.Write([]byte{1}); !errors.Is(err, mem.ErrNotEnoughMemory) { + t.Fatalf("Write() error = %v, want ErrNotEnoughMemory", err) + } +} + +func TestHybridCacheStrictMemoryDoesNotSpill(t *testing.T) { + withCacheConfig(t, 0) + hc, err := NewHybridCache(4, 8, cache.PolicyMemory) + if err != nil { + t.Fatalf("NewHybridCache() error = %v", err) + } + t.Cleanup(func() { _ = hc.Close() }) + if _, err := hc.AllocBlock(4); err != nil { + t.Fatalf("AllocBlock() error = %v", err) + } + if _, err := hc.AllocBlock(5); !errors.Is(err, mem.ErrNotEnoughMemory) { + t.Fatalf("AllocBlock() error = %v, want ErrNotEnoughMemory", err) + } + if hc.backingStore != nil { + t.Fatalf("strict memory policy spilled to %T", hc.backingStore) + } +} + +func TestHybridCacheUnknownMemory(t *testing.T) { + withCacheConfig(t, 0) + hc, err := NewHybridCache(3, -1, cache.PolicyMemory) + if err != nil { + t.Fatalf("NewHybridCache() error = %v", err) + } + t.Cleanup(func() { _ = hc.Close() }) + if _, ok := hc.backingStore.(*BufferStore); !ok { + t.Fatalf("unknown memory cache backing = %T, want BufferStore", hc.backingStore) + } + if _, err := hc.Write([]byte("abcdefg")); err != nil { + t.Fatalf("Write() error = %v", err) + } + got := make([]byte, 7) + if _, err := hc.ReadAt(got, 0); err != nil { + t.Fatalf("ReadAt() error = %v", err) + } + if string(got) != "abcdefg" || hc.Size() != 7 { + t.Fatalf("cache = %q size=%d", got, hc.Size()) + } +} + +func TestHybridCacheReservationsPreventOversubscription(t *testing.T) { + withCacheConfig(t, 8) + budget := mem.NewBudget(8) + first, err := newHybridCache(4, 8, cache.PolicyAuto, budget) + if err != nil { + t.Fatalf("first newHybridCache() error = %v", err) + } + if budget.Reserved() != 8 || !first.memoryBacking { + t.Fatalf("first cache = memory:%v reserved:%d", first.memoryBacking, budget.Reserved()) + } + + second, err := newHybridCache(4, 8, cache.PolicyAuto, budget) + if err != nil { + t.Fatalf("second newHybridCache() error = %v", err) + } + if second.memoryBacking || budget.Reserved() != 8 { + t.Fatalf("second cache = memory:%v reserved:%d", second.memoryBacking, budget.Reserved()) + } + if err := second.Close(); err != nil { + t.Fatalf("second Close() error = %v", err) + } + if err := first.Close(); err != nil { + t.Fatalf("first Close() error = %v", err) + } + if budget.Reserved() != 0 { + t.Fatalf("reserved after close = %d", budget.Reserved()) + } + + third, err := newHybridCache(4, 8, cache.PolicyAuto, budget) + if err != nil { + t.Fatalf("third newHybridCache() error = %v", err) + } + if !third.memoryBacking { + t.Fatal("third cache did not reuse released memory budget") + } + if err := third.Close(); err != nil { + t.Fatalf("third Close() error = %v", err) + } +} + +func TestHybridCacheStrictMemoryIgnoresBudget(t *testing.T) { + withCacheConfig(t, 8) + budget := mem.NewBudget(0) + hc, err := newHybridCache(4, 8, cache.PolicyMemory, budget) + if err != nil { + t.Fatalf("newHybridCache() error = %v", err) + } + if budget.Reserved() != 0 { + t.Fatalf("strict memory reserved budget = %d, want 0", budget.Reserved()) + } + if err := hc.Close(); err != nil { + t.Fatalf("Close() error = %v", err) + } +} + +func TestHybridCacheUnknownMemoryIgnoresBudget(t *testing.T) { + withCacheConfig(t, 8) + budget := mem.NewBudget(0) + hc, err := newHybridCache(3, -1, cache.PolicyMemory, budget) + if err != nil { + t.Fatalf("newHybridCache() error = %v", err) + } + if _, err := hc.Write([]byte("1234567")); err != nil { + t.Fatalf("Write() error = %v", err) + } + if err := hc.Close(); err != nil { + t.Fatalf("Close() error = %v", err) + } + if budget.Reserved() != 0 { + t.Fatalf("reserved after close = %d", budget.Reserved()) + } +} + +func TestHybridCacheAutoSpillKeepsMemoryPrefix(t *testing.T) { + withCacheConfig(t, 0) + budget := mem.NewBudget(4) + reservation, ok := budget.Reserve(4) + if !ok { + t.Fatal("Reserve(4) failed") + } + hc := &HybridCache{ + blockSize: 2, + memoryStore: &limitedMemory{buf: make([]byte, 0, 2)}, + spillOnMemoryFailure: true, + memoryCeiling: 2, + memoryReservation: reservation, + } + t.Cleanup(func() { _ = hc.Close() }) + if _, err := hc.Write([]byte("abcd")); err != nil { + t.Fatalf("Write() error = %v", err) + } + if hc.memoryOffset != 2 || hc.backingOffset != 2 { + t.Fatalf("offsets = memory:%d disk:%d, want 2 and 2", hc.memoryOffset, hc.backingOffset) + } + if budget.Reserved() != 2 { + t.Fatalf("reserved after spill = %d, want 2", budget.Reserved()) + } + got := make([]byte, 4) + if _, err := hc.ReadAt(got, 0); err != nil { + t.Fatalf("ReadAt() error = %v", err) + } + if string(got) != "abcd" { + t.Fatalf("cache = %q, want abcd", got) + } +} + +type limitedMemory struct { + buf []byte +} + +func (m *limitedMemory) Reallocate(size uint64) ([]byte, error) { + if size > uint64(cap(m.buf)) { + return nil, mem.ErrNotEnoughMemory + } + m.buf = m.buf[:size] + return m.buf, nil +} + +func (m *limitedMemory) Free() error { + m.buf = nil + return nil +} + +func withCacheConfig(t *testing.T, autoMemoryLimit uint64) { + t.Helper() + oldConf := conf.Conf + oldLimit := conf.AutoMemoryLimit + oldBudgetCapacity := mem.CacheMemoryBudget.Capacity() + oldBudgetReserved := mem.CacheMemoryBudget.Reserved() + conf.Conf = &conf.Config{TempDir: t.TempDir()} + conf.AutoMemoryLimit = autoMemoryLimit + mem.CacheMemoryBudget.SetCapacity(1 << 40) + t.Cleanup(func() { + conf.Conf = oldConf + conf.AutoMemoryLimit = oldLimit + if reserved := mem.CacheMemoryBudget.Reserved(); reserved != oldBudgetReserved { + t.Errorf("cache memory reserved after test = %d, want %d", reserved, oldBudgetReserved) + } + mem.CacheMemoryBudget.SetCapacity(oldBudgetCapacity) + }) +} + +var _ io.Writer = (*HybridCache)(nil) diff --git a/internal/mem/budget.go b/internal/mem/budget.go new file mode 100644 index 0000000000..439ab0b3de --- /dev/null +++ b/internal/mem/budget.go @@ -0,0 +1,160 @@ +package mem + +import ( + "runtime" + "sync" +) + +// CacheMemoryBudget accounts for memory reserved by active HybridCache instances. +// Bootstrap configures its capacity from the effective available memory. +var CacheMemoryBudget = NewBudget(0) + +// Budget is a process-local capacity shared by cache reservations. +type Budget struct { + mu sync.Mutex + capacity uint64 + reserved uint64 +} + +// NewBudget creates an empty budget with the given capacity. +func NewBudget(capacity uint64) *Budget { + return &Budget{capacity: capacity} +} + +// SetCapacity changes the maximum number of bytes that reservations may hold. +// Existing reservations remain valid if the capacity is reduced below them. +func (b *Budget) SetCapacity(capacity uint64) { + b.mu.Lock() + b.capacity = capacity + b.mu.Unlock() +} + +// Capacity returns the maximum number of bytes the budget can reserve. +func (b *Budget) Capacity() uint64 { + b.mu.Lock() + defer b.mu.Unlock() + return b.capacity +} + +// Reserved returns the number of bytes held by active reservations. +func (b *Budget) Reserved() uint64 { + b.mu.Lock() + defer b.mu.Unlock() + return b.reserved +} + +// Available returns the number of bytes that remain reservable. +func (b *Budget) Available() uint64 { + b.mu.Lock() + defer b.mu.Unlock() + if b.reserved >= b.capacity { + return 0 + } + return b.capacity - b.reserved +} + +// Reserve atomically claims size bytes until the returned reservation is +// released. A zero-sized reservation may be grown later. +func (b *Budget) Reserve(size uint64) (*Reservation, bool) { + if !b.resize(0, size) { + return nil, false + } + state := &reservationState{budget: b, size: size} + reservation := &Reservation{state: state} + reservation.cleanup = runtime.AddCleanup(reservation, releaseReservation, state) + return reservation, true +} + +func (b *Budget) release(size uint64) { + b.mu.Lock() + defer b.mu.Unlock() + if size > b.reserved { + panic("cache memory budget reservation underflow") + } + b.reserved -= size +} + +func (b *Budget) resize(current, target uint64) bool { + b.mu.Lock() + defer b.mu.Unlock() + if current > b.reserved { + panic("cache memory budget reservation exceeds total reserved memory") + } + if target > current { + growth := target - current + if b.reserved > b.capacity || growth > b.capacity-b.reserved { + return false + } + b.reserved += growth + } else { + b.reserved -= current - target + } + return true +} + +type reservationState struct { + mu sync.Mutex + budget *Budget + size uint64 + released bool +} + +func releaseReservation(state *reservationState) { + state.mu.Lock() + defer state.mu.Unlock() + if state.released { + return + } + state.budget.release(state.size) + state.size = 0 + state.released = true +} + +// Reservation is an idempotently releasable claim on a Budget. +type Reservation struct { + state *reservationState + cleanup runtime.Cleanup +} + +// Resize changes the reservation while preserving atomic budget accounting. +// It returns false when growing would exceed the budget or after Release. +func (r *Reservation) Resize(size uint64) bool { + if r == nil || r.state == nil { + return false + } + state := r.state + defer runtime.KeepAlive(r) + state.mu.Lock() + defer state.mu.Unlock() + if state.released { + return false + } + current := state.size + if !state.budget.resize(current, size) { + return false + } + state.size = size + return true +} + +// Size returns the number of bytes currently held by the reservation. +func (r *Reservation) Size() uint64 { + if r == nil || r.state == nil { + return 0 + } + state := r.state + defer runtime.KeepAlive(r) + state.mu.Lock() + defer state.mu.Unlock() + return state.size +} + +// Release returns the reserved capacity to the budget. It is idempotent. +func (r *Reservation) Release() { + if r == nil || r.state == nil { + return + } + defer runtime.KeepAlive(r) + r.cleanup.Stop() + releaseReservation(r.state) +} diff --git a/internal/mem/budget_test.go b/internal/mem/budget_test.go new file mode 100644 index 0000000000..890f125245 --- /dev/null +++ b/internal/mem/budget_test.go @@ -0,0 +1,106 @@ +package mem + +import ( + "math" + "sync" + "testing" +) + +func TestBudgetReservationLifecycle(t *testing.T) { + budget := NewBudget(10) + reservation, ok := budget.Reserve(6) + if !ok { + t.Fatal("Reserve(6) failed") + } + if budget.Reserved() != 6 || budget.Available() != 4 { + t.Fatalf("budget after reserve = reserved:%d available:%d", budget.Reserved(), budget.Available()) + } + if reservation.Resize(11) { + t.Fatal("Resize(11) succeeded beyond capacity") + } + if !reservation.Resize(8) || budget.Reserved() != 8 { + t.Fatalf("budget after grow = reserved:%d", budget.Reserved()) + } + if !reservation.Resize(3) || budget.Reserved() != 3 { + t.Fatalf("budget after shrink = reserved:%d", budget.Reserved()) + } + reservation.Release() + reservation.Release() + if budget.Reserved() != 0 || budget.Available() != 10 { + t.Fatalf("budget after release = reserved:%d available:%d", budget.Reserved(), budget.Available()) + } + if reservation.Resize(1) { + t.Fatal("Resize() succeeded after release") + } +} + +func TestBudgetConcurrentReservationsDoNotOversubscribe(t *testing.T) { + const ( + capacity = 32 + claim = 4 + workers = 64 + ) + budget := NewBudget(capacity) + start := make(chan struct{}) + results := make(chan *Reservation, workers) + var wg sync.WaitGroup + for range workers { + wg.Add(1) + go func() { + defer wg.Done() + <-start + reservation, ok := budget.Reserve(claim) + if !ok { + results <- nil + return + } + results <- reservation + }() + } + close(start) + wg.Wait() + close(results) + + var reservations []*Reservation + for reservation := range results { + if reservation != nil { + reservations = append(reservations, reservation) + } + } + if len(reservations) != capacity/claim { + t.Fatalf("successful reservations = %d, want %d", len(reservations), capacity/claim) + } + if budget.Reserved() != capacity || budget.Available() != 0 { + t.Fatalf("full budget = reserved:%d available:%d", budget.Reserved(), budget.Available()) + } + for _, reservation := range reservations { + reservation.Release() + } + if budget.Reserved() != 0 { + t.Fatalf("reserved after releases = %d", budget.Reserved()) + } +} + +func TestBudgetBoundaries(t *testing.T) { + budget := NewBudget(math.MaxUint64) + reservation, ok := budget.Reserve(math.MaxUint64 - 5) + if !ok { + t.Fatal("large reservation failed") + } + if _, ok := budget.Reserve(6); ok { + t.Fatal("overflowing reservation succeeded") + } + if budget.Available() != 5 { + t.Fatalf("available = %d, want 5", budget.Available()) + } + budget.SetCapacity(1) + if budget.Available() != 0 { + t.Fatalf("available below existing reservations = %d, want 0", budget.Available()) + } + reservation.Release() + if zero, ok := budget.Reserve(0); !ok { + t.Fatal("zero reservation failed") + } else { + zero.Release() + } +} diff --git a/internal/mem/mem_unix.go b/internal/mem/mem_unix.go index bd1979298e..a031ad8581 100644 --- a/internal/mem/mem_unix.go +++ b/internal/mem/mem_unix.go @@ -39,12 +39,7 @@ func NewMemory(cap, max uint64) (LinearMemory, error) { // - len(buf) is the already committed memory, // - cap(buf) is the reserved address space. type mmappedMemory struct { - buf []byte - growCheck GrowCheck -} - -func (m *mmappedMemory) SetGrowCheck(c GrowCheck) { - m.growCheck = c + buf []byte } func (m *mmappedMemory) Reallocate(size uint64) ([]byte, error) { @@ -58,12 +53,6 @@ func (m *mmappedMemory) Reallocate(size uint64) ([]byte, error) { new = min(max(size, new), res) new = (new + rnd) &^ rnd - if m.growCheck != nil { - if err := m.growCheck(new - com); err != nil { - return nil, err - } - } - // Commit additional memory up to new bytes. err := unix.Mprotect(m.buf[com:new], unix.PROT_READ|unix.PROT_WRITE) if err != nil { diff --git a/internal/mem/mem_windows.go b/internal/mem/mem_windows.go index e7a4bb27ef..ba75aeb41a 100644 --- a/internal/mem/mem_windows.go +++ b/internal/mem/mem_windows.go @@ -39,13 +39,8 @@ func NewMemory(cap, max uint64) (LinearMemory, error) { // - len(buf) is the already committed memory, // - cap(buf) is the reserved address space. type virtualMemory struct { - buf []byte - addr uintptr - growCheck GrowCheck -} - -func (m *virtualMemory) SetGrowCheck(c GrowCheck) { - m.growCheck = c + buf []byte + addr uintptr } func (m *virtualMemory) Reallocate(size uint64) ([]byte, error) { @@ -59,12 +54,6 @@ func (m *virtualMemory) Reallocate(size uint64) ([]byte, error) { new = min(max(size, new), res) new = (new + rnd) &^ rnd - if m.growCheck != nil { - if err := m.growCheck(new - com); err != nil { - return nil, err - } - } - // Commit additional memory up to new bytes. _, err := windows.VirtualAlloc(m.addr, uintptr(new), windows.MEM_COMMIT, windows.PAGE_READWRITE) if err != nil { diff --git a/internal/mem/snapshot.go b/internal/mem/snapshot.go new file mode 100644 index 0000000000..92bba2c160 --- /dev/null +++ b/internal/mem/snapshot.go @@ -0,0 +1,77 @@ +package mem + +import ( + "os" + + gopsutilmem "github.com/shirou/gopsutil/v4/mem" +) + +const ( + MemorySourceHost = "host" + MemorySourceCgroupV1 = "cgroup_v1" + MemorySourceCgroupV2 = "cgroup_v2" +) + +// MemorySnapshot describes the effective memory boundary visible to the +// current process. Available is always capped by Limit when Limit is known. +type MemorySnapshot struct { + Limit uint64 + Used uint64 + Available uint64 + Source string +} + +type memoryFileReader func(string) ([]byte, error) + +// GetMemorySnapshot returns a zero-value snapshot with any error so callers +// can fail closed instead of using incomplete memory information. +func GetMemorySnapshot() (MemorySnapshot, error) { + virtualMemory, hostErr := gopsutilmem.VirtualMemory() + var host *MemorySnapshot + if virtualMemory != nil { + available := min(virtualMemory.Available, virtualMemory.Total) + host = &MemorySnapshot{ + Limit: virtualMemory.Total, + Used: virtualMemory.Used, + Available: available, + Source: MemorySourceHost, + } + } + return platformMemorySnapshot(host, hostErr, os.ReadFile) +} + +type cgroupMemory struct { + limit uint64 + available uint64 + source string +} + +func combineMemorySnapshots(host *MemorySnapshot, cgroup cgroupMemory) MemorySnapshot { + if host == nil { + available := min(cgroup.available, cgroup.limit) + return MemorySnapshot{ + Limit: cgroup.limit, + Used: saturatingSub(cgroup.limit, available), + Available: available, + Source: cgroup.source, + } + } + if cgroup.limit >= host.Limit && cgroup.available >= host.Available { + return *host + } + limit := min(host.Limit, cgroup.limit) + available := min(host.Available, cgroup.available, limit) + return MemorySnapshot{ + Limit: limit, + Used: saturatingSub(limit, available), + Available: available, + Source: cgroup.source, + } +} + +func saturatingSub(left, right uint64) uint64 { + if right >= left { + return 0 + } + return left - right +} diff --git a/internal/mem/snapshot_cgroup.go b/internal/mem/snapshot_cgroup.go new file mode 100644 index 0000000000..0f179904c4 --- /dev/null +++ b/internal/mem/snapshot_cgroup.go @@ -0,0 +1,266 @@ +package mem + +import ( + "bufio" + "errors" + "fmt" + "path" + "strconv" + "strings" +) + +// This parser intentionally builds on every platform so its proc/cgroup +// fixtures can be tested outside Linux. Only snapshot_linux.go calls it in +// production. + +const ( + procSelfCgroup = "/proc/self/cgroup" + procSelfMountInfo = "/proc/self/mountinfo" +) + +type cgroupMembership struct { + v2Path string + v2Found bool + memoryPath string + memoryFound bool +} + +type cgroupMount struct { + root string + mountPoint string + fsType string + memory bool +} + +func cgroupAwareMemorySnapshot(host *MemorySnapshot, hostErr error, readFile memoryFileReader) (MemorySnapshot, error) { + cgroup, found, err := readCgroupMemory(readFile) + if err == nil && found { + return combineMemorySnapshots(host, cgroup), nil + } + if err != nil { + return MemorySnapshot{}, errors.Join(hostErr, fmt.Errorf("read cgroup memory: %w", err)) + } + if host != nil { + return *host, nil + } + if hostErr != nil { + return MemorySnapshot{}, hostErr + } + return MemorySnapshot{}, errors.New("memory information is unavailable") +} + +func readCgroupMemory(readFile memoryFileReader) (cgroupMemory, bool, error) { + cgroupData, err := readFile(procSelfCgroup) + if err != nil { + return cgroupMemory{}, false, err + } + mountInfo, err := readFile(procSelfMountInfo) + if err != nil { + return cgroupMemory{}, false, err + } + membership, err := parseCgroupMembership(string(cgroupData)) + if err != nil { + return cgroupMemory{}, false, err + } + mounts, err := parseCgroupMounts(string(mountInfo)) + if err != nil { + return cgroupMemory{}, false, err + } + + if membership.v2Found { + for _, mount := range mounts { + if mount.fsType != "cgroup2" { + continue + } + base, ok := resolveCgroupPath(mount, membership.v2Path) + if !ok { + continue + } + return readCgroupHierarchy(readFile, base, mount.mountPoint, "memory.max", "memory.current", "max", MemorySourceCgroupV2) + } + if !membership.memoryFound { + return cgroupMemory{}, false, errors.New("cgroup v2 membership found without a matching mount") + } + } + + if membership.memoryFound { + for _, mount := range mounts { + if mount.fsType != "cgroup" || !mount.memory { + continue + } + base, ok := resolveCgroupPath(mount, membership.memoryPath) + if !ok { + continue + } + return readCgroupHierarchy(readFile, base, mount.mountPoint, "memory.limit_in_bytes", "memory.usage_in_bytes", "", MemorySourceCgroupV1) + } + return cgroupMemory{}, false, errors.New("cgroup v1 memory membership found without a matching mount") + } + + return cgroupMemory{}, false, nil +} + +func readCgroupHierarchy( + readFile memoryFileReader, + base string, + mountPoint string, + limitFile string, + usedFile string, + unlimitedValue string, + source string, +) (cgroupMemory, bool, error) { + base = path.Clean(base) + mountPoint = path.Clean(mountPoint) + if base != mountPoint && !strings.HasPrefix(base, mountPoint+"/") { + return cgroupMemory{}, false, fmt.Errorf("cgroup path %q is outside mount point %q", base, mountPoint) + } + + var result cgroupMemory + found := false + for current := base; ; current = path.Dir(current) { + limitData, err := readFile(path.Join(current, limitFile)) + if err != nil { + return cgroupMemory{}, false, err + } + if strings.TrimSpace(string(limitData)) != unlimitedValue { + limit, err := parseMemoryValue(limitFile, limitData) + if err != nil { + return cgroupMemory{}, false, err + } + usedData, err := readFile(path.Join(current, usedFile)) + if err != nil { + return cgroupMemory{}, false, err + } + used, err := parseMemoryValue(usedFile, usedData) + if err != nil { + return cgroupMemory{}, false, err + } + available := saturatingSub(limit, used) + if !found { + result = cgroupMemory{limit: limit, available: available, source: source} + found = true + } else { + result.limit = min(result.limit, limit) + result.available = min(result.available, available) + } + } + if current == mountPoint { + break + } + } + return result, found, nil +} + +func parseCgroupMembership(data string) (cgroupMembership, error) { + var membership cgroupMembership + scanner := bufio.NewScanner(strings.NewReader(data)) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" { + continue + } + parts := strings.SplitN(line, ":", 3) + if len(parts) != 3 { + return cgroupMembership{}, fmt.Errorf("invalid cgroup entry %q", line) + } + if parts[0] == "0" && parts[1] == "" { + membership.v2Path = parts[2] + membership.v2Found = true + continue + } + for _, controller := range strings.Split(parts[1], ",") { + if controller == "memory" { + membership.memoryPath = parts[2] + membership.memoryFound = true + break + } + } + } + if err := scanner.Err(); err != nil { + return cgroupMembership{}, err + } + return membership, nil +} + +func parseCgroupMounts(data string) ([]cgroupMount, error) { + var mounts []cgroupMount + scanner := bufio.NewScanner(strings.NewReader(data)) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" { + continue + } + fields := strings.Fields(line) + separator := -1 + for i, field := range fields { + if field == "-" { + separator = i + break + } + } + if len(fields) < 6 || separator < 6 || separator+3 >= len(fields) { + return nil, fmt.Errorf("invalid mountinfo entry %q", line) + } + fsType := fields[separator+1] + if fsType != "cgroup" && fsType != "cgroup2" { + continue + } + mount := cgroupMount{ + root: unescapeMountInfoPath(fields[3]), + mountPoint: unescapeMountInfoPath(fields[4]), + fsType: fsType, + } + if fsType == "cgroup" { + for _, option := range strings.Split(fields[separator+3], ",") { + if option == "memory" { + mount.memory = true + break + } + } + } + mounts = append(mounts, mount) + } + if err := scanner.Err(); err != nil { + return nil, err + } + return mounts, nil +} + +func resolveCgroupPath(mount cgroupMount, membershipPath string) (string, bool) { + root := path.Clean(mount.root) + membershipPath = path.Clean(membershipPath) + var relative string + switch { + case membershipPath == root: + relative = "." + case root == "/": + relative = strings.TrimPrefix(membershipPath, "/") + case membershipPath == "/": + // In a cgroup namespace, / names the root exposed at the mount point. + relative = "." + case strings.HasPrefix(membershipPath, root+"/"): + relative = strings.TrimPrefix(membershipPath, root+"/") + default: + return "", false + } + return path.Join(mount.mountPoint, relative), true +} + +func parseMemoryValue(name string, data []byte) (uint64, error) { + value := strings.TrimSpace(string(data)) + parsed, err := strconv.ParseUint(value, 10, 64) + if err != nil { + return 0, fmt.Errorf("parse %s value %q: %w", name, value, err) + } + return parsed, nil +} + +func unescapeMountInfoPath(value string) string { + replacer := strings.NewReplacer( + `\040`, " ", + `\011`, "\t", + `\012`, "\n", + `\134`, `\`, + ) + return replacer.Replace(value) +} diff --git a/internal/mem/snapshot_cgroup_test.go b/internal/mem/snapshot_cgroup_test.go new file mode 100644 index 0000000000..a7c4d34422 --- /dev/null +++ b/internal/mem/snapshot_cgroup_test.go @@ -0,0 +1,221 @@ +package mem + +import ( + "errors" + "os" + "testing" +) + +const ( + megabyte = uint64(1024 * 1024) + gigabyte = uint64(1024 * 1024 * 1024) +) + +type memoryFixture map[string]string + +func (fixture memoryFixture) readFile(name string) ([]byte, error) { + value, ok := fixture[name] + if !ok { + return nil, os.ErrNotExist + } + return []byte(value), nil +} + +func TestPlatformMemorySnapshotCgroupV2(t *testing.T) { + host := &MemorySnapshot{ + Limit: 128 * gigabyte, + Used: 64 * gigabyte, + Available: 64 * gigabyte, + Source: MemorySourceHost, + } + tests := []struct { + name string + fixture memoryFixture + wantLimit uint64 + wantAvailable uint64 + wantSource string + }{ + { + name: "docker leaf limit", + fixture: memoryFixture{ + procSelfCgroup: "0::/docker/container-id\n", + procSelfMountInfo: "29 23 0:26 / /sys/fs/cgroup rw,nosuid,nodev,noexec,relatime - cgroup2 cgroup rw\n", + "/sys/fs/cgroup/docker/container-id/memory.max": "4294967296\n", + "/sys/fs/cgroup/docker/container-id/memory.current": "1073741824\n", + "/sys/fs/cgroup/docker/memory.max": "max\n", + "/sys/fs/cgroup/memory.max": "max\n", + }, + wantLimit: 4 * gigabyte, + wantAvailable: 3 * gigabyte, + wantSource: MemorySourceCgroupV2, + }, + { + name: "parent limit", + fixture: memoryFixture{ + procSelfCgroup: "0::/docker/container-id\n", + procSelfMountInfo: "29 23 0:26 / /sys/fs/cgroup rw,nosuid,nodev,noexec,relatime - cgroup2 cgroup rw\n", + "/sys/fs/cgroup/docker/container-id/memory.max": "max\n", + "/sys/fs/cgroup/docker/memory.max": "2147483648\n", + "/sys/fs/cgroup/docker/memory.current": "536870912\n", + "/sys/fs/cgroup/memory.max": "max\n", + }, + wantLimit: 2 * gigabyte, + wantAvailable: 1536 * megabyte, + wantSource: MemorySourceCgroupV2, + }, + { + name: "parent has tighter available memory", + fixture: memoryFixture{ + procSelfCgroup: "0::/docker/container-id\n", + procSelfMountInfo: "29 23 0:26 / /sys/fs/cgroup rw,nosuid,nodev,noexec,relatime - cgroup2 cgroup rw\n", + "/sys/fs/cgroup/docker/container-id/memory.max": "4294967296\n", + "/sys/fs/cgroup/docker/container-id/memory.current": "1073741824\n", + "/sys/fs/cgroup/docker/memory.max": "8589934592\n", + "/sys/fs/cgroup/docker/memory.current": "7516192768\n", + "/sys/fs/cgroup/memory.max": "max\n", + }, + wantLimit: 4 * gigabyte, + wantAvailable: gigabyte, + wantSource: MemorySourceCgroupV2, + }, + { + name: "unlimited falls back to host", + fixture: memoryFixture{ + procSelfCgroup: "0::/\n", + procSelfMountInfo: "29 23 0:26 / /sys/fs/cgroup rw,nosuid,nodev,noexec,relatime - cgroup2 cgroup rw\n", + "/sys/fs/cgroup/memory.max": "max\n", + }, + wantLimit: host.Limit, + wantAvailable: host.Available, + wantSource: MemorySourceHost, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + snapshot, err := cgroupAwareMemorySnapshot(host, nil, tt.fixture.readFile) + if err != nil { + t.Fatalf("cgroupAwareMemorySnapshot() error = %v", err) + } + if snapshot.Limit != tt.wantLimit || snapshot.Available != tt.wantAvailable || snapshot.Source != tt.wantSource { + t.Fatalf("snapshot = %+v, want limit=%d available=%d source=%q", snapshot, tt.wantLimit, tt.wantAvailable, tt.wantSource) + } + }) + } +} + +func TestPlatformMemorySnapshotCgroupV1(t *testing.T) { + host := &MemorySnapshot{ + Limit: 64 * gigabyte, + Used: 32 * gigabyte, + Available: 32 * gigabyte, + Source: MemorySourceHost, + } + const unlimited = "9223372036854771712\n" + tests := []struct { + name string + leafLimit string + leafUsage string + wantLimit uint64 + wantAvailable uint64 + wantSource string + }{ + { + name: "finite docker limit", + leafLimit: "2147483648\n", + leafUsage: "536870912\n", + wantLimit: 2 * gigabyte, + wantAvailable: 1536 * megabyte, + wantSource: MemorySourceCgroupV1, + }, + { + name: "kernel unlimited value is capped by host", + leafLimit: unlimited, + leafUsage: "1073741824\n", + wantLimit: host.Limit, + wantAvailable: host.Available, + wantSource: MemorySourceHost, + }, + { + name: "usage above limit saturates available", + leafLimit: "1073741824\n", + leafUsage: "2147483648\n", + wantLimit: gigabyte, + wantAvailable: 0, + wantSource: MemorySourceCgroupV1, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + fixture := memoryFixture{ + procSelfCgroup: "5:cpu,memory:/docker/container-id\n", + procSelfMountInfo: "35 23 0:31 / /sys/fs/cgroup/memory rw,nosuid,nodev,noexec,relatime - cgroup cgroup rw,memory\n", + "/sys/fs/cgroup/memory/docker/container-id/memory.limit_in_bytes": tt.leafLimit, + "/sys/fs/cgroup/memory/docker/container-id/memory.usage_in_bytes": tt.leafUsage, + "/sys/fs/cgroup/memory/docker/memory.limit_in_bytes": unlimited, + "/sys/fs/cgroup/memory/docker/memory.usage_in_bytes": "1073741824\n", + "/sys/fs/cgroup/memory/memory.limit_in_bytes": unlimited, + "/sys/fs/cgroup/memory/memory.usage_in_bytes": "1073741824\n", + } + snapshot, err := cgroupAwareMemorySnapshot(host, nil, fixture.readFile) + if err != nil { + t.Fatalf("cgroupAwareMemorySnapshot() error = %v", err) + } + if snapshot.Limit != tt.wantLimit || snapshot.Available != tt.wantAvailable || snapshot.Source != tt.wantSource { + t.Fatalf("snapshot = %+v, want limit=%d available=%d source=%q", snapshot, tt.wantLimit, tt.wantAvailable, tt.wantSource) + } + }) + } +} + +func TestPlatformMemorySnapshotResolvesMountRoot(t *testing.T) { + host := &MemorySnapshot{Limit: 16 * gigabyte, Available: 8 * gigabyte, Source: MemorySourceHost} + fixture := memoryFixture{ + procSelfCgroup: "0::/tenant/workload\n", + procSelfMountInfo: "29 23 0:26 /tenant /sys/fs/cgroup rw,nosuid,nodev,noexec,relatime - cgroup2 cgroup rw\n", + "/sys/fs/cgroup/workload/memory.max": "1073741824\n", + "/sys/fs/cgroup/workload/memory.current": "268435456\n", + "/sys/fs/cgroup/memory.max": "max\n", + } + snapshot, err := cgroupAwareMemorySnapshot(host, nil, fixture.readFile) + if err != nil { + t.Fatalf("cgroupAwareMemorySnapshot() error = %v", err) + } + if snapshot.Limit != gigabyte || snapshot.Available != 768*megabyte || snapshot.Source != MemorySourceCgroupV2 { + t.Fatalf("snapshot = %+v", snapshot) + } +} + +func TestPlatformMemorySnapshotFailsClosedOnCgroupReadError(t *testing.T) { + host := &MemorySnapshot{Limit: 16 * gigabyte, Available: 8 * gigabyte, Source: MemorySourceHost} + fixture := memoryFixture{ + procSelfCgroup: "0::/docker/container-id\n", + procSelfMountInfo: "29 23 0:26 / /sys/fs/cgroup rw,nosuid,nodev,noexec,relatime - cgroup2 cgroup rw\n", + } + snapshot, err := cgroupAwareMemorySnapshot(host, nil, fixture.readFile) + if err == nil { + t.Fatal("cgroupAwareMemorySnapshot() expected an error") + } + if snapshot != (MemorySnapshot{}) { + t.Fatalf("snapshot = %+v, want zero value", snapshot) + } + if !errors.Is(err, os.ErrNotExist) { + t.Fatalf("error = %v, want os.ErrNotExist", err) + } +} + +func TestPlatformMemorySnapshotFailsClosedWithoutCgroupMount(t *testing.T) { + host := &MemorySnapshot{Limit: 16 * gigabyte, Available: 8 * gigabyte, Source: MemorySourceHost} + fixture := memoryFixture{ + procSelfCgroup: "0::/docker/container-id\n", + procSelfMountInfo: "29 23 0:26 / /proc rw,nosuid,nodev,noexec,relatime - proc proc rw\n", + } + snapshot, err := cgroupAwareMemorySnapshot(host, nil, fixture.readFile) + if err == nil { + t.Fatal("cgroupAwareMemorySnapshot() expected an error") + } + if snapshot != (MemorySnapshot{}) { + t.Fatalf("snapshot = %+v, want zero value", snapshot) + } +} diff --git a/internal/mem/snapshot_linux.go b/internal/mem/snapshot_linux.go new file mode 100644 index 0000000000..e0d5c7b6db --- /dev/null +++ b/internal/mem/snapshot_linux.go @@ -0,0 +1,7 @@ +//go:build linux + +package mem + +func platformMemorySnapshot(host *MemorySnapshot, hostErr error, readFile memoryFileReader) (MemorySnapshot, error) { + return cgroupAwareMemorySnapshot(host, hostErr, readFile) +} diff --git a/internal/mem/snapshot_other.go b/internal/mem/snapshot_other.go new file mode 100644 index 0000000000..94ddad873b --- /dev/null +++ b/internal/mem/snapshot_other.go @@ -0,0 +1,15 @@ +//go:build !linux + +package mem + +import "errors" + +func platformMemorySnapshot(host *MemorySnapshot, hostErr error, _ memoryFileReader) (MemorySnapshot, error) { + if host != nil { + return *host, nil + } + if hostErr != nil { + return MemorySnapshot{}, hostErr + } + return MemorySnapshot{}, errors.New("memory information is unavailable") +} diff --git a/internal/mem/type.go b/internal/mem/type.go index 3ba8e35ca0..af4ba83488 100644 --- a/internal/mem/type.go +++ b/internal/mem/type.go @@ -5,5 +5,3 @@ type LinearMemory interface { Reallocate(size uint64) (all []byte, err error) Free() error } - -type GrowCheck func(growSize uint64) error diff --git a/internal/mem/utils.go b/internal/mem/utils.go index ef82f99589..d84892cfab 100644 --- a/internal/mem/utils.go +++ b/internal/mem/utils.go @@ -4,50 +4,12 @@ import ( "errors" "fmt" "runtime" - "sync/atomic" - - "github.com/OpenListTeam/OpenList/v4/internal/conf" - "github.com/OpenListTeam/OpenList/v4/pkg/singleflight" - "github.com/shirou/gopsutil/v4/mem" ) var ErrNotEnoughMemory = errors.New("not enough memory") -func MemoryGrowCheck(growSize uint64) error { - if conf.MinFreeMemory == 0 { - return ErrNotEnoughMemory - } - r, err, _ := singleflight.AnyGroup.Do("MemoryGrowCheck", func() (any, error) { - m, err := mem.VirtualMemory() - if err != nil { - return nil, err - } - if m.Available < conf.MinFreeMemory { - return nil, ErrNotEnoughMemory - } - var res atomic.Uint64 - res.Store(m.Available) - return &res, nil - }) - if err != nil { - return err - } - res := r.(*atomic.Uint64) - for { - available := res.Load() - if available < growSize || available-growSize < conf.MinFreeMemory { - return ErrNotEnoughMemory - } - if res.CompareAndSwap(available, available-growSize) { - return nil - } - } -} - -func NewGuardedMemory(cap, max uint64) (m LinearMemory, err error) { - if err := MemoryGrowCheck(cap); err != nil { - return nil, err - } +// NewManagedMemory creates memory with panic recovery and lifecycle cleanup. +func NewManagedMemory(cap, max uint64) (m LinearMemory, err error) { defer func() { if r := recover(); r != nil { err = fmt.Errorf("%w: %v", ErrNotEnoughMemory, r) @@ -57,9 +19,6 @@ func NewGuardedMemory(cap, max uint64) (m LinearMemory, err error) { if err != nil { return nil, err } - if s, ok := m.(interface{ SetGrowCheck(GrowCheck) }); ok { - s.SetGrowCheck(MemoryGrowCheck) - } gm := &guardedMemory{LinearMemory: m} gm.cleanup = runtime.AddCleanup(gm, func(m LinearMemory) { m.Free() diff --git a/internal/net/request.go b/internal/net/request.go index 8c5794e204..1ae6cfe511 100644 --- a/internal/net/request.go +++ b/internal/net/request.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "io" + "math" "math/rand/v2" "net/http" stdpath "path" @@ -13,6 +14,7 @@ import ( "sync/atomic" "time" + "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" hcache "github.com/OpenListTeam/OpenList/v4/internal/hybrid_cache" @@ -40,6 +42,8 @@ var DefaultConcurrencyLimit *ConcurrencyLimit type Downloader struct { PartSize int + // CachePolicy overrides the global cache policy. The zero value inherits it. + CachePolicy cache.Policy // PartBodyMaxRetries is the number of retry attempts to make for failed part downloads. PartBodyMaxRetries int @@ -92,9 +96,20 @@ func (d Downloader) Download(ctx context.Context, p *HttpRequestParams) (readClo if conf.MinFreeMemory > 0 && impl.cfg.PartSize > int(conf.MaxBlockLimit) { impl.cfg.PartSize = int(conf.MaxBlockLimit) } + if impl.cfg.Concurrency <= 0 { + return nil, fmt.Errorf("download concurrency must be positive") + } + if impl.cfg.PartSize <= 0 { + return nil, fmt.Errorf("download part size must be positive") + } if impl.cfg.HttpClient == nil { impl.cfg.HttpClient = DefaultHttpRequestFunc } + policy, err := cache.ResolvePolicy(impl.cfg.CachePolicy, conf.CachePolicy) + if err != nil { + return nil, err + } + impl.cfg.CachePolicy = policy return impl.download() } @@ -197,8 +212,9 @@ func (d *downloader) download() (io.ReadCloser, error) { d.maxPos = d.params.Range.Start + d.params.Range.Length d.concurrency = d.cfg.Concurrency + memoryCeiling := downloaderMemoryCeiling(d.params.Range.Length, d.cfg.Concurrency, d.cfg.PartSize) var err error - d.hc, err = hcache.NewHybridCache(uint64(d.cfg.PartSize), uint64(d.params.Range.Length)) + d.hc, err = hcache.NewHybridCache(uint64(d.cfg.PartSize), memoryCeiling, d.cfg.CachePolicy) if err == nil { d.bufMap = make(map[int]*buffer.PipeBuffer, d.cfg.Concurrency) err = d.sendChunkTask(true) @@ -215,6 +231,19 @@ func (d *downloader) download() (io.ReadCloser, error) { return &multiReadCloser{d: d, curBuf: d.popBuf(0), maxPos: maxPart}, nil } +func downloaderMemoryCeiling(rangeLength int64, concurrency, partSize int) int64 { + if rangeLength <= 0 || concurrency <= 0 || partSize <= 0 { + return 0 + } + partSize64 := int64(partSize) + parts := 1 + (rangeLength-1)/partSize64 + activeBlocks := min(parts, int64(concurrency)) + if activeBlocks > math.MaxInt64/partSize64 { + return math.MaxInt64 + } + return activeBlocks * partSize64 +} + func (d *downloader) sendChunkTask(newConcurrency bool) (err error) { d.mu.Lock() defer d.mu.Unlock() @@ -334,7 +363,7 @@ func (d *downloader) popBuf(id int) *buffer.PipeBuffer { return br } -func (d *downloader) finishBuf(nextId int, prev *buffer.PipeBuffer) (next *buffer.PipeBuffer) { +func (d *downloader) finishBuf(nextId int, prev *buffer.PipeBuffer) (next *buffer.PipeBuffer, err error) { d.readingID.Store(int64(nextId)) d.mu.Lock() @@ -345,14 +374,17 @@ func (d *downloader) finishBuf(nextId int, prev *buffer.PipeBuffer) (next *buffe d.mu.Unlock() if shouldSendTask { - _ = d.sendChunkTask(false) + if err := d.sendChunkTask(false); err != nil { + d.cancel(err) + return nil, err + } } else { _ = prev.Close() } d.mu.Lock() defer d.mu.Unlock() - return d.popBuf(nextId) + return d.popBuf(nextId), nil } // downloadPart is an individual goroutine worker reading from the ch channel @@ -499,7 +531,9 @@ func (d *downloader) tryDownloadChunk(params *HttpRequestParams, ch *chunk) (int return 0, err } } - _ = d.sendChunkTask(true) + if err := d.sendChunkTask(true); err != nil && !errors.Is(err, ErrExceedMaxConcurrency) { + return 0, err + } n, err := utils.CopyWithBuffer(ch.buf, resp.Body) if err != nil { @@ -634,8 +668,8 @@ func (mr *multiReadCloser) Read(p []byte) (n int, err error) { if mr.pos >= mr.maxPos { return n, io.EOF } - mr.curBuf = mr.d.finishBuf(mr.pos, mr.curBuf) - return n, nil + mr.curBuf, err = mr.d.finishBuf(mr.pos, mr.curBuf) + return n, err } return n, err } diff --git a/internal/net/request_test.go b/internal/net/request_test.go index 0fdc56eb33..bd6f2e92c0 100644 --- a/internal/net/request_test.go +++ b/internal/net/request_test.go @@ -6,13 +6,21 @@ package net import ( "bytes" "context" + "errors" "fmt" "io" + "math" "net/http" + "strconv" "sync" "testing" "time" + "github.com/OpenListTeam/OpenList/v4/internal/cache" + "github.com/OpenListTeam/OpenList/v4/internal/conf" + hcache "github.com/OpenListTeam/OpenList/v4/internal/hybrid_cache" + "github.com/OpenListTeam/OpenList/v4/internal/mem" + "github.com/OpenListTeam/OpenList/v4/pkg/buffer" "github.com/OpenListTeam/OpenList/v4/pkg/http_range" "github.com/sirupsen/logrus" ) @@ -26,13 +34,104 @@ func containsString(slice []string, val string) bool { return false } +func withDownloaderConfig(t *testing.T) { + t.Helper() + previousConf := conf.Conf + previousPolicy := conf.CachePolicy + conf.Conf = &conf.Config{TempDir: t.TempDir()} + conf.CachePolicy = cache.PolicyMemory + t.Cleanup(func() { + conf.Conf = previousConf + conf.CachePolicy = previousPolicy + }) +} + +func TestDownloaderMemoryCeiling(t *testing.T) { + tests := []struct { + name string + rangeLength int64 + concurrency int + partSize int + want int64 + }{ + {"working set", 100 << 20, 2, 8 << 20, 16 << 20}, + {"partial final block", 6, 2, 4, 8}, + {"range smaller than pool", 10, 4, 8, 16}, + {"exact block", 8, 2, 4, 8}, + {"multiplication overflow", 100, int(^uint(0) >> 1), 2, 100}, + {"negative range", -1, 2, 4, 0}, + {"invalid concurrency", 8, 0, 4, 0}, + {"invalid part size", 8, 2, 0, 0}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := downloaderMemoryCeiling(tt.rangeLength, tt.concurrency, tt.partSize); got != tt.want { + t.Fatalf("downloaderMemoryCeiling() = %d, want %d", got, tt.want) + } + }) + } + if strconv.IntSize == 64 { + maxInt := int(^uint(0) >> 1) + if got := downloaderMemoryCeiling(math.MaxInt64, maxInt, 2); got != math.MaxInt64 { + t.Fatalf("downloaderMemoryCeiling() overflow result = %d, want MaxInt64", got) + } + } +} + +func TestTryDownloadChunkReturnsPrefetchAllocationError(t *testing.T) { + withDownloaderConfig(t) + hc, err := hcache.NewHybridCache(4, 4, cache.PolicyMemory) + if err != nil { + t.Fatalf("NewHybridCache() error = %v", err) + } + t.Cleanup(func() { _ = hc.Close() }) + block, err := hc.NextBlock() + if err != nil { + t.Fatalf("NextBlock() error = %v", err) + } + + ctx, cancel := context.WithCancelCause(context.Background()) + d := &downloader{ + ctx: ctx, + cancel: cancel, + cfg: Downloader{ + PartSize: 4, + Concurrency: 2, + HttpClient: func(context.Context, *HttpRequestParams) (*http.Response, error) { + return &http.Response{ + Body: io.NopCloser(bytes.NewReader([]byte("1234"))), + Header: http.Header{"Content-Range": {"bytes 0-3/8"}}, + ContentLength: 4, + }, nil + }, + }, + params: &HttpRequestParams{ + Range: http_range.Range{Length: 8}, + Size: 8, + }, + chunkCh: make(chan chunk, 2), + bufMap: make(map[int]*buffer.PipeBuffer), + concurrency: 1, + pos: 4, + maxPos: 8, + nextChunk: 1, + hc: hc, + } + current := &chunk{start: 0, size: 4, id: 0, buf: buffer.NewPipeBuffer(ctx, block)} + if _, err := d.tryDownloadChunk(d.getParamsFromChunk(current), current); !errors.Is(err, mem.ErrNotEnoughMemory) { + t.Fatalf("tryDownloadChunk() error = %v, want ErrNotEnoughMemory", err) + } +} + func TestDownloadOrder(t *testing.T) { + withDownloaderConfig(t) buff := []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15} downloader, invocations, ranges := newDownloadRangeClient(buff) con, partSize := 3, 3 d := NewDownloader(func(d *Downloader) { d.Concurrency = con d.PartSize = partSize + d.CachePolicy = cache.PolicyMemory d.HttpClient = downloader.HttpRequest }) @@ -77,6 +176,7 @@ func TestDownloadOrder(t *testing.T) { } func TestDownloadInterrupt(t *testing.T) { + withDownloaderConfig(t) buff := []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15} buff = append(buff, buff...) downloader, _, _ := newDownloadRangeClient(buff) @@ -112,6 +212,7 @@ func TestDownloadInterrupt(t *testing.T) { } func TestHighConcurrency(t *testing.T) { + withDownloaderConfig(t) buff := make([]byte, 8<<10) for i := range len(buff) { buff[i] = byte(i % 256) @@ -122,6 +223,7 @@ func TestHighConcurrency(t *testing.T) { d := NewDownloader(func(d *Downloader) { d.Concurrency = con d.PartSize = partSize + d.CachePolicy = cache.PolicyMemory d.HttpClient = downloader.HttpRequest d.ConcurrencyLimit = &ConcurrencyLimit{ Limit: concurrencyLimit, @@ -176,6 +278,7 @@ func init() { } func TestDownloadSingle(t *testing.T) { + withDownloaderConfig(t) buff := []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15} downloader, invocations, ranges := newDownloadRangeClient(buff) con, partSize := 1, 4 diff --git a/internal/stream/policy_test.go b/internal/stream/policy_test.go new file mode 100644 index 0000000000..d0b0630b75 --- /dev/null +++ b/internal/stream/policy_test.go @@ -0,0 +1,28 @@ +package stream + +import ( + "testing" + + "github.com/OpenListTeam/OpenList/v4/internal/cache" + "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/model" +) + +type fileStreamerWithoutCachePolicy struct { + model.FileStreamer +} + +func TestGetCachePolicyOptionalInterface(t *testing.T) { + previous := conf.CachePolicy + conf.CachePolicy = cache.PolicyDisk + t.Cleanup(func() { conf.CachePolicy = previous }) + + file := fileStreamerWithoutCachePolicy{FileStreamer: &FileStream{}} + policy, err := getCachePolicy(file) + if err != nil { + t.Fatalf("getCachePolicy() error = %v", err) + } + if policy != cache.PolicyDisk { + t.Fatalf("getCachePolicy() = %q, want disk", policy) + } +} diff --git a/internal/stream/stream.go b/internal/stream/stream.go index 09db2ea226..11305001ea 100644 --- a/internal/stream/stream.go +++ b/internal/stream/stream.go @@ -1,7 +1,6 @@ package stream import ( - "bytes" "context" "errors" "fmt" @@ -11,6 +10,7 @@ import ( "sort" "sync" + "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/internal/conf" hcache "github.com/OpenListTeam/OpenList/v4/internal/hybrid_cache" "github.com/OpenListTeam/OpenList/v4/internal/model" @@ -29,14 +29,17 @@ type FileStream struct { ForceStreamUpload bool Exist model.Obj //the file existed in the destination, we can reuse some info since we wil overwrite it utils.Closers - size int64 - oriReader io.Reader // the original reader, used for caching - hc *hcache.HybridCache - peek buffer.SizedReadAtSeeker + size int64 + sizeSet bool + cachePolicy cache.Policy + cachePolicyLocked bool + oriReader io.Reader // the original reader, used for caching + hc *hcache.HybridCache + peek buffer.SizedReadAtSeeker } func (f *FileStream) GetSize() int64 { - if f.size > 0 { + if f.sizeSet { return f.size } return f.Obj.GetSize() @@ -61,6 +64,33 @@ func (f *FileStream) SetExist(obj model.Obj) { f.Exist = obj } +// GetCachePolicy returns the stream override, or the global policy when the +// stream inherits it. +func (f *FileStream) GetCachePolicy() (cache.Policy, error) { + if f.cachePolicyLocked { + return f.cachePolicy, nil + } + return cache.ResolvePolicy(f.cachePolicy, conf.CachePolicy) +} + +func (f *FileStream) lockCachePolicy(policy cache.Policy) { + f.cachePolicy = policy + f.cachePolicyLocked = true +} + +// SetCachePolicy overrides the global policy before this stream initializes a +// cache. PolicyInherit clears the override. +func (f *FileStream) SetCachePolicy(policy cache.Policy) error { + if policy != cache.PolicyInherit && !policy.IsConcrete() { + return fmt.Errorf("invalid cache policy %q", policy) + } + if f.cachePolicyLocked { + return errors.New("cache policy cannot be changed after cache initialization") + } + f.cachePolicy = policy + return nil +} + // CacheFullAndWriter save all data into tmpFile or memory. // It's not thread-safe! func (f *FileStream) CacheFullAndWriter(up *model.UpdateProgress, writer io.Writer) (model.File, error) { @@ -110,36 +140,7 @@ func (f *FileStream) CacheFullAndWriter(up *model.UpdateProgress, writer io.Writ reader = io.TeeReader(reader, writer) } - // 如果文件大小未知,直接缓存到磁盘 - if f.GetSize() < 0 { - // 检查是否有数据 - buf := []byte{0} - n, err := io.ReadFull(reader, buf) - br := bytes.NewReader(buf[:n]) - if err == io.ErrUnexpectedEOF || err == io.EOF { - f.size = br.Size() - f.Reader = br - return br, nil - } else if err != nil { - return nil, err - } - tmpF, err := utils.CreateTempFile(io.MultiReader(br, reader), 0) - if err != nil { - return nil, err - } - f.Add(utils.CloseFunc(func() error { - return errors.Join(tmpF.Close(), os.RemoveAll(tmpF.Name())) - })) - stat, err := tmpF.Stat() - if err != nil { - return nil, err - } - f.size = stat.Size() - f.Reader = tmpF - return tmpF, nil - } - - if up != nil { + if up != nil && f.GetSize() >= 0 { cacheProgress := model.UpdateProgressWithRange(*up, 0, 50) *up = model.UpdateProgressWithRange(*up, 50, 100) size := f.GetSize() @@ -198,17 +199,38 @@ func (f *FileStream) RangeRead(httpRange http_range.Range) (io.Reader, error) { // 确保指定大小的数据被缓存 func (f *FileStream) ensureCache(size int64) (model.File, error) { if f.peek == nil { - blockSize := min(size, f.GetSize(), int64(conf.MaxBlockLimit)) - var err error - f.hc, err = hcache.NewHybridCache(uint64(blockSize), uint64(f.GetSize())) + memoryCeiling := f.GetSize() + blockSize := int64(conf.MaxBlockLimit) + if memoryCeiling >= 0 { + blockSize = min(memoryCeiling, int64(conf.MaxBlockLimit)) + if size > 0 { + blockSize = min(blockSize, size) + } + } + policy, err := f.GetCachePolicy() if err != nil { return nil, err } + f.hc, err = hcache.NewHybridCache(uint64(blockSize), memoryCeiling, policy) + if err != nil { + return nil, err + } + f.lockCachePolicy(policy) f.peek = buffer.NewDynamicReadAtSeeker(f.hc) f.oriReader = f.Reader f.Reader = io.MultiReader(f.peek, f.oriReader) f.Add(f.hc) } + if size < 0 { + _, err := utils.CopyWithBuffer(f.hc, f.oriReader) + if err != nil { + return nil, err + } + f.size = f.peek.Size() + f.sizeSet = true + f.Reader = f.peek + return f.peek, nil + } size = size - f.peek.Size() if size <= 0 { return f.peek, nil @@ -265,6 +287,7 @@ func NewSeekableStream(fs *FileStream, link *model.Link) (*SeekableStream, error fs.Add(rc) } fs.size = size + fs.sizeSet = true fs.Add(link) return &SeekableStream{FileStream: fs, rangeReader: rr}, nil } diff --git a/internal/stream/stream_test.go b/internal/stream/stream_test.go index 1d8d002e2d..b6d56c5c30 100644 --- a/internal/stream/stream_test.go +++ b/internal/stream/stream_test.go @@ -5,16 +5,222 @@ import ( "errors" "fmt" "io" + "os" "testing" + "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/internal/conf" + "github.com/OpenListTeam/OpenList/v4/internal/mem" "github.com/OpenListTeam/OpenList/v4/internal/model" "github.com/OpenListTeam/OpenList/v4/internal/stream" "github.com/OpenListTeam/OpenList/v4/pkg/http_range" "github.com/OpenListTeam/OpenList/v4/pkg/utils" ) +func withStreamCacheConfig(t *testing.T) { + t.Helper() + oldConf := conf.Conf + oldPolicy := conf.CachePolicy + oldBlockLimit := conf.MaxBlockLimit + oldAutoMemoryLimit := conf.AutoMemoryLimit + oldBudgetCapacity := mem.CacheMemoryBudget.Capacity() + oldBudgetReserved := mem.CacheMemoryBudget.Reserved() + t.Cleanup(func() { + conf.Conf = oldConf + conf.CachePolicy = oldPolicy + conf.MaxBlockLimit = oldBlockLimit + conf.AutoMemoryLimit = oldAutoMemoryLimit + if reserved := mem.CacheMemoryBudget.Reserved(); reserved != oldBudgetReserved { + t.Errorf("cache memory reserved after test = %d, want %d", reserved, oldBudgetReserved) + } + mem.CacheMemoryBudget.SetCapacity(oldBudgetCapacity) + }) + conf.Conf = &conf.Config{TempDir: t.TempDir()} + conf.CachePolicy = cache.PolicyMemory + conf.MaxBlockLimit = 16 << 20 + conf.AutoMemoryLimit = 4 << 20 + mem.CacheMemoryBudget.SetCapacity(0) +} + +func TestFileStreamCachePolicy(t *testing.T) { + withStreamCacheConfig(t) + conf.MaxBlockLimit = 4 + conf.AutoMemoryLimit = 0 + + t.Run("inherit and override", func(t *testing.T) { + withStreamCacheConfig(t) + conf.CachePolicy = cache.PolicyDisk + f := &stream.FileStream{} + got, err := f.GetCachePolicy() + if err != nil || got != cache.PolicyDisk { + t.Fatalf("GetCachePolicy() = %q, want disk", got) + } + if err := f.SetCachePolicy(cache.PolicyMemory); err != nil { + t.Fatalf("SetCachePolicy() error = %v", err) + } + got, err = f.GetCachePolicy() + if err != nil || got != cache.PolicyMemory { + t.Fatalf("GetCachePolicy() = %q, want memory", got) + } + if err := f.SetCachePolicy(cache.PolicyInherit); err != nil { + t.Fatalf("SetCachePolicy(inherit) error = %v", err) + } + got, err = f.GetCachePolicy() + if err != nil || got != cache.PolicyDisk { + t.Fatalf("GetCachePolicy() = %q after inherit, want disk", got) + } + conf.CachePolicy = cache.Policy("invalid") + if _, err := f.GetCachePolicy(); err == nil { + t.Fatal("GetCachePolicy() expected an error for an invalid global policy") + } + }) + + for _, tt := range []struct { + policy cache.Policy + wantFile bool + }{ + {cache.PolicyAuto, true}, + {cache.PolicyDisk, true}, + {cache.PolicyMemory, false}, + } { + t.Run(string(tt.policy)+" unknown size", func(t *testing.T) { + withStreamCacheConfig(t) + conf.MaxBlockLimit = 4 + conf.AutoMemoryLimit = 0 + tempDir := t.TempDir() + conf.Conf = &conf.Config{TempDir: tempDir} + conf.CachePolicy = cache.PolicyAuto + input := []byte("unknown-size-stream") + f := &stream.FileStream{ + Obj: &model.Object{Size: -1}, + Reader: io.NopCloser(bytes.NewReader(input)), + } + if err := f.SetCachePolicy(tt.policy); err != nil { + t.Fatalf("SetCachePolicy() error = %v", err) + } + cached, err := f.CacheFullAndWriter(nil, nil) + if err != nil { + t.Fatalf("CacheFullAndWriter() error = %v", err) + } + if f.GetSize() != int64(len(input)) { + t.Fatalf("GetSize() = %d, want %d", f.GetSize(), len(input)) + } + got, err := io.ReadAll(cached) + if err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if !bytes.Equal(got, input) { + t.Fatalf("cached content = %q, want %q", got, input) + } + entries, err := os.ReadDir(tempDir) + if err != nil { + t.Fatalf("ReadDir() error = %v", err) + } + if gotFile := len(entries) > 0; gotFile != tt.wantFile { + t.Fatalf("temporary file present = %v, want %v", gotFile, tt.wantFile) + } + if err := f.SetCachePolicy(cache.PolicyDisk); err == nil { + t.Fatal("SetCachePolicy() expected an error after cache initialization") + } + if err := f.Close(); err != nil { + t.Fatalf("Close() error = %v", err) + } + entries, err = os.ReadDir(tempDir) + if err != nil { + t.Fatalf("ReadDir() after close error = %v", err) + } + if len(entries) != 0 { + t.Fatalf("temporary files remain after close: %v", entries) + } + }) + } + + for _, tt := range []struct { + name string + policy cache.Policy + wantFile bool + }{ + {name: "memory keeps known stream in memory", policy: cache.PolicyMemory}, + {name: "disk keeps known stream on disk", policy: cache.PolicyDisk, wantFile: true}, + {name: "auto rejected keeps known stream on disk", policy: cache.PolicyAuto, wantFile: true}, + } { + t.Run(tt.name, func(t *testing.T) { + withStreamCacheConfig(t) + conf.MaxBlockLimit = 4 + conf.AutoMemoryLimit = 0 + tempDir := t.TempDir() + conf.Conf = &conf.Config{TempDir: tempDir} + conf.CachePolicy = cache.PolicyAuto + conf.MinFreeMemory = 0 + input := []byte("known-size-stream") + f := &stream.FileStream{ + Obj: &model.Object{Size: int64(len(input))}, + Reader: io.NopCloser(bytes.NewReader(input)), + } + if err := f.SetCachePolicy(tt.policy); err != nil { + t.Fatalf("SetCachePolicy() error = %v", err) + } + cached, err := f.CacheFullAndWriter(nil, nil) + if err != nil { + t.Fatalf("CacheFullAndWriter() error = %v", err) + } + got, err := io.ReadAll(cached) + if err != nil { + t.Fatalf("ReadAll() error = %v", err) + } + if !bytes.Equal(got, input) { + t.Fatalf("cached content = %q, want %q", got, input) + } + entries, err := os.ReadDir(tempDir) + if err != nil { + t.Fatalf("ReadDir() error = %v", err) + } + if gotFile := len(entries) > 0; gotFile != tt.wantFile { + t.Fatalf("temporary file present = %v, want %v", gotFile, tt.wantFile) + } + if err := f.Close(); err != nil { + t.Fatalf("Close() error = %v", err) + } + entries, err = os.ReadDir(tempDir) + if err != nil { + t.Fatalf("ReadDir() after close error = %v", err) + } + if len(entries) != 0 { + t.Fatalf("temporary files remain after close: %v", entries) + } + }) + } + + t.Run("inherited policy is frozen after cache initialization", func(t *testing.T) { + withStreamCacheConfig(t) + conf.MaxBlockLimit = 4 + conf.AutoMemoryLimit = 0 + conf.CachePolicy = cache.PolicyDisk + input := []byte("frozen-policy") + f := &stream.FileStream{ + Obj: &model.Object{Size: int64(len(input))}, + Reader: io.NopCloser(bytes.NewReader(input)), + } + if _, err := f.CacheFullAndWriter(nil, nil); err != nil { + t.Fatalf("CacheFullAndWriter() error = %v", err) + } + t.Cleanup(func() { _ = f.Close() }) + conf.CachePolicy = cache.PolicyMemory + policy, err := f.GetCachePolicy() + if err != nil { + t.Fatalf("GetCachePolicy() error = %v", err) + } + if policy != cache.PolicyDisk { + t.Fatalf("GetCachePolicy() = %q, want frozen disk policy", policy) + } + if err := f.SetCachePolicy(cache.PolicyMemory); err == nil { + t.Fatal("SetCachePolicy() expected an error after cache initialization") + } + }) +} + func TestRangeRead(t *testing.T) { + withStreamCacheConfig(t) type args struct { httpRange http_range.Range } @@ -25,12 +231,7 @@ func TestRangeRead(t *testing.T) { }, Reader: io.NopCloser(bytes.NewReader(buf)), } - prevAutoMemoryLimit := conf.AutoMemoryLimit - prevMaxBlockLimit := conf.MaxBlockLimit - t.Cleanup(func() { - conf.AutoMemoryLimit = prevAutoMemoryLimit - conf.MaxBlockLimit = prevMaxBlockLimit - }) + t.Cleanup(func() { _ = f.Close() }) conf.AutoMemoryLimit = 0 conf.MaxBlockLimit = 15 tests := []struct { @@ -95,6 +296,7 @@ func TestRangeRead(t *testing.T) { } func TestPreHash(t *testing.T) { + withStreamCacheConfig(t) buf := []byte("github.com/OpenListTeam/OpenList") f := &stream.FileStream{ Obj: &model.Object{ @@ -102,12 +304,7 @@ func TestPreHash(t *testing.T) { }, Reader: io.NopCloser(bytes.NewReader(buf)), } - prevAutoMemoryLimit := conf.AutoMemoryLimit - prevMaxBlockLimit := conf.MaxBlockLimit - t.Cleanup(func() { - conf.AutoMemoryLimit = prevAutoMemoryLimit - conf.MaxBlockLimit = prevMaxBlockLimit - }) + t.Cleanup(func() { _ = f.Close() }) conf.AutoMemoryLimit = 0 conf.MaxBlockLimit = 15 @@ -127,6 +324,7 @@ func TestPreHash(t *testing.T) { } func TestStreamSectionReader(t *testing.T) { + withStreamCacheConfig(t) buf := make([]byte, 8<<10) for i := range len(buf) { buf[i] = byte(i % 256) @@ -137,18 +335,10 @@ func TestStreamSectionReader(t *testing.T) { }, Reader: io.NopCloser(bytes.NewReader(buf)), } - prevAutoMemoryLimit := conf.AutoMemoryLimit - prevMaxBlockLimit := conf.MaxBlockLimit - prevConf := conf.Conf - t.Cleanup(func() { - conf.AutoMemoryLimit = prevAutoMemoryLimit - conf.MaxBlockLimit = prevMaxBlockLimit - conf.Conf = prevConf - }) + t.Cleanup(func() { _ = f.Close() }) conf.AutoMemoryLimit = 0 conf.MaxBlockLimit = 2 << 10 partSize := 3 << 10 - conf.Conf = &conf.Config{} ss, err := stream.NewStreamSectionReader(f, partSize, nil) if err != nil { t.Errorf("NewStreamSectionReader() error = %v", err) @@ -174,12 +364,5 @@ func TestStreamSectionReader(t *testing.T) { if !bytes.Equal(buf[i:i+length], b1) { t.Errorf("StreamSectionReader.Read() = %s, want %s", b1, buf[i:i+length]) } - if i == 0 { - prevMinFreeMemory := conf.MinFreeMemory - conf.MinFreeMemory = 0 // 强制使用文件缓存 - t.Cleanup(func() { - conf.MinFreeMemory = prevMinFreeMemory - }) - } } } diff --git a/internal/stream/util.go b/internal/stream/util.go index 2947fcbc2b..ed569e36f9 100644 --- a/internal/stream/util.go +++ b/internal/stream/util.go @@ -9,6 +9,7 @@ import ( "net/http" "sync" + "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/internal/conf" "github.com/OpenListTeam/OpenList/v4/internal/errs" hcache "github.com/OpenListTeam/OpenList/v4/internal/hybrid_cache" @@ -182,16 +183,45 @@ type StreamSectionReader interface { DiscardSection(off int64, length int64) error } +type cachePolicyGetter interface { + GetCachePolicy() (cache.Policy, error) +} + +type cachePolicyLocker interface { + lockCachePolicy(cache.Policy) +} + +func getCachePolicy(file model.FileStreamer) (cache.Policy, error) { + if getter, ok := file.(cachePolicyGetter); ok { + return getter.GetCachePolicy() + } + return cache.ResolvePolicy(cache.PolicyInherit, conf.CachePolicy) +} + func NewStreamSectionReader(file model.FileStreamer, sectionSize int, up *model.UpdateProgress) (StreamSectionReader, error) { if file.GetFile() != nil { return &cachedSectionReader{file.GetFile()}, nil } + if sectionSize <= 0 { + return nil, fmt.Errorf("section size must be positive") + } - blockSize := min(uint64(sectionSize), uint64(file.GetSize()), conf.MaxBlockLimit) - hc, err := hcache.NewHybridCache(blockSize, uint64(file.GetSize())) + fileSize := file.GetSize() + blockSize := min(uint64(sectionSize), conf.MaxBlockLimit) + if fileSize >= 0 { + blockSize = min(blockSize, uint64(fileSize)) + } + policy, err := getCachePolicy(file) + if err != nil { + return nil, err + } + hc, err := hcache.NewHybridCache(blockSize, fileSize, policy) if err != nil { return nil, err } + if locker, ok := file.(cachePolicyLocker); ok { + locker.lockCachePolicy(policy) + } file.Add(hc) return &hybridSectionReader{file: file, hc: hc}, nil }