From ba90e9f452741e91020d41157fdf7dcbc3e94b00 Mon Sep 17 00:00:00 2001 From: vxtls <187420201+vxtls@users.noreply.github.com> Date: Fri, 11 Sep 2026 22:31:44 -0400 Subject: [PATCH 1/5] feat(cache): add configurable HybridCache policies - Add auto, memory, and disk cache policies with global configuration and per-instance overrides. - Select the backing store in auto mode using the complete workload memory ceiling while preserving runtime disk spill. - Enforce memory ceilings independently of page-aligned allocations and prevent strict memory mode from falling back to disk. - Route unknown-size streams through HybridCache and support sequential writes without unnecessary file preallocation. - Bound downloader cache decisions by its concurrent working set and validate cache policy configuration. - Add policy, backing-store selection, stream override, cleanup, and downloader ceiling tests. --- internal/bootstrap/config.go | 2 + internal/cache/policy.go | 62 +++++++++ internal/cache/policy_test.go | 82 ++++++++++++ internal/conf/config.go | 57 ++++---- internal/conf/config_test.go | 21 +++ internal/conf/var.go | 7 +- internal/hybrid_cache/hybrid_cache.go | 179 +++++++++++++++++++++---- internal/hybrid_cache/policy_test.go | 184 ++++++++++++++++++++++++++ internal/mem/utils.go | 18 ++- internal/model/obj.go | 3 + internal/net/request.go | 25 +++- internal/net/request_test.go | 24 ++++ internal/stream/stream.go | 87 ++++++------ internal/stream/stream_test.go | 94 +++++++++++++ internal/stream/util.go | 11 +- 15 files changed, 753 insertions(+), 103 deletions(-) create mode 100644 internal/cache/policy.go create mode 100644 internal/cache/policy_test.go create mode 100644 internal/conf/config_test.go create mode 100644 internal/hybrid_cache/policy_test.go diff --git a/internal/bootstrap/config.go b/internal/bootstrap/config.go index 8304468080..9a5c0d7fbd 100644 --- a/internal/bootstrap/config.go +++ b/internal/bootstrap/config.go @@ -96,6 +96,8 @@ func InitConfig() { if !conf.Conf.Force { confFromEnv() } + conf.CachePolicy = conf.Conf.CachePolicy + log.Infof("cache policy: %s", conf.CachePolicy) if conf.Conf.MaxConcurrency > math.MaxInt32 { net.DefaultConcurrencyLimit = &net.ConcurrencyLimit{Limit: math.MaxInt32} diff --git a/internal/cache/policy.go b/internal/cache/policy.go new file mode 100644 index 0000000000..5a10220e53 --- /dev/null +++ b/internal/cache/policy.go @@ -0,0 +1,62 @@ +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" +) + +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 { + 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..c982a47abe --- /dev/null +++ b/internal/cache/policy_test.go @@ -0,0 +1,82 @@ +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") + } +} 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/hybrid_cache.go b/internal/hybrid_cache/hybrid_cache.go index c69147937e..027e750950 100644 --- a/internal/hybrid_cache/hybrid_cache.go +++ b/internal/hybrid_cache/hybrid_cache.go @@ -2,9 +2,11 @@ package hybrid_cache import ( "errors" + "fmt" "io" "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" @@ -13,12 +15,15 @@ import ( // 线程不安全,单线程使用,或者外部加锁保护 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 } // HybridCache本身是一个大的Block,支持分块成多个小的Block @@ -27,6 +32,9 @@ type HybridCache struct { func (hc *HybridCache) AllocBlock(size uint64) (buffer.Block, error) { retry: if hc.backingStore != nil { + if hc.memoryBacking && hc.exceedsMemoryCeiling(hc.backingOffset, size) { + return nil, mem.ErrNotEnoughMemory + } if err := hc.backingStore.GrowTo(int64(hc.backingOffset + size)); err != nil { return nil, err } @@ -38,12 +46,21 @@ retry: ) 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) } @@ -53,6 +70,9 @@ retry: func (hc *HybridCache) allocWriteAtSeeker(size uint64) (buffer.WriteAtSeeker, error) { retry: if hc.backingStore != nil { + if hc.memoryBacking && hc.exceedsMemoryCeiling(hc.backingOffset, size) { + return nil, mem.ErrNotEnoughMemory + } if err := hc.backingStore.GrowTo(int64(hc.backingOffset + size)); err != nil { return nil, err } @@ -60,12 +80,21 @@ retry: 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) } @@ -95,8 +124,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 } @@ -187,6 +228,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 +276,87 @@ 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 +type memoryCheck func(uint64) error + +func selectPolicy(requested cache.Policy, memoryCeiling int64, check memoryCheck) (cache.Policy, error) { + if !requested.IsConcrete() { + return cache.PolicyInherit, fmt.Errorf("invalid cache policy %q", requested) + } + switch requested { + case cache.PolicyMemory, cache.PolicyDisk: + return requested, nil + case cache.PolicyAuto: + if memoryCeiling < 0 { + return cache.PolicyDisk, nil + } + if memoryCeiling == 0 { + return cache.PolicyMemory, nil + } + if err := check(uint64(memoryCeiling)); err != nil { + return cache.PolicyDisk, nil + } + return cache.PolicyMemory, nil + default: + panic("unreachable") + } +} + +// SelectPolicy resolves auto to a concrete memory or disk policy for a cache +// whose maximum simultaneous memory footprint is memoryCeiling. A negative +// ceiling means that the upper bound is unknown. +func SelectPolicy(requested cache.Policy, memoryCeiling int64) (cache.Policy, error) { + return selectPolicy(requested, memoryCeiling, mem.MemoryGrowCheck) +} + +// NewHybridCache creates a non-thread-safe cache using the requested policy. +func NewHybridCache(blockSize uint64, memoryCeiling int64, requested cache.Policy) (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") } + } - // 策略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} - } + selected, err := SelectPolicy(requested, memoryCeiling) + if err != nil { + return nil, err + } + hc = &HybridCache{blockSize: blockSize, memoryCeiling: memoryCeiling} + 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 } - // 策略3: 文件后备 - if hc == nil { - hc = &HybridCache{blockSize: blockSize} - // 文件 - if err2 := hc.initFileCache(); err2 != nil { - return nil, errors.Join(err, err2) + + if requested == cache.PolicyMemory { + hc.memoryStore, err = mem.NewManagedMemory(blockSize, uint64(memoryCeiling), nil) + if err != nil { + return nil, err } + return hc, nil + } + + hc.memoryStore, err = mem.NewGuardedMemory(blockSize, uint64(memoryCeiling)) + if err == nil { + hc.spillOnMemoryFailure = true + return hc, nil + } + + 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..5ae505d1c9 --- /dev/null +++ b/internal/hybrid_cache/policy_test.go @@ -0,0 +1,184 @@ +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) { + errNoMemory := errors.New("no memory") + tests := []struct { + name string + requested cache.Policy + ceiling int64 + checkErr error + want cache.Policy + checks int + }{ + {"explicit memory", cache.PolicyMemory, -1, errNoMemory, cache.PolicyMemory, 0}, + {"explicit disk", cache.PolicyDisk, 1024, nil, cache.PolicyDisk, 0}, + {"auto unknown", cache.PolicyAuto, -1, nil, cache.PolicyDisk, 0}, + {"auto empty", cache.PolicyAuto, 0, nil, cache.PolicyMemory, 0}, + {"auto admitted", cache.PolicyAuto, 1024, nil, cache.PolicyMemory, 1}, + {"auto rejected", cache.PolicyAuto, 1024, errNoMemory, cache.PolicyDisk, 1}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + checks := 0 + got, err := selectPolicy(tt.requested, tt.ceiling, func(size uint64) error { + checks++ + if size != uint64(tt.ceiling) { + t.Fatalf("memory check size = %d, want %d", size, tt.ceiling) + } + return tt.checkErr + }) + if err != nil { + t.Fatalf("selectPolicy() error = %v", err) + } + if got != tt.want || checks != tt.checks { + t.Fatalf("selectPolicy() = %q with %d checks, want %q with %d", got, checks, tt.want, tt.checks) + } + }) + } + if _, err := selectPolicy(cache.PolicyInherit, 1, func(uint64) error { return nil }); 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 TestHybridCacheAutoRejectsWholeCeiling(t *testing.T) { + withCacheConfig(t, 1024) + hc, err := NewHybridCache(4, 8, cache.PolicyAuto) + 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 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 TestHybridCacheAutoSpillKeepsMemoryPrefix(t *testing.T) { + withCacheConfig(t, 0) + hc := &HybridCache{ + blockSize: 2, + memoryStore: &limitedMemory{buf: make([]byte, 0, 2)}, + spillOnMemoryFailure: true, + memoryCeiling: 2, + } + 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) + } + 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 + oldMinFreeMemory := conf.MinFreeMemory + conf.Conf = &conf.Config{TempDir: t.TempDir()} + conf.AutoMemoryLimit = autoMemoryLimit + conf.MinFreeMemory = 0 + t.Cleanup(func() { + conf.Conf = oldConf + conf.AutoMemoryLimit = oldLimit + conf.MinFreeMemory = oldMinFreeMemory + }) +} + +var _ io.Writer = (*HybridCache)(nil) diff --git a/internal/mem/utils.go b/internal/mem/utils.go index ef82f99589..9f7bd82e90 100644 --- a/internal/mem/utils.go +++ b/internal/mem/utils.go @@ -45,8 +45,16 @@ func MemoryGrowCheck(growSize uint64) error { } func NewGuardedMemory(cap, max uint64) (m LinearMemory, err error) { - if err := MemoryGrowCheck(cap); err != nil { - return nil, err + return NewManagedMemory(cap, max, MemoryGrowCheck) +} + +// NewManagedMemory creates memory with panic recovery and lifecycle cleanup. +// A nil growCheck intentionally permits growth without an availability check. +func NewManagedMemory(cap, max uint64, growCheck GrowCheck) (m LinearMemory, err error) { + if growCheck != nil { + if err := growCheck(cap); err != nil { + return nil, err + } } defer func() { if r := recover(); r != nil { @@ -57,8 +65,10 @@ 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) + if growCheck != nil { + if s, ok := m.(interface{ SetGrowCheck(GrowCheck) }); ok { + s.SetGrowCheck(growCheck) + } } gm := &guardedMemory{LinearMemory: m} gm.cleanup = runtime.AddCleanup(gm, func(m LinearMemory) { diff --git a/internal/model/obj.go b/internal/model/obj.go index 1269b5b797..ef7689d522 100644 --- a/internal/model/obj.go +++ b/internal/model/obj.go @@ -6,6 +6,7 @@ import ( "strings" "time" + "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/pkg/http_range" "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/dlclark/regexp2" @@ -47,6 +48,8 @@ type FileStreamer interface { IsForceStreamUpload() bool GetExist() Obj SetExist(Obj) + GetCachePolicy() cache.Policy + SetCachePolicy(cache.Policy) error // for a non-seekable Stream, RangeRead supports peeking some data, and CacheFullAndWriter still works RangeRead(http_range.Range) (io.Reader, error) // for a non-seekable Stream, if Read is called, this function won't work. diff --git a/internal/net/request.go b/internal/net/request.go index 0cfa7942ea..152bcc138d 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) @@ -214,6 +230,13 @@ 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 int64(concurrency) > math.MaxInt64/int64(partSize) { + return rangeLength + } + return min(rangeLength, int64(concurrency)*int64(partSize)) +} + func (d *downloader) sendChunkTask(newConcurrency bool) (err error) { d.mu.Lock() defer d.mu.Unlock() diff --git a/internal/net/request_test.go b/internal/net/request_test.go index 0fdc56eb33..afd1771a6f 100644 --- a/internal/net/request_test.go +++ b/internal/net/request_test.go @@ -13,6 +13,7 @@ import ( "testing" "time" + "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/pkg/http_range" "github.com/sirupsen/logrus" ) @@ -26,6 +27,27 @@ func containsString(slice []string, val string) bool { return false } +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}, + {"range smaller than pool", 10, 4, 8, 10}, + {"multiplication overflow", 100, int(^uint(0) >> 1), 2, 100}, + } + 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) + } + }) + } +} + func TestDownloadOrder(t *testing.T) { buff := []byte{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15} downloader, invocations, ranges := newDownloadRangeClient(buff) @@ -33,6 +55,7 @@ func TestDownloadOrder(t *testing.T) { d := NewDownloader(func(d *Downloader) { d.Concurrency = con d.PartSize = partSize + d.CachePolicy = cache.PolicyMemory d.HttpClient = downloader.HttpRequest }) @@ -122,6 +145,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, diff --git a/internal/stream/stream.go b/internal/stream/stream.go index b1e6fd5faf..189bd93508 100644 --- a/internal/stream/stream.go +++ b/internal/stream/stream.go @@ -1,15 +1,14 @@ package stream import ( - "bytes" "context" "errors" "fmt" "io" "math" - "os" "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" @@ -28,14 +27,16 @@ 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 + cachePolicyOverride cache.Policy + 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() @@ -60,6 +61,25 @@ func (f *FileStream) SetExist(obj model.Obj) { f.Exist = obj } +func (f *FileStream) GetCachePolicy() cache.Policy { + policy, err := cache.ResolvePolicy(f.cachePolicyOverride, conf.CachePolicy) + if err != nil { + panic(err) + } + return policy +} + +func (f *FileStream) SetCachePolicy(policy cache.Policy) error { + if policy != cache.PolicyInherit && !policy.IsConcrete() { + return fmt.Errorf("invalid cache policy %q", policy) + } + if f.peek != nil { + return errors.New("cache policy cannot be changed after cache initialization") + } + f.cachePolicyOverride = 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) { @@ -109,36 +129,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() @@ -197,9 +188,16 @@ 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)) + memoryCeiling := f.GetSize() + blockSize := int64(conf.MaxBlockLimit) + if memoryCeiling >= 0 { + blockSize = min(memoryCeiling, int64(conf.MaxBlockLimit)) + if size > 0 { + blockSize = min(blockSize, size) + } + } var err error - f.hc, err = hcache.NewHybridCache(uint64(blockSize), uint64(f.GetSize())) + f.hc, err = hcache.NewHybridCache(uint64(blockSize), memoryCeiling, f.GetCachePolicy()) if err != nil { return nil, err } @@ -208,6 +206,16 @@ func (f *FileStream) ensureCache(size int64) (model.File, error) { 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 @@ -264,6 +272,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..22e6b1866f 100644 --- a/internal/stream/stream_test.go +++ b/internal/stream/stream_test.go @@ -5,8 +5,10 @@ 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/model" "github.com/OpenListTeam/OpenList/v4/internal/stream" @@ -14,6 +16,98 @@ import ( "github.com/OpenListTeam/OpenList/v4/pkg/utils" ) +func TestFileStreamCachePolicy(t *testing.T) { + oldConf := conf.Conf + oldPolicy := conf.CachePolicy + oldBlockLimit := conf.MaxBlockLimit + oldAutoMemoryLimit := conf.AutoMemoryLimit + t.Cleanup(func() { + conf.Conf = oldConf + conf.CachePolicy = oldPolicy + conf.MaxBlockLimit = oldBlockLimit + conf.AutoMemoryLimit = oldAutoMemoryLimit + }) + conf.MaxBlockLimit = 4 + conf.AutoMemoryLimit = 0 + + t.Run("inherit and override", func(t *testing.T) { + conf.CachePolicy = cache.PolicyDisk + f := &stream.FileStream{} + if got := f.GetCachePolicy(); got != cache.PolicyDisk { + t.Fatalf("GetCachePolicy() = %q, want disk", got) + } + if err := f.SetCachePolicy(cache.PolicyMemory); err != nil { + t.Fatalf("SetCachePolicy() error = %v", err) + } + if got := f.GetCachePolicy(); 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) + } + if got := f.GetCachePolicy(); got != cache.PolicyDisk { + t.Fatalf("GetCachePolicy() = %q after inherit, want disk", got) + } + }) + + 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) { + 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) + } + }) + } +} + func TestRangeRead(t *testing.T) { type args struct { httpRange http_range.Range diff --git a/internal/stream/util.go b/internal/stream/util.go index 2947fcbc2b..adbfa428bb 100644 --- a/internal/stream/util.go +++ b/internal/stream/util.go @@ -186,9 +186,16 @@ func NewStreamSectionReader(file model.FileStreamer, sectionSize int, up *model. 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)) + } + hc, err := hcache.NewHybridCache(blockSize, fileSize, file.GetCachePolicy()) if err != nil { return nil, err } From 1d85ae68418ffef305dc2c1d140009070599c6b4 Mon Sep 17 00:00:00 2001 From: vxtls <187420201+vxtls@users.noreply.github.com> Date: Sun, 13 Sep 2026 17:22:36 -0400 Subject: [PATCH 2/5] test(cache): verify HybridCache backing policy selection - Add deterministic coverage for memory, disk, auto-admitted, auto-rejected, and unknown-size cache policies. - Verify cached data is written entirely to the selected backing store and remains readable. - Add known-size FileStream integration tests for memory, disk, and rejected auto policies. - Verify disk-backed temporary files are created when expected and removed when streams close. - Inject the memory checker through a private constructor to test auto decisions without changing the public API or production behavior. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --- internal/hybrid_cache/hybrid_cache.go | 6 ++- internal/hybrid_cache/policy_test.go | 68 +++++++++++++++++++++++++++ internal/stream/stream_test.go | 55 ++++++++++++++++++++++ 3 files changed, 128 insertions(+), 1 deletion(-) diff --git a/internal/hybrid_cache/hybrid_cache.go b/internal/hybrid_cache/hybrid_cache.go index 027e750950..fedb84bba7 100644 --- a/internal/hybrid_cache/hybrid_cache.go +++ b/internal/hybrid_cache/hybrid_cache.go @@ -310,6 +310,10 @@ func SelectPolicy(requested cache.Policy, memoryCeiling int64) (cache.Policy, er // 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.MemoryGrowCheck) +} + +func newHybridCache(blockSize uint64, memoryCeiling int64, requested cache.Policy, check memoryCheck) (hc *HybridCache, err error) { if memoryCeiling < 0 && blockSize == 0 { return nil, fmt.Errorf("block size must be positive when memory ceiling is unknown") } @@ -320,7 +324,7 @@ func NewHybridCache(blockSize uint64, memoryCeiling int64, requested cache.Polic } } - selected, err := SelectPolicy(requested, memoryCeiling) + selected, err := selectPolicy(requested, memoryCeiling, check) if err != nil { return nil, err } diff --git a/internal/hybrid_cache/policy_test.go b/internal/hybrid_cache/policy_test.go index 5ae505d1c9..6817188305 100644 --- a/internal/hybrid_cache/policy_test.go +++ b/internal/hybrid_cache/policy_test.go @@ -70,6 +70,74 @@ func TestHybridCacheDiskPolicy(t *testing.T) { } } +func TestHybridCachePolicyBackingSelection(t *testing.T) { + tests := []struct { + name string + requested cache.Policy + ceiling int64 + checkErr error + wantMemory bool + wantCheckCalls int + }{ + {name: "memory uses memory", 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, wantMemory: true, wantCheckCalls: 1}, + {name: "auto rejected uses disk", requested: cache.PolicyAuto, ceiling: 8, checkErr: mem.ErrNotEnoughMemory, wantCheckCalls: 1}, + {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) + checkCalls := 0 + hc, err := newHybridCache(4, tt.ceiling, tt.requested, func(uint64) error { + checkCalls++ + return tt.checkErr + }) + if err != nil { + t.Fatalf("newHybridCache() error = %v", err) + } + if checkCalls != tt.wantCheckCalls { + t.Fatalf("memory check calls = %d, want %d", checkCalls, tt.wantCheckCalls) + } + + 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) diff --git a/internal/stream/stream_test.go b/internal/stream/stream_test.go index 22e6b1866f..553056a6e7 100644 --- a/internal/stream/stream_test.go +++ b/internal/stream/stream_test.go @@ -21,11 +21,13 @@ func TestFileStreamCachePolicy(t *testing.T) { oldPolicy := conf.CachePolicy oldBlockLimit := conf.MaxBlockLimit oldAutoMemoryLimit := conf.AutoMemoryLimit + oldMinFreeMemory := conf.MinFreeMemory t.Cleanup(func() { conf.Conf = oldConf conf.CachePolicy = oldPolicy conf.MaxBlockLimit = oldBlockLimit conf.AutoMemoryLimit = oldAutoMemoryLimit + conf.MinFreeMemory = oldMinFreeMemory }) conf.MaxBlockLimit = 4 conf.AutoMemoryLimit = 0 @@ -106,6 +108,59 @@ func TestFileStreamCachePolicy(t *testing.T) { } }) } + + 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) { + 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) + } + }) + } } func TestRangeRead(t *testing.T) { From 844866bbc1074ba74221b4da6449ac20f64365ca Mon Sep 17 00:00:00 2001 From: vxtls <187420201+vxtls@users.noreply.github.com> Date: Sun, 13 Sep 2026 20:02:40 -0400 Subject: [PATCH 3/5] fix(cache): improve memory detection in cgroups - Detect effective memory limits and usage from cgroup v1 and v2. - Resolve container cgroup paths through proc membership and mount information. - Check parent cgroups and use the most restrictive available memory boundary. - Cap cgroup memory against host limits and use saturating arithmetic for invalid usage states. - Fall back to disk when detected cgroup memory information cannot be read reliably. - Use effective memory limits for cache growth checks and automatic memory configuration. - Add deterministic coverage for Docker limits, nested cgroups, unlimited values, and failure cases. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --- internal/bootstrap/config.go | 15 +- internal/mem/snapshot.go | 74 ++++++++ internal/mem/snapshot_cgroup.go | 262 +++++++++++++++++++++++++++ internal/mem/snapshot_cgroup_test.go | 221 ++++++++++++++++++++++ internal/mem/snapshot_linux.go | 7 + internal/mem/snapshot_other.go | 15 ++ internal/mem/utils.go | 7 +- 7 files changed, 591 insertions(+), 10 deletions(-) create mode 100644 internal/mem/snapshot.go create mode 100644 internal/mem/snapshot_cgroup.go create mode 100644 internal/mem/snapshot_cgroup_test.go create mode 100644 internal/mem/snapshot_linux.go create mode 100644 internal/mem/snapshot_other.go diff --git a/internal/bootstrap/config.go b/internal/bootstrap/config.go index 9a5c0d7fbd..2ba6a3637b 100644 --- a/internal/bootstrap/config.go +++ b/internal/bootstrap/config.go @@ -10,10 +10,10 @@ import ( "github.com/OpenListTeam/OpenList/v4/cmd/flags" "github.com/OpenListTeam/OpenList/v4/drivers/base" "github.com/OpenListTeam/OpenList/v4/internal/conf" + internalmem "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" ) @@ -105,15 +105,18 @@ 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 := internalmem.GetMemorySnapshot() + if memErr != nil { + log.Warnf("memory detection warning: %v", memErr) + } + 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 { 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 @@ -122,7 +125,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 diff --git a/internal/mem/snapshot.go b/internal/mem/snapshot.go new file mode 100644 index 0000000000..54c2bcb443 --- /dev/null +++ b/internal/mem/snapshot.go @@ -0,0 +1,74 @@ +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) + +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 { + return MemorySnapshot{ + Limit: cgroup.limit, + Used: cgroup.limit - min(cgroup.available, cgroup.limit), + Available: min(cgroup.available, cgroup.limit), + 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: 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..f701ed343c --- /dev/null +++ b/internal/mem/snapshot_cgroup.go @@ -0,0 +1,262 @@ +package mem + +import ( + "bufio" + "errors" + "fmt" + "path" + "strconv" + "strings" +) + +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/utils.go b/internal/mem/utils.go index 9f7bd82e90..8e1ffd8725 100644 --- a/internal/mem/utils.go +++ b/internal/mem/utils.go @@ -8,7 +8,6 @@ import ( "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") @@ -18,15 +17,15 @@ func MemoryGrowCheck(growSize uint64) error { return ErrNotEnoughMemory } r, err, _ := singleflight.AnyGroup.Do("MemoryGrowCheck", func() (any, error) { - m, err := mem.VirtualMemory() + snapshot, err := GetMemorySnapshot() if err != nil { return nil, err } - if m.Available < conf.MinFreeMemory { + if snapshot.Available < conf.MinFreeMemory { return nil, ErrNotEnoughMemory } var res atomic.Uint64 - res.Store(m.Available) + res.Store(snapshot.Available) return &res, nil }) if err != nil { From ffa96cd6d30ac550db0e6d206714a97d999313bd Mon Sep 17 00:00:00 2001 From: vxtls <187420201+vxtls@users.noreply.github.com> Date: Mon, 14 Sep 2026 10:36:48 -0400 Subject: [PATCH 4/5] fix(cache): review and fix cache policy integration issues - Align downloader memory ceilings with block-sized allocations and propagate allocation failures. - Return cache policy resolution errors instead of panicking in request paths. - Make per-stream cache policy support optional to preserve FileStreamer compatibility. - Preserve inherit semantics when empty cache policies are serialized and parsed. - Isolate cache-related global configuration across stream tests. - Calculate downloader memory ceilings safely for invalid inputs and integer overflow. --- internal/bootstrap/config.go | 8 ++- internal/cache/policy.go | 8 +++ internal/cache/policy_test.go | 13 ++++ internal/hybrid_cache/buffer.go | 14 +++++ internal/hybrid_cache/hybrid_cache.go | 2 + internal/hybrid_cache/policy_test.go | 15 +++++ internal/model/obj.go | 3 - internal/net/request.go | 29 ++++++--- internal/net/request_test.go | 63 ++++++++++++++++++- internal/stream/policy_test.go | 28 +++++++++ internal/stream/stream.go | 26 +++++--- internal/stream/stream_test.go | 90 +++++++++++++++++---------- internal/stream/util.go | 25 +++++++- 13 files changed, 269 insertions(+), 55 deletions(-) create mode 100644 internal/stream/policy_test.go diff --git a/internal/bootstrap/config.go b/internal/bootstrap/config.go index 2ba6a3637b..e72869c9d0 100644 --- a/internal/bootstrap/config.go +++ b/internal/bootstrap/config.go @@ -9,6 +9,7 @@ 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" internalmem "github.com/OpenListTeam/OpenList/v4/internal/mem" "github.com/OpenListTeam/OpenList/v4/internal/net" @@ -96,7 +97,12 @@ func InitConfig() { if !conf.Conf.Force { confFromEnv() } - conf.CachePolicy = conf.Conf.CachePolicy + 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 { diff --git a/internal/cache/policy.go b/internal/cache/policy.go index 5a10220e53..4d0dc22c54 100644 --- a/internal/cache/policy.go +++ b/internal/cache/policy.go @@ -16,6 +16,8 @@ const ( 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 { @@ -53,6 +55,12 @@ func (p Policy) MarshalText() ([]byte, error) { } 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 diff --git a/internal/cache/policy_test.go b/internal/cache/policy_test.go index c982a47abe..5996e27e52 100644 --- a/internal/cache/policy_test.go +++ b/internal/cache/policy_test.go @@ -79,4 +79,17 @@ func TestPolicyJSON(t *testing.T) { 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/hybrid_cache/buffer.go b/internal/hybrid_cache/buffer.go index 0023996e63..00e543ecf7 100644 --- a/internal/hybrid_cache/buffer.go +++ b/internal/hybrid_cache/buffer.go @@ -3,24 +3,32 @@ package hybrid_cache import ( "fmt" "io" + "sync" ) type BufferStore struct { + mu sync.RWMutex blocks [][]byte size int64 } func (m *BufferStore) Size() int64 { + m.mu.RLock() + defer m.mu.RUnlock() return m.size } // 用于存储不复用的[]byte func (m *BufferStore) Append(buf []byte) { + m.mu.Lock() + defer m.mu.Unlock() m.size += int64(len(buf)) m.blocks = append(m.blocks, buf) } func (m *BufferStore) Close() error { + m.mu.Lock() + defer m.mu.Unlock() if len(m.blocks) > 0 { clear(m.blocks) m.blocks = m.blocks[:0] @@ -30,6 +38,8 @@ func (m *BufferStore) Close() error { } func (m *BufferStore) ReadAt(p []byte, off int64) (int, error) { + m.mu.RLock() + defer m.mu.RUnlock() if len(p) == 0 { return 0, nil } @@ -55,6 +65,8 @@ func (m *BufferStore) ReadAt(p []byte, off int64) (int, error) { } func (m *BufferStore) WriteAt(p []byte, off int64) (int, error) { + m.mu.RLock() + defer m.mu.RUnlock() if len(p) == 0 { return 0, nil } @@ -80,6 +92,8 @@ func (m *BufferStore) WriteAt(p []byte, off int64) (int, error) { } func (m *BufferStore) GrowTo(size int64) (err error) { + m.mu.Lock() + defer m.mu.Unlock() if size <= m.size { return nil } diff --git a/internal/hybrid_cache/hybrid_cache.go b/internal/hybrid_cache/hybrid_cache.go index fedb84bba7..51d3223c26 100644 --- a/internal/hybrid_cache/hybrid_cache.go +++ b/internal/hybrid_cache/hybrid_cache.go @@ -290,6 +290,8 @@ func selectPolicy(requested cache.Policy, memoryCeiling int64, check memoryCheck return cache.PolicyDisk, nil } if memoryCeiling == 0 { + // Zero is a known empty workload, unlike a negative unknown ceiling. + // The hard ceiling still rejects any unexpected writes. return cache.PolicyMemory, nil } if err := check(uint64(memoryCeiling)); err != nil { diff --git a/internal/hybrid_cache/policy_test.go b/internal/hybrid_cache/policy_test.go index 6817188305..db5ee74fea 100644 --- a/internal/hybrid_cache/policy_test.go +++ b/internal/hybrid_cache/policy_test.go @@ -153,6 +153,21 @@ func TestHybridCacheAutoRejectsWholeCeiling(t *testing.T) { } } +func TestHybridCacheZeroCeilingRejectsUnexpectedWrites(t *testing.T) { + withCacheConfig(t, 8) + hc, err := newHybridCache(4, 0, cache.PolicyAuto, func(uint64) error { + t.Fatal("zero ceiling must not check memory") + return nil + }) + 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) diff --git a/internal/model/obj.go b/internal/model/obj.go index ef7689d522..1269b5b797 100644 --- a/internal/model/obj.go +++ b/internal/model/obj.go @@ -6,7 +6,6 @@ import ( "strings" "time" - "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/pkg/http_range" "github.com/OpenListTeam/OpenList/v4/pkg/utils" "github.com/dlclark/regexp2" @@ -48,8 +47,6 @@ type FileStreamer interface { IsForceStreamUpload() bool GetExist() Obj SetExist(Obj) - GetCachePolicy() cache.Policy - SetCachePolicy(cache.Policy) error // for a non-seekable Stream, RangeRead supports peeking some data, and CacheFullAndWriter still works RangeRead(http_range.Range) (io.Reader, error) // for a non-seekable Stream, if Read is called, this function won't work. diff --git a/internal/net/request.go b/internal/net/request.go index 152bcc138d..e5647ac247 100644 --- a/internal/net/request.go +++ b/internal/net/request.go @@ -231,10 +231,16 @@ func (d *downloader) download() (io.ReadCloser, error) { } func downloaderMemoryCeiling(rangeLength int64, concurrency, partSize int) int64 { - if int64(concurrency) > math.MaxInt64/int64(partSize) { - return rangeLength + if rangeLength <= 0 || concurrency <= 0 || partSize <= 0 { + return 0 } - return min(rangeLength, int64(concurrency)*int64(partSize)) + 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) { @@ -355,7 +361,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() @@ -366,14 +372,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 @@ -520,7 +529,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 { @@ -655,8 +666,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 afd1771a6f..6bf9b0071f 100644 --- a/internal/net/request_test.go +++ b/internal/net/request_test.go @@ -6,14 +6,20 @@ package net import ( "bytes" "context" + "errors" "fmt" "io" + "math" "net/http" + "strconv" "sync" "testing" "time" "github.com/OpenListTeam/OpenList/v4/internal/cache" + 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" ) @@ -36,8 +42,13 @@ func TestDownloaderMemoryCeiling(t *testing.T) { want int64 }{ {"working set", 100 << 20, 2, 8 << 20, 16 << 20}, - {"range smaller than pool", 10, 4, 8, 10}, + {"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) { @@ -46,6 +57,56 @@ func TestDownloaderMemoryCeiling(t *testing.T) { } }) } + 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) { + 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) { 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 189bd93508..37df320113 100644 --- a/internal/stream/stream.go +++ b/internal/stream/stream.go @@ -30,6 +30,8 @@ type FileStream struct { size int64 sizeSet bool cachePolicyOverride cache.Policy + cachePolicyResolved cache.Policy + cachePolicyLocked bool oriReader io.Reader // the original reader, used for caching hc *hcache.HybridCache peek buffer.SizedReadAtSeeker @@ -61,19 +63,23 @@ func (f *FileStream) SetExist(obj model.Obj) { f.Exist = obj } -func (f *FileStream) GetCachePolicy() cache.Policy { - policy, err := cache.ResolvePolicy(f.cachePolicyOverride, conf.CachePolicy) - if err != nil { - panic(err) +func (f *FileStream) GetCachePolicy() (cache.Policy, error) { + if f.cachePolicyLocked { + return f.cachePolicyResolved, nil } - return policy + return cache.ResolvePolicy(f.cachePolicyOverride, conf.CachePolicy) +} + +func (f *FileStream) freezeCachePolicy(policy cache.Policy) { + f.cachePolicyResolved = policy + f.cachePolicyLocked = true } func (f *FileStream) SetCachePolicy(policy cache.Policy) error { if policy != cache.PolicyInherit && !policy.IsConcrete() { return fmt.Errorf("invalid cache policy %q", policy) } - if f.peek != nil { + if f.cachePolicyLocked { return errors.New("cache policy cannot be changed after cache initialization") } f.cachePolicyOverride = policy @@ -196,11 +202,15 @@ func (f *FileStream) ensureCache(size int64) (model.File, error) { blockSize = min(blockSize, size) } } - var err error - f.hc, err = hcache.NewHybridCache(uint64(blockSize), memoryCeiling, f.GetCachePolicy()) + 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.freezeCachePolicy(policy) f.peek = buffer.NewDynamicReadAtSeeker(f.hc) f.oriReader = f.Reader f.Reader = io.MultiReader(f.peek, f.oriReader) diff --git a/internal/stream/stream_test.go b/internal/stream/stream_test.go index 553056a6e7..c3e1792152 100644 --- a/internal/stream/stream_test.go +++ b/internal/stream/stream_test.go @@ -16,7 +16,8 @@ import ( "github.com/OpenListTeam/OpenList/v4/pkg/utils" ) -func TestFileStreamCachePolicy(t *testing.T) { +func withStreamCacheConfig(t *testing.T) { + t.Helper() oldConf := conf.Conf oldPolicy := conf.CachePolicy oldBlockLimit := conf.MaxBlockLimit @@ -29,27 +30,44 @@ func TestFileStreamCachePolicy(t *testing.T) { conf.AutoMemoryLimit = oldAutoMemoryLimit conf.MinFreeMemory = oldMinFreeMemory }) + conf.Conf = &conf.Config{TempDir: t.TempDir()} + conf.CachePolicy = cache.PolicyMemory + conf.MaxBlockLimit = 16 << 20 + conf.AutoMemoryLimit = 4 << 20 + conf.MinFreeMemory = 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{} - if got := f.GetCachePolicy(); got != cache.PolicyDisk { + 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) } - if got := f.GetCachePolicy(); got != cache.PolicyMemory { + 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) } - if got := f.GetCachePolicy(); got != cache.PolicyDisk { + 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 { @@ -61,6 +79,9 @@ func TestFileStreamCachePolicy(t *testing.T) { {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 @@ -119,6 +140,9 @@ func TestFileStreamCachePolicy(t *testing.T) { {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 @@ -161,9 +185,37 @@ func TestFileStreamCachePolicy(t *testing.T) { } }) } + + 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 } @@ -174,12 +226,6 @@ 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 - }) conf.AutoMemoryLimit = 0 conf.MaxBlockLimit = 15 tests := []struct { @@ -244,6 +290,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{ @@ -251,12 +298,6 @@ 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 - }) conf.AutoMemoryLimit = 0 conf.MaxBlockLimit = 15 @@ -276,6 +317,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) @@ -286,18 +328,9 @@ 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 - }) 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) @@ -323,12 +356,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 adbfa428bb..13d05f6e3e 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,6 +183,21 @@ type StreamSectionReader interface { DiscardSection(off int64, length int64) error } +type cachePolicyGetter interface { + GetCachePolicy() (cache.Policy, error) +} + +type cachePolicyFreezer interface { + freezeCachePolicy(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 @@ -195,10 +211,17 @@ func NewStreamSectionReader(file model.FileStreamer, sectionSize int, up *model. if fileSize >= 0 { blockSize = min(blockSize, uint64(fileSize)) } - hc, err := hcache.NewHybridCache(blockSize, fileSize, file.GetCachePolicy()) + policy, err := getCachePolicy(file) if err != nil { return nil, err } + hc, err := hcache.NewHybridCache(blockSize, fileSize, policy) + if err != nil { + return nil, err + } + if freezer, ok := file.(cachePolicyFreezer); ok { + freezer.freezeCachePolicy(policy) + } file.Add(hc) return &hybridSectionReader{file: file, hc: hc}, nil } From 6595659c4e1bac3edfdb00c3638bca629463d21d Mon Sep 17 00:00:00 2001 From: vxtls <187420201+vxtls@users.noreply.github.com> Date: Sun, 20 Sep 2026 18:12:58 -0400 Subject: [PATCH 5/5] fix(cache): review and fix cache memory accounting issues - Derive a process-wide cache budget from effective available memory while preserving configured free-memory headroom. - Reserve known-size automatic caches for their full lifetime so concurrent HybridCache instances cannot oversubscribe memory. - Include direct BufferStore allocations in admission decisions, release reservations on close, and retain only the memory prefix after disk spill. - Keep explicit memory policy independent of automatic admission and route unknown-size automatic caches to disk. - Replace transient growth checks with lifecycle accounting and centralize backing-store ceiling and overflow validation. - Simplify per-stream policy state and clarify BufferStore locking behavior. - Add deterministic coverage for budget concurrency, policy selection, bootstrap capacity, stream cleanup, and downloader behavior. Co-authored-by: Codex <267193182+codex@users.noreply.github.com> --- internal/bootstrap/config.go | 16 ++- internal/bootstrap/config_test.go | 25 ++++ internal/hybrid_cache/buffer.go | 36 +++--- internal/hybrid_cache/hybrid_cache.go | 112 ++++++++++-------- internal/hybrid_cache/policy_test.go | 146 ++++++++++++++++------- internal/mem/budget.go | 160 ++++++++++++++++++++++++++ internal/mem/budget_test.go | 106 +++++++++++++++++ internal/mem/mem_unix.go | 13 +-- internal/mem/mem_windows.go | 15 +-- internal/mem/snapshot.go | 9 +- internal/mem/snapshot_cgroup.go | 4 + internal/mem/type.go | 2 - internal/mem/utils.go | 52 +-------- internal/net/request_test.go | 18 +++ internal/stream/stream.go | 31 ++--- internal/stream/stream_test.go | 14 ++- internal/stream/util.go | 8 +- 17 files changed, 555 insertions(+), 212 deletions(-) create mode 100644 internal/bootstrap/config_test.go create mode 100644 internal/mem/budget.go create mode 100644 internal/mem/budget_test.go diff --git a/internal/bootstrap/config.go b/internal/bootstrap/config.go index e72869c9d0..f634b053fa 100644 --- a/internal/bootstrap/config.go +++ b/internal/bootstrap/config.go @@ -11,7 +11,7 @@ import ( "github.com/OpenListTeam/OpenList/v4/drivers/base" "github.com/OpenListTeam/OpenList/v4/internal/cache" "github.com/OpenListTeam/OpenList/v4/internal/conf" - internalmem "github.com/OpenListTeam/OpenList/v4/internal/mem" + "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" @@ -111,16 +111,18 @@ func InitConfig() { net.DefaultConcurrencyLimit = &net.ConcurrencyLimit{Limit: uint32(conf.Conf.MaxConcurrency)} } - memStat, memErr := internalmem.GetMemorySnapshot() + 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.Limit >> 20) / 10 conf.MinFreeMemory = max(16, min(t, 1024)) << 20 @@ -141,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 @@ -180,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/hybrid_cache/buffer.go b/internal/hybrid_cache/buffer.go index 00e543ecf7..1bd93c775f 100644 --- a/internal/hybrid_cache/buffer.go +++ b/internal/hybrid_cache/buffer.go @@ -6,29 +6,33 @@ import ( "sync" ) +// BufferStore is a growable in-memory backing store split into stable blocks. type BufferStore struct { - mu sync.RWMutex - 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.mu.RLock() - defer m.mu.RUnlock() + 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.mu.Lock() - defer m.mu.Unlock() + m.layoutMu.Lock() + defer m.layoutMu.Unlock() m.size += int64(len(buf)) m.blocks = append(m.blocks, buf) } func (m *BufferStore) Close() error { - m.mu.Lock() - defer m.mu.Unlock() + m.layoutMu.Lock() + defer m.layoutMu.Unlock() if len(m.blocks) > 0 { clear(m.blocks) m.blocks = m.blocks[:0] @@ -38,8 +42,8 @@ func (m *BufferStore) Close() error { } func (m *BufferStore) ReadAt(p []byte, off int64) (int, error) { - m.mu.RLock() - defer m.mu.RUnlock() + m.layoutMu.RLock() + defer m.layoutMu.RUnlock() if len(p) == 0 { return 0, nil } @@ -65,8 +69,8 @@ func (m *BufferStore) ReadAt(p []byte, off int64) (int, error) { } func (m *BufferStore) WriteAt(p []byte, off int64) (int, error) { - m.mu.RLock() - defer m.mu.RUnlock() + m.layoutMu.RLock() + defer m.layoutMu.RUnlock() if len(p) == 0 { return 0, nil } @@ -92,8 +96,8 @@ func (m *BufferStore) WriteAt(p []byte, off int64) (int, error) { } func (m *BufferStore) GrowTo(size int64) (err error) { - m.mu.Lock() - defer m.mu.Unlock() + 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 51d3223c26..645de9b6d0 100644 --- a/internal/hybrid_cache/hybrid_cache.go +++ b/internal/hybrid_cache/hybrid_cache.go @@ -4,6 +4,7 @@ import ( "errors" "fmt" "io" + "math" "runtime" "github.com/OpenListTeam/OpenList/v4/internal/cache" @@ -13,7 +14,8 @@ import ( "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 @@ -24,22 +26,17 @@ type HybridCache struct { 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 hc.memoryBacking && hc.exceedsMemoryCeiling(hc.backingOffset, size) { - return nil, mem.ErrNotEnoughMemory - } - 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)), @@ -64,20 +61,17 @@ retry: 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 hc.memoryBacking && hc.exceedsMemoryCeiling(hc.backingOffset, size) { - return nil, mem.ErrNotEnoughMemory - } - 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 } var all []byte @@ -98,9 +92,32 @@ retry: 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) } @@ -161,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 } @@ -276,12 +297,7 @@ func (hc *HybridCache) CopyFromN(src io.Reader, n int64) (written int64, err err return written, nil } -type memoryCheck func(uint64) error - -func selectPolicy(requested cache.Policy, memoryCeiling int64, check memoryCheck) (cache.Policy, error) { - if !requested.IsConcrete() { - return cache.PolicyInherit, fmt.Errorf("invalid cache policy %q", requested) - } +func selectPolicy(requested cache.Policy, memoryCeiling int64) (cache.Policy, error) { switch requested { case cache.PolicyMemory, cache.PolicyDisk: return requested, nil @@ -289,33 +305,18 @@ func selectPolicy(requested cache.Policy, memoryCeiling int64, check memoryCheck if memoryCeiling < 0 { return cache.PolicyDisk, nil } - if memoryCeiling == 0 { - // Zero is a known empty workload, unlike a negative unknown ceiling. - // The hard ceiling still rejects any unexpected writes. - return cache.PolicyMemory, nil - } - if err := check(uint64(memoryCeiling)); err != nil { - return cache.PolicyDisk, nil - } return cache.PolicyMemory, nil default: - panic("unreachable") + return cache.PolicyInherit, fmt.Errorf("invalid cache policy %q", requested) } } -// SelectPolicy resolves auto to a concrete memory or disk policy for a cache -// whose maximum simultaneous memory footprint is memoryCeiling. A negative -// ceiling means that the upper bound is unknown. -func SelectPolicy(requested cache.Policy, memoryCeiling int64) (cache.Policy, error) { - return selectPolicy(requested, memoryCeiling, mem.MemoryGrowCheck) -} - // 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.MemoryGrowCheck) + return newHybridCache(blockSize, memoryCeiling, requested, mem.CacheMemoryBudget) } -func newHybridCache(blockSize uint64, memoryCeiling int64, requested cache.Policy, check memoryCheck) (hc *HybridCache, err error) { +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") } @@ -326,11 +327,21 @@ func newHybridCache(blockSize uint64, memoryCeiling int64, requested cache.Polic } } - selected, err := selectPolicy(requested, memoryCeiling, check) + 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 + } + } if selected == cache.PolicyDisk { if err := hc.initFileCache(); err != nil { return nil, err @@ -344,20 +355,19 @@ func newHybridCache(blockSize uint64, memoryCeiling int64, requested cache.Polic return hc, nil } - if requested == cache.PolicyMemory { - hc.memoryStore, err = mem.NewManagedMemory(blockSize, uint64(memoryCeiling), nil) - if err != nil { - return nil, err - } - return hc, nil - } - - hc.memoryStore, err = mem.NewGuardedMemory(blockSize, uint64(memoryCeiling)) + hc.memoryStore, err = mem.NewManagedMemory(blockSize, uint64(memoryCeiling)) if err == nil { - hc.spillOnMemoryFailure = true + 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) } diff --git a/internal/hybrid_cache/policy_test.go b/internal/hybrid_cache/policy_test.go index db5ee74fea..82328e2cee 100644 --- a/internal/hybrid_cache/policy_test.go +++ b/internal/hybrid_cache/policy_test.go @@ -12,41 +12,30 @@ import ( ) func TestSelectPolicy(t *testing.T) { - errNoMemory := errors.New("no memory") tests := []struct { name string requested cache.Policy ceiling int64 - checkErr error want cache.Policy - checks int }{ - {"explicit memory", cache.PolicyMemory, -1, errNoMemory, cache.PolicyMemory, 0}, - {"explicit disk", cache.PolicyDisk, 1024, nil, cache.PolicyDisk, 0}, - {"auto unknown", cache.PolicyAuto, -1, nil, cache.PolicyDisk, 0}, - {"auto empty", cache.PolicyAuto, 0, nil, cache.PolicyMemory, 0}, - {"auto admitted", cache.PolicyAuto, 1024, nil, cache.PolicyMemory, 1}, - {"auto rejected", cache.PolicyAuto, 1024, errNoMemory, cache.PolicyDisk, 1}, + {"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) { - checks := 0 - got, err := selectPolicy(tt.requested, tt.ceiling, func(size uint64) error { - checks++ - if size != uint64(tt.ceiling) { - t.Fatalf("memory check size = %d, want %d", size, tt.ceiling) - } - return tt.checkErr - }) + got, err := selectPolicy(tt.requested, tt.ceiling) if err != nil { t.Fatalf("selectPolicy() error = %v", err) } - if got != tt.want || checks != tt.checks { - t.Fatalf("selectPolicy() = %q with %d checks, want %q with %d", got, checks, tt.want, tt.checks) + if got != tt.want { + t.Fatalf("selectPolicy() = %q, want %q", got, tt.want) } }) } - if _, err := selectPolicy(cache.PolicyInherit, 1, func(uint64) error { return nil }); err == nil { + if _, err := selectPolicy(cache.PolicyInherit, 1); err == nil { t.Fatal("selectPolicy() expected an error for inherit") } } @@ -75,31 +64,23 @@ func TestHybridCachePolicyBackingSelection(t *testing.T) { name string requested cache.Policy ceiling int64 - checkErr error + budgetCapacity uint64 wantMemory bool - wantCheckCalls int }{ - {name: "memory uses memory", requested: cache.PolicyMemory, ceiling: 8, wantMemory: true}, + {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, wantMemory: true, wantCheckCalls: 1}, - {name: "auto rejected uses disk", requested: cache.PolicyAuto, ceiling: 8, checkErr: mem.ErrNotEnoughMemory, wantCheckCalls: 1}, + {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) - checkCalls := 0 - hc, err := newHybridCache(4, tt.ceiling, tt.requested, func(uint64) error { - checkCalls++ - return tt.checkErr - }) + hc, err := newHybridCache(4, tt.ceiling, tt.requested, mem.NewBudget(tt.budgetCapacity)) if err != nil { t.Fatalf("newHybridCache() error = %v", err) } - if checkCalls != tt.wantCheckCalls { - t.Fatalf("memory check calls = %d, want %d", checkCalls, tt.wantCheckCalls) - } if hc.memoryBacking != tt.wantMemory { t.Fatalf("memory backing = %v, want %v", hc.memoryBacking, tt.wantMemory) @@ -140,7 +121,7 @@ func TestHybridCachePolicyBackingSelection(t *testing.T) { func TestHybridCacheAutoRejectsWholeCeiling(t *testing.T) { withCacheConfig(t, 1024) - hc, err := NewHybridCache(4, 8, cache.PolicyAuto) + hc, err := newHybridCache(4, 8, cache.PolicyAuto, mem.NewBudget(7)) if err != nil { t.Fatalf("NewHybridCache() error = %v", err) } @@ -155,10 +136,7 @@ func TestHybridCacheAutoRejectsWholeCeiling(t *testing.T) { func TestHybridCacheZeroCeilingRejectsUnexpectedWrites(t *testing.T) { withCacheConfig(t, 8) - hc, err := newHybridCache(4, 0, cache.PolicyAuto, func(uint64) error { - t.Fatal("zero ceiling must not check memory") - return nil - }) + hc, err := newHybridCache(4, 0, cache.PolicyAuto, mem.NewBudget(0)) if err != nil { t.Fatalf("newHybridCache() error = %v", err) } @@ -208,13 +186,92 @@ func TestHybridCacheUnknownMemory(t *testing.T) { } } +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 { @@ -223,6 +280,9 @@ func TestHybridCacheAutoSpillKeepsMemoryPrefix(t *testing.T) { 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) @@ -253,14 +313,18 @@ func withCacheConfig(t *testing.T, autoMemoryLimit uint64) { t.Helper() oldConf := conf.Conf oldLimit := conf.AutoMemoryLimit - oldMinFreeMemory := conf.MinFreeMemory + oldBudgetCapacity := mem.CacheMemoryBudget.Capacity() + oldBudgetReserved := mem.CacheMemoryBudget.Reserved() conf.Conf = &conf.Config{TempDir: t.TempDir()} conf.AutoMemoryLimit = autoMemoryLimit - conf.MinFreeMemory = 0 + mem.CacheMemoryBudget.SetCapacity(1 << 40) t.Cleanup(func() { conf.Conf = oldConf conf.AutoMemoryLimit = oldLimit - conf.MinFreeMemory = oldMinFreeMemory + if reserved := mem.CacheMemoryBudget.Reserved(); reserved != oldBudgetReserved { + t.Errorf("cache memory reserved after test = %d, want %d", reserved, oldBudgetReserved) + } + mem.CacheMemoryBudget.SetCapacity(oldBudgetCapacity) }) } 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 index 54c2bcb443..92bba2c160 100644 --- a/internal/mem/snapshot.go +++ b/internal/mem/snapshot.go @@ -23,6 +23,8 @@ type MemorySnapshot struct { 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 @@ -46,10 +48,11 @@ type cgroupMemory struct { func combineMemorySnapshots(host *MemorySnapshot, cgroup cgroupMemory) MemorySnapshot { if host == nil { + available := min(cgroup.available, cgroup.limit) return MemorySnapshot{ Limit: cgroup.limit, - Used: cgroup.limit - min(cgroup.available, cgroup.limit), - Available: min(cgroup.available, cgroup.limit), + Used: saturatingSub(cgroup.limit, available), + Available: available, Source: cgroup.source, } } @@ -60,7 +63,7 @@ func combineMemorySnapshots(host *MemorySnapshot, cgroup cgroupMemory) MemorySna available := min(host.Available, cgroup.available, limit) return MemorySnapshot{ Limit: limit, - Used: limit - available, + Used: saturatingSub(limit, available), Available: available, Source: cgroup.source, } diff --git a/internal/mem/snapshot_cgroup.go b/internal/mem/snapshot_cgroup.go index f701ed343c..0f179904c4 100644 --- a/internal/mem/snapshot_cgroup.go +++ b/internal/mem/snapshot_cgroup.go @@ -9,6 +9,10 @@ import ( "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" 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 8e1ffd8725..d84892cfab 100644 --- a/internal/mem/utils.go +++ b/internal/mem/utils.go @@ -4,57 +4,12 @@ import ( "errors" "fmt" "runtime" - "sync/atomic" - - "github.com/OpenListTeam/OpenList/v4/internal/conf" - "github.com/OpenListTeam/OpenList/v4/pkg/singleflight" ) 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) { - snapshot, err := GetMemorySnapshot() - if err != nil { - return nil, err - } - if snapshot.Available < conf.MinFreeMemory { - return nil, ErrNotEnoughMemory - } - var res atomic.Uint64 - res.Store(snapshot.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) { - return NewManagedMemory(cap, max, MemoryGrowCheck) -} - // NewManagedMemory creates memory with panic recovery and lifecycle cleanup. -// A nil growCheck intentionally permits growth without an availability check. -func NewManagedMemory(cap, max uint64, growCheck GrowCheck) (m LinearMemory, err error) { - if growCheck != nil { - if err := growCheck(cap); err != nil { - return nil, err - } - } +func NewManagedMemory(cap, max uint64) (m LinearMemory, err error) { defer func() { if r := recover(); r != nil { err = fmt.Errorf("%w: %v", ErrNotEnoughMemory, r) @@ -64,11 +19,6 @@ func NewManagedMemory(cap, max uint64, growCheck GrowCheck) (m LinearMemory, err if err != nil { return nil, err } - if growCheck != nil { - if s, ok := m.(interface{ SetGrowCheck(GrowCheck) }); ok { - s.SetGrowCheck(growCheck) - } - } gm := &guardedMemory{LinearMemory: m} gm.cleanup = runtime.AddCleanup(gm, func(m LinearMemory) { m.Free() diff --git a/internal/net/request_test.go b/internal/net/request_test.go index 6bf9b0071f..bd6f2e92c0 100644 --- a/internal/net/request_test.go +++ b/internal/net/request_test.go @@ -17,6 +17,7 @@ import ( "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" @@ -33,6 +34,18 @@ 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 @@ -66,6 +79,7 @@ func TestDownloaderMemoryCeiling(t *testing.T) { } func TestTryDownloadChunkReturnsPrefetchAllocationError(t *testing.T) { + withDownloaderConfig(t) hc, err := hcache.NewHybridCache(4, 4, cache.PolicyMemory) if err != nil { t.Fatalf("NewHybridCache() error = %v", err) @@ -110,6 +124,7 @@ func TestTryDownloadChunkReturnsPrefetchAllocationError(t *testing.T) { } 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 @@ -161,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) @@ -196,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) @@ -261,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/stream.go b/internal/stream/stream.go index 37df320113..80b6a15876 100644 --- a/internal/stream/stream.go +++ b/internal/stream/stream.go @@ -27,14 +27,13 @@ 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 - sizeSet bool - cachePolicyOverride cache.Policy - cachePolicyResolved cache.Policy - cachePolicyLocked bool - 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 { @@ -63,18 +62,22 @@ 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.cachePolicyResolved, nil + return f.cachePolicy, nil } - return cache.ResolvePolicy(f.cachePolicyOverride, conf.CachePolicy) + return cache.ResolvePolicy(f.cachePolicy, conf.CachePolicy) } -func (f *FileStream) freezeCachePolicy(policy cache.Policy) { - f.cachePolicyResolved = policy +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) @@ -82,7 +85,7 @@ func (f *FileStream) SetCachePolicy(policy cache.Policy) error { if f.cachePolicyLocked { return errors.New("cache policy cannot be changed after cache initialization") } - f.cachePolicyOverride = policy + f.cachePolicy = policy return nil } @@ -210,7 +213,7 @@ func (f *FileStream) ensureCache(size int64) (model.File, error) { if err != nil { return nil, err } - f.freezeCachePolicy(policy) + f.lockCachePolicy(policy) f.peek = buffer.NewDynamicReadAtSeeker(f.hc) f.oriReader = f.Reader f.Reader = io.MultiReader(f.peek, f.oriReader) diff --git a/internal/stream/stream_test.go b/internal/stream/stream_test.go index c3e1792152..b6d56c5c30 100644 --- a/internal/stream/stream_test.go +++ b/internal/stream/stream_test.go @@ -10,6 +10,7 @@ import ( "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" @@ -22,19 +23,23 @@ func withStreamCacheConfig(t *testing.T) { oldPolicy := conf.CachePolicy oldBlockLimit := conf.MaxBlockLimit oldAutoMemoryLimit := conf.AutoMemoryLimit - oldMinFreeMemory := conf.MinFreeMemory + oldBudgetCapacity := mem.CacheMemoryBudget.Capacity() + oldBudgetReserved := mem.CacheMemoryBudget.Reserved() t.Cleanup(func() { conf.Conf = oldConf conf.CachePolicy = oldPolicy conf.MaxBlockLimit = oldBlockLimit conf.AutoMemoryLimit = oldAutoMemoryLimit - conf.MinFreeMemory = oldMinFreeMemory + 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 - conf.MinFreeMemory = 0 + mem.CacheMemoryBudget.SetCapacity(0) } func TestFileStreamCachePolicy(t *testing.T) { @@ -226,6 +231,7 @@ func TestRangeRead(t *testing.T) { }, Reader: io.NopCloser(bytes.NewReader(buf)), } + t.Cleanup(func() { _ = f.Close() }) conf.AutoMemoryLimit = 0 conf.MaxBlockLimit = 15 tests := []struct { @@ -298,6 +304,7 @@ func TestPreHash(t *testing.T) { }, Reader: io.NopCloser(bytes.NewReader(buf)), } + t.Cleanup(func() { _ = f.Close() }) conf.AutoMemoryLimit = 0 conf.MaxBlockLimit = 15 @@ -328,6 +335,7 @@ func TestStreamSectionReader(t *testing.T) { }, Reader: io.NopCloser(bytes.NewReader(buf)), } + t.Cleanup(func() { _ = f.Close() }) conf.AutoMemoryLimit = 0 conf.MaxBlockLimit = 2 << 10 partSize := 3 << 10 diff --git a/internal/stream/util.go b/internal/stream/util.go index 13d05f6e3e..ed569e36f9 100644 --- a/internal/stream/util.go +++ b/internal/stream/util.go @@ -187,8 +187,8 @@ type cachePolicyGetter interface { GetCachePolicy() (cache.Policy, error) } -type cachePolicyFreezer interface { - freezeCachePolicy(cache.Policy) +type cachePolicyLocker interface { + lockCachePolicy(cache.Policy) } func getCachePolicy(file model.FileStreamer) (cache.Policy, error) { @@ -219,8 +219,8 @@ func NewStreamSectionReader(file model.FileStreamer, sectionSize int, up *model. if err != nil { return nil, err } - if freezer, ok := file.(cachePolicyFreezer); ok { - freezer.freezeCachePolicy(policy) + if locker, ok := file.(cachePolicyLocker); ok { + locker.lockCachePolicy(policy) } file.Add(hc) return &hybridSectionReader{file: file, hc: hc}, nil