From cf9191bef6daf7f4b92d4321298ca9d86808ff6e Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Mon, 24 Apr 2023 13:16:24 -0700 Subject: [PATCH 01/67] add start of cache tool packages, runner Migrated-from: bradfitz/go-tool-cache@c09b10de4d296896fbd25141cb7d03ad79137b77 --- go.mod | 2 + go.sum | 0 gocache/.gitignore | 1 + gocache/cacheproc/cacheproc.go | 218 +++++++++++++++++++++++++++++++++ gocache/cachers/disk.go | 106 ++++++++++++++++ gocache/wire/wire.go | 88 +++++++++++++ 6 files changed, 415 insertions(+) create mode 100644 go.mod create mode 100644 go.sum create mode 100644 gocache/.gitignore create mode 100644 gocache/cacheproc/cacheproc.go create mode 100644 gocache/cachers/disk.go create mode 100644 gocache/wire/wire.go diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..04a1142 --- /dev/null +++ b/go.mod @@ -0,0 +1,2 @@ +module github.com/tailscale/tb +go 1.19 diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..e69de29 diff --git a/gocache/.gitignore b/gocache/.gitignore new file mode 100644 index 0000000..b25c15b --- /dev/null +++ b/gocache/.gitignore @@ -0,0 +1 @@ +*~ diff --git a/gocache/cacheproc/cacheproc.go b/gocache/cacheproc/cacheproc.go new file mode 100644 index 0000000..da2b3f7 --- /dev/null +++ b/gocache/cacheproc/cacheproc.go @@ -0,0 +1,218 @@ +// Copyright 2023 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +// Package cacheproc implements the mechanics of talking to cmd/go's GOCACHE protocol +// so you can write a caching child process at a higher level. +package cacheproc + +import ( + "bufio" + "bytes" + "context" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "os" + "sync" + "sync/atomic" + + "github.com/tailscale/tb/gocache/wire" +) + +// Process implements the cmd/go JSON protocol over stdin & stdout via three +// funcs that callers can optionally implement. +type Process struct { + // Get optionally specifies a func to look up something from the cache. If + // nil, all gets are treated as cache misses.touch + // + // The actionID is a lowercase hex string of unspecified format or length. + // + // The returned outputID must be the same outputID provided to Put earlier; + // it will be a lowercase hex string of unspecified hash function or length. + // + // On cache miss, return all zero values (no error). On cache hit, diskPath + // must be the absolute path to a regular file; its size and modtime are + // returned to cmd/go. + // + // If the returned diskPath doesn't exist, it's treated as a cache miss. + Get func(ctx context.Context, actionID string) (outputID, diskPath string, _ error) + + // Put optionally specifies a func to add something to the cache. + // The actionID and objectID is a lowercase hex string of unspecified format or length. + // On success, diskPath must be the absolute path to a regular file. + // If nil, cmd/go may write to disk elsewhere as needed. + Put func(ctx context.Context, actionID, objectID string, size int64, r io.Reader) (diskPath string, _ error) + + // Close optionally specifies a func to run when the cmd/go tool is + // shutting down. + Close func() error + + Gets atomic.Int64 + GetHits atomic.Int64 + GetMisses atomic.Int64 + GetErrors atomic.Int64 + Puts atomic.Int64 + PutErrors atomic.Int64 +} + +func (p *Process) Run() error { + br := bufio.NewReader(os.Stdin) + jd := json.NewDecoder(br) + + bw := bufio.NewWriter(os.Stdout) + je := json.NewEncoder(bw) + + var caps []wire.Cmd + if p.Get != nil { + caps = append(caps, "get") + } + if p.Put != nil { + caps = append(caps, "put") + } + if p.Close != nil { + caps = append(caps, "close") + } + je.Encode(&wire.Response{KnownCommands: caps}) + if err := bw.Flush(); err != nil { + return err + } + + var wmu sync.Mutex // guards writing responses + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + for { + var req wire.Request + if err := jd.Decode(&req); err != nil { + if errors.Is(err, io.EOF) { + return nil + } + return err + } + if req.Command == wire.CmdPut && req.BodySize > 0 { + // TODO(bradfitz): stream this and pass a checksum-validating + // io.Reader that validates on EOF. + var bodyb []byte + if err := jd.Decode(&bodyb); err != nil { + log.Fatal(err) + } + if int64(len(bodyb)) != req.BodySize { + log.Fatalf("only got %d bytes of declared %d", len(bodyb), req.BodySize) + } + req.Body = bytes.NewReader(bodyb) + } + go func() { + res := &wire.Response{ID: req.ID} + ctx := ctx // TODO: include req ID as a context.Value for tracing? + if err := p.handleRequest(ctx, &req, res); err != nil { + res.Err = err.Error() + } + wmu.Lock() + defer wmu.Unlock() + je.Encode(res) + bw.Flush() + }() + } +} + +func (p *Process) handleRequest(ctx context.Context, req *wire.Request, res *wire.Response) error { + switch req.Command { + default: + return errors.New("unknown command") + case "close": + if p.Close != nil { + return p.Close() + } + return nil + case "get": + return p.handleGet(ctx, req, res) + case "put": + return p.handlePut(ctx, req, res) + } +} + +func (p *Process) handleGet(ctx context.Context, req *wire.Request, res *wire.Response) (retErr error) { + p.Gets.Add(1) + defer func() { + if retErr != nil { + p.GetErrors.Add(1) + } else if res.Miss { + p.GetMisses.Add(1) + } else { + p.GetHits.Add(1) + } + }() + if p.Get == nil { + res.Miss = true + return nil + } + outputID, diskPath, err := p.Get(ctx, fmt.Sprintf("%x", req.ActionID)) + if err != nil { + return err + } + if outputID == "" && diskPath == "" { + res.Miss = true + return nil + } + if outputID == "" { + return errors.New("no outputID") + } + res.OutputID, err = hex.DecodeString(outputID) + if err != nil { + return fmt.Errorf("invalid OutputID: %v", err) + } + fi, err := os.Stat(diskPath) + if err != nil { + if os.IsNotExist(err) { + res.Miss = true + return nil + } + return err + } + if !fi.Mode().IsRegular() { + return fmt.Errorf("not a regular file") + } + res.Size = fi.Size() + res.TimeNanos = fi.ModTime().UnixNano() + res.DiskPath = diskPath + return nil +} + +func (p *Process) handlePut(ctx context.Context, req *wire.Request, res *wire.Response) (retErr error) { + actionID, objectID := fmt.Sprintf("%x", req.ActionID), fmt.Sprintf("%x", req.ObjectID) + p.Puts.Add(1) + defer func() { + if retErr != nil { + p.PutErrors.Add(1) + log.Printf("put(action %s, obj %s, %v bytes): %v", actionID, objectID, req.BodySize, retErr) + } + }() + if p.Put == nil { + if req.Body != nil { + io.Copy(io.Discard, req.Body) + } + return nil + } + var body io.Reader = req.Body + if body == nil { + body = bytes.NewReader(nil) + } + diskPath, err := p.Put(ctx, actionID, objectID, req.BodySize, body) + if err != nil { + return err + } + fi, err := os.Stat(diskPath) + if err != nil { + return fmt.Errorf("stat after successful Put: %w", err) + } + if fi.Size() != req.BodySize { + return fmt.Errorf("failed to write file to disk with right size: disk=%v; wanted=%v", fi.Size(), req.BodySize) + } + res.DiskPath = diskPath + return nil +} diff --git a/gocache/cachers/disk.go b/gocache/cachers/disk.go new file mode 100644 index 0000000..cd612e1 --- /dev/null +++ b/gocache/cachers/disk.go @@ -0,0 +1,106 @@ +package cachers + +import ( + "bytes" + "context" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "log" + "os" + "path/filepath" + "time" +) + +// indexEntry is the metadata that DiskCache stores on disk for an ActionID. +type indexEntry struct { + Version int `json:"v"` + OutputID string `json:"o"` + Size int64 `json:"n"` + TimeNanos int64 `json:"t"` +} + +type DiskCache struct { + Dir string +} + +func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPath string, err error) { + actionFile := filepath.Join(dc.Dir, fmt.Sprintf("a-%s", actionID)) + ij, err := os.ReadFile(actionFile) + if err != nil { + if os.IsNotExist(err) { + err = nil + log.Printf("Miss: %v", actionID) + } + return "", "", err + } + var ie indexEntry + if err := json.Unmarshal(ij, &ie); err != nil { + log.Printf("Warning: JSON error for action %q: %v", actionID, err) + return "", "", nil + } + if _, err := hex.DecodeString(ie.OutputID); err != nil { + // Protect against malicious non-hex OutputID on disk + return "", "", nil + } + return ie.OutputID, filepath.Join(dc.Dir, fmt.Sprintf("o-%v", ie.OutputID)), nil +} + +func (dc *DiskCache) Put(ctx context.Context, actionID, objectID string, size int64, body io.Reader) (diskPath string, _ error) { + file := filepath.Join(dc.Dir, fmt.Sprintf("o-%s", objectID)) + + // Special case empty files; they're both common and easier to do race-free. + if size == 0 { + zf, err := os.OpenFile(file, os.O_CREATE|os.O_RDWR, 0644) + if err != nil { + return "", err + } + zf.Close() + } else { + wrote, err := writeAtomic(file, body) + if err != nil { + return "", err + } + if wrote != size { + return "", fmt.Errorf("wrote %d bytes, expected %d", wrote, size) + } + } + + ij, err := json.Marshal(indexEntry{ + Version: 1, + OutputID: objectID, + Size: size, + TimeNanos: time.Now().UnixNano(), + }) + if err != nil { + return "", err + } + actionFile := filepath.Join(dc.Dir, fmt.Sprintf("a-%s", actionID)) + if _, err := writeAtomic(actionFile, bytes.NewReader(ij)); err != nil { + return "", err + } + return file, nil +} + +func writeAtomic(dest string, r io.Reader) (int64, error) { + tf, err := os.CreateTemp(filepath.Dir(dest), filepath.Base(dest)+".*") + if err != nil { + return 0, err + } + size, err := io.Copy(tf, r) + if err != nil { + tf.Close() + os.Remove(tf.Name()) + return 0, err + } + if err := tf.Close(); err != nil { + os.Remove(tf.Name()) + return 0, err + } + if err := os.Rename(tf.Name(), dest); err != nil { + os.Remove(tf.Name()) + return 0, err + } + return size, nil +} diff --git a/gocache/wire/wire.go b/gocache/wire/wire.go new file mode 100644 index 0000000..36f4810 --- /dev/null +++ b/gocache/wire/wire.go @@ -0,0 +1,88 @@ +// Copyright 2023 The Go Authors. All rights reserved. +// Use of this source code is governed by a BSD-style +// license that can be found in the LICENSE file. + +// Package wire contains the JSON types that cmd/go uses +// to communicate with child processes implementing +// the cache interface. +package wire + +import "io" + +// Cmd is a command that can be issued to a child process. +// +// If the interface needs to grow, we can add new commands or new versioned +// commands like "get2". +type Cmd string + +const ( + CmdGet = Cmd("get") + CmdPut = Cmd("put") + CmdClose = Cmd("close") +) + +// Request is the JSON-encoded message that's sent from cmd/go to +// the GOCACHEPROG child process over stdin. Each JSON object is on its +// own line. A Request of Type "put" with BodySize > 0 will be followed +// by a line containing a base64-encoded JSON string literal of the body. +type Request struct { + // ID is a unique number per process across all requests. + // It must be echoed in the Response from the child. + ID int64 + + // Command is the type of request. + // The cmd/go tool will only send commands that were declared + // as supported by the child. + Command Cmd + + // ActionID is non-nil for get and puts. + ActionID []byte `json:",omitempty"` // or nil if not used + + // ObjectID is set for Type "put" and "output-file". + ObjectID []byte `json:",omitempty"` // or nil if not used + + // Body is the body for "put" requests. It's sent after the JSON object + // as a base64-encoded JSON string when BodySize is non-zero. + // It's sent as a separate JSON value instead of being a struct field + // send in this JSON object so large values can be streamed in both directions. + // The base64 string body of a Request will always be written + // immediately after the JSON object and a newline. + Body io.Reader `json:"-"` + + // BodySize is the number of bytes of Body. If zero, the body isn't written. + BodySize int64 `json:",omitempty"` +} + +// Response is the JSON response from the child process to cmd/go. +// +// With the exception of the first protocol message that the child writes to its +// stdout with ID==0 and KnownCommands populated, these are only sent in +// response to a Request from cmd/go. +// +// Responses can be sent in any order. The ID must match the request they're +// replying to. +type Response struct { + ID int64 // that corresponds to Request; they can be answered out of order + Err string `json:",omitempty"` // if non-empty, the error + + // KnownCommands is included in the first message that cache helper program + // writes to stdout on startup (with ID==0). It includes the + // Request.Command types that are supported by the program. + // + // This lets us extend the gracefully over time (adding "get2", etc), or + // fail gracefully when needed. It also lets us verify the program + // wants to be a cache helper. + KnownCommands []Cmd `json:",omitempty"` + + // For Get requests. + + Miss bool `json:",omitempty"` // cache miss + OutputID []byte `json:",omitempty"` + Size int64 `json:",omitempty"` + TimeNanos int64 `json:",omitempty"` // TODO(bradfitz): document + + // DiskPath is the absolute path on disk of the ObjectID corresponding + // a "get" request's ActionID (on cache hit) or a "put" request's + // provided ObjectID. + DiskPath string `json:",omitempty"` +} From 70bd57f1c569142d58d7d2afe4b259c9e03acc63 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Mon, 24 Apr 2023 13:26:02 -0700 Subject: [PATCH 02/67] README.md: add Migrated-from: bradfitz/go-tool-cache@4fed81dfb5bb21520a11de6981ca626f7ccd016a --- gocache/README.md | 12 ++++++++++++ 1 file changed, 12 insertions(+) create mode 100644 gocache/README.md diff --git a/gocache/README.md b/gocache/README.md new file mode 100644 index 0000000..6fbea9e --- /dev/null +++ b/gocache/README.md @@ -0,0 +1,12 @@ +# go-tool-cache + +Like Go's built-in build/test caching but wish it weren't purely stored on local disk in the `$GOCACHE` directory? + +Want to share your cache over the network between your various machines, coworkers, and CI runs without all that GitHub actions/caches tarring and untarring? + +Along with a [modification to Go's `cmd/go` tool](https://go-review.googlesource.com/c/go/+/486715) ([open proposal](https://github.com/golang/go/issues/59719)), this repo lets you write +custom `GOCACHE` implementations to handle the cache however you'd like. + +## Status + +Currently you need to build your own Go toolchain to use this. As of 2023-04-24 it's still an open proposal & work in progress. From 474d34a383ae083e0b73c06d3012f727b3fe36e6 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Mon, 24 Apr 2023 13:31:12 -0700 Subject: [PATCH 03/67] README.md: usage, examples Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@e435233295d95f4caf2f69267e8b85162262f6d3 --- gocache/README.md | 30 ++++++++++++++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/gocache/README.md b/gocache/README.md index 6fbea9e..cc5fff8 100644 --- a/gocache/README.md +++ b/gocache/README.md @@ -10,3 +10,33 @@ custom `GOCACHE` implementations to handle the cache however you'd like. ## Status Currently you need to build your own Go toolchain to use this. As of 2023-04-24 it's still an open proposal & work in progress. + +## Using + +First, build your cache child process. For example, + +```sh +$ go install github.com/bradfitz/cmd/go-cacher +``` + +Then tell Go to use it: + +```sh +$ GOCACHEPROG=$HOME/go/bin/go-cacher go install std +``` + +See some stats: + +```sh +$ GOCACHEPROG="$HOME/go/bin/go-cacher --verbose" go install std +Defaulting to cache dir /home/bradfitz/.cache/go-cacher ... +cacher: closing; 548 gets (0 hits, 548 misses, 0 errors); 1090 puts (0 errors) +``` + +Run it again and watch the hit rate go up: + +```sh +$ GOCACHEPROG="$HOME/go/bin/go-cacher --verbose" go install std +Defaulting to cache dir /home/bradfitz/.cache/go-cacher ... +cacher: closing; 808 gets (808 hits, 0 misses, 0 errors); 0 puts (0 errors) +``` \ No newline at end of file From dbcd5c90ee2b9149aae078be18b645094f92c34e Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Mon, 24 Apr 2023 16:29:58 -0700 Subject: [PATCH 04/67] add HTTP client+server support Migrated-from: bradfitz/go-tool-cache@fa12dc10f93c90d6f3e99ea1e15c4775816207aa --- gocache/cachers/disk.go | 21 +++++- gocache/cachers/http.go | 145 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 164 insertions(+), 2 deletions(-) create mode 100644 gocache/cachers/http.go diff --git a/gocache/cachers/disk.go b/gocache/cachers/disk.go index cd612e1..3259b80 100644 --- a/gocache/cachers/disk.go +++ b/gocache/cachers/disk.go @@ -22,7 +22,8 @@ type indexEntry struct { } type DiskCache struct { - Dir string + Dir string + Verbose bool } func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPath string, err error) { @@ -31,7 +32,9 @@ func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPa if err != nil { if os.IsNotExist(err) { err = nil - log.Printf("Miss: %v", actionID) + if dc.Verbose { + log.Printf("disk miss: %v", actionID) + } } return "", "", err } @@ -47,6 +50,20 @@ func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPa return ie.OutputID, filepath.Join(dc.Dir, fmt.Sprintf("o-%v", ie.OutputID)), nil } +func (dc *DiskCache) OutputFilename(objectID string) string { + if len(objectID) < 4 || len(objectID) > 1000 { + return "" + } + for i := range objectID { + b := objectID[i] + if b >= '0' && b <= '9' || b >= 'a' && b <= 'f' { + continue + } + return "" + } + return filepath.Join(dc.Dir, fmt.Sprintf("o-%s", objectID)) +} + func (dc *DiskCache) Put(ctx context.Context, actionID, objectID string, size int64, body io.Reader) (diskPath string, _ error) { file := filepath.Join(dc.Dir, fmt.Sprintf("o-%s", objectID)) diff --git a/gocache/cachers/http.go b/gocache/cachers/http.go new file mode 100644 index 0000000..2022b04 --- /dev/null +++ b/gocache/cachers/http.go @@ -0,0 +1,145 @@ +package cachers + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "log" + "net/http" + "os" +) + +// ActionValue is the JSON value returned by the cacher server for an GET /action request. +type ActionValue struct { + OutputID string `json:"outputID"` + Size int64 `json:"size"` +} + +type HTTPClient struct { + // BaseURL is the base URL of the cacher server, like "http://localhost:31364". + BaseURL string + + // Disk is where to write the output files to local disk, as required by the + // cache protocol. + Disk *DiskCache + + // HTTPClient optionally specifies the http.Client to use. + // If nil, http.DefaultClient is used. + HTTPClient *http.Client + + // Verbose optionally specifies whether to log verbose messages. + Verbose bool +} + +func (c *HTTPClient) httpClient() *http.Client { + if c.HTTPClient != nil { + return c.HTTPClient + } + return http.DefaultClient +} + +func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPath string, err error) { + outputID, diskPath, err = c.Disk.Get(ctx, actionID) + if err == nil && outputID != "" { + return outputID, diskPath, nil + } + + req, _ := http.NewRequestWithContext(ctx, "GET", c.BaseURL+"/action/"+actionID, nil) + + res, err := c.httpClient().Do(req) + if err != nil { + return "", "", err + } + defer res.Body.Close() + if res.StatusCode == http.StatusNotFound { + return "", "", nil + } + if res.StatusCode != http.StatusOK { + return "", "", fmt.Errorf("unexpected GET /action/%s status %v", actionID, res.Status) + } + var av ActionValue + if err := json.NewDecoder(res.Body).Decode(&av); err != nil { + return "", "", err + } + outputID = av.OutputID + + diskPath = c.Disk.OutputFilename(outputID) + if diskPath == "" { + return "", "", fmt.Errorf("invalid outputID %q", av.OutputID) + } + + // See if it's already on disk. + fi, err := os.Stat(diskPath) + if err == nil && fi.Size() == av.Size { + return av.OutputID, diskPath, nil + } + + // If not on disk, download it to disk. + req, _ = http.NewRequestWithContext(ctx, "GET", c.BaseURL+"/output/"+outputID, nil) + res, err = c.httpClient().Do(req) + if err != nil { + return "", "", err + } + defer res.Body.Close() + if res.StatusCode == http.StatusNotFound { + return "", "", nil + } + if res.StatusCode != http.StatusOK { + return "", "", fmt.Errorf("unexpected GET /output/%s status %v", outputID, res.Status) + } + if res.ContentLength == -1 { + return "", "", fmt.Errorf("no Content-Length from server") + } + diskPath, err = c.Disk.Put(ctx, actionID, outputID, res.ContentLength, res.Body) + return outputID, diskPath, err +} + +func (c *HTTPClient) Put(ctx context.Context, actionID, outputID string, size int64, body io.Reader) (diskPath string, _ error) { + // Write to disk locally as we write it remotely, as we need to guarantee + // it's on disk locally for the caller. + pr, pw := io.Pipe() + diskPutCh := make(chan any, 1) + go func() { + var putBody io.Reader = pr + if size == 0 { + putBody = bytes.NewReader(nil) + } + diskPath, err := c.Disk.Put(ctx, actionID, outputID, size, putBody) + if err != nil { + diskPutCh <- err + } else { + diskPutCh <- diskPath + } + }() + + var putBody io.Reader + if size == 0 { + // Special case the empty file so NewRequest sets "Content-Length: 0", + // as opposed to thinking we didn't set it and not being able to sniff its size + // from the type. + putBody = bytes.NewReader(nil) + } else { + putBody = io.TeeReader(body, pw) + } + req, _ := http.NewRequestWithContext(ctx, "PUT", c.BaseURL+"/"+actionID+"/"+outputID, putBody) + req.ContentLength = size + res, err := c.httpClient().Do(req) + pw.Close() + if err != nil { + log.Printf("error PUT /%s/%s: %v", actionID, outputID, err) + return "", err + } + defer res.Body.Close() + if res.StatusCode != http.StatusNoContent { + all, _ := io.ReadAll(io.LimitReader(res.Body, 4<<10)) + return "", fmt.Errorf("unexpected PUT /%s/%s status %v: %s", actionID, outputID, res.Status, all) + } + v := <-diskPutCh + if err, ok := v.(error); ok { + log.Printf("HTTPClient.Put local disk error: %v", err) + return "", err + } + return v.(string), nil +} From ff29b8f657bdebfa403b7d56cd3174628a3b6033 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Mon, 24 Apr 2023 17:26:07 -0700 Subject: [PATCH 05/67] fix an HTTP issue with multiple actionIDs sharing same 0-length outputID Migrated-from: bradfitz/go-tool-cache@76f4068d03b63c9c79d0f1b26ebc56cadc985e67 --- gocache/cachers/http.go | 48 ++++++++++++++++++----------------------- 1 file changed, 21 insertions(+), 27 deletions(-) diff --git a/gocache/cachers/http.go b/gocache/cachers/http.go index 2022b04..1eed0bd 100644 --- a/gocache/cachers/http.go +++ b/gocache/cachers/http.go @@ -8,7 +8,6 @@ import ( "io" "log" "net/http" - "os" ) // ActionValue is the JSON value returned by the cacher server for an GET /action request. @@ -65,34 +64,29 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa } outputID = av.OutputID - diskPath = c.Disk.OutputFilename(outputID) - if diskPath == "" { - return "", "", fmt.Errorf("invalid outputID %q", av.OutputID) - } - - // See if it's already on disk. - fi, err := os.Stat(diskPath) - if err == nil && fi.Size() == av.Size { - return av.OutputID, diskPath, nil - } - // If not on disk, download it to disk. - req, _ = http.NewRequestWithContext(ctx, "GET", c.BaseURL+"/output/"+outputID, nil) - res, err = c.httpClient().Do(req) - if err != nil { - return "", "", err - } - defer res.Body.Close() - if res.StatusCode == http.StatusNotFound { - return "", "", nil - } - if res.StatusCode != http.StatusOK { - return "", "", fmt.Errorf("unexpected GET /output/%s status %v", outputID, res.Status) - } - if res.ContentLength == -1 { - return "", "", fmt.Errorf("no Content-Length from server") + var putBody io.Reader + if av.Size == 0 { + putBody = bytes.NewReader(nil) + } else { + req, _ = http.NewRequestWithContext(ctx, "GET", c.BaseURL+"/output/"+outputID, nil) + res, err = c.httpClient().Do(req) + if err != nil { + return "", "", err + } + defer res.Body.Close() + if res.StatusCode == http.StatusNotFound { + return "", "", nil + } + if res.StatusCode != http.StatusOK { + return "", "", fmt.Errorf("unexpected GET /output/%s status %v", outputID, res.Status) + } + if res.ContentLength == -1 { + return "", "", fmt.Errorf("no Content-Length from server") + } + putBody = res.Body } - diskPath, err = c.Disk.Put(ctx, actionID, outputID, res.ContentLength, res.Body) + diskPath, err = c.Disk.Put(ctx, actionID, outputID, av.Size, putBody) return outputID, diskPath, err } From a05d6add653a8fd16749a571c839ea2fdd2da197 Mon Sep 17 00:00:00 2001 From: xieyuschen Date: Tue, 25 Apr 2023 14:48:04 +0800 Subject: [PATCH 06/67] chore: fix go version to 1.20 Migrated-from: bradfitz/go-tool-cache@a411ed23e0a46ad9ef951e214b3ab5ba8f9b55dd --- go.mod | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/go.mod b/go.mod index 04a1142..60b4913 100644 --- a/go.mod +++ b/go.mod @@ -1,2 +1,2 @@ module github.com/tailscale/tb -go 1.19 +go 1.20 From 45052200c600147a8d9a7e3da0dccf4405fe294f Mon Sep 17 00:00:00 2001 From: Sal Sal <0xack13@gmail.com> Date: Tue, 25 Apr 2023 04:22:23 +0300 Subject: [PATCH 07/67] Fix go install cmd package path Migrated-from: bradfitz/go-tool-cache@ef6c7b1b26e95d55cf68d316459e90de46979d2c --- gocache/README.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/gocache/README.md b/gocache/README.md index cc5fff8..e06e42a 100644 --- a/gocache/README.md +++ b/gocache/README.md @@ -16,7 +16,7 @@ Currently you need to build your own Go toolchain to use this. As of 2023-04-24 First, build your cache child process. For example, ```sh -$ go install github.com/bradfitz/cmd/go-cacher +$ go install github.com/bradfitz/go-tool-cache/cmd/go-cacher ``` Then tell Go to use it: @@ -39,4 +39,4 @@ Run it again and watch the hit rate go up: $ GOCACHEPROG="$HOME/go/bin/go-cacher --verbose" go install std Defaulting to cache dir /home/bradfitz/.cache/go-cacher ... cacher: closing; 808 gets (808 hits, 0 misses, 0 errors); 0 puts (0 errors) -``` \ No newline at end of file +``` From fdc86ae8d3be50e0466c19e25a6965c5403c41ba Mon Sep 17 00:00:00 2001 From: j178 <10510431+j178@users.noreply.github.com> Date: Mon, 16 Dec 2024 14:15:52 +0800 Subject: [PATCH 08/67] Add `OutputID` for compatibility with Go1.24 Migrated-from: bradfitz/go-tool-cache@785169e9157ed70db072ac6c68d83d7a4d61dab5 --- gocache/cacheproc/cacheproc.go | 14 +++++++++----- gocache/cachers/disk.go | 16 ++++++++-------- gocache/wire/wire.go | 12 ++++++++---- 3 files changed, 25 insertions(+), 17 deletions(-) diff --git a/gocache/cacheproc/cacheproc.go b/gocache/cacheproc/cacheproc.go index da2b3f7..ed53e70 100644 --- a/gocache/cacheproc/cacheproc.go +++ b/gocache/cacheproc/cacheproc.go @@ -42,10 +42,10 @@ type Process struct { Get func(ctx context.Context, actionID string) (outputID, diskPath string, _ error) // Put optionally specifies a func to add something to the cache. - // The actionID and objectID is a lowercase hex string of unspecified format or length. + // The actionID and outputID is a lowercase hex string of unspecified format or length. // On success, diskPath must be the absolute path to a regular file. // If nil, cmd/go may write to disk elsewhere as needed. - Put func(ctx context.Context, actionID, objectID string, size int64, r io.Reader) (diskPath string, _ error) + Put func(ctx context.Context, actionID, outputID string, size int64, r io.Reader) (diskPath string, _ error) // Close optionally specifies a func to run when the cmd/go tool is // shutting down. @@ -94,6 +94,10 @@ func (p *Process) Run() error { } return err } + // For Go1.23 backward compatibility, remove in Go1.25. + if len(req.OutputID) == 0 && len(req.ObjectID) != 0 { + req.OutputID = req.ObjectID + } if req.Command == wire.CmdPut && req.BodySize > 0 { // TODO(bradfitz): stream this and pass a checksum-validating // io.Reader that validates on EOF. @@ -184,12 +188,12 @@ func (p *Process) handleGet(ctx context.Context, req *wire.Request, res *wire.Re } func (p *Process) handlePut(ctx context.Context, req *wire.Request, res *wire.Response) (retErr error) { - actionID, objectID := fmt.Sprintf("%x", req.ActionID), fmt.Sprintf("%x", req.ObjectID) + actionID, outputID := fmt.Sprintf("%x", req.ActionID), fmt.Sprintf("%x", req.OutputID) p.Puts.Add(1) defer func() { if retErr != nil { p.PutErrors.Add(1) - log.Printf("put(action %s, obj %s, %v bytes): %v", actionID, objectID, req.BodySize, retErr) + log.Printf("put(action %s, obj %s, %v bytes): %v", actionID, outputID, req.BodySize, retErr) } }() if p.Put == nil { @@ -202,7 +206,7 @@ func (p *Process) handlePut(ctx context.Context, req *wire.Request, res *wire.Re if body == nil { body = bytes.NewReader(nil) } - diskPath, err := p.Put(ctx, actionID, objectID, req.BodySize, body) + diskPath, err := p.Put(ctx, actionID, outputID, req.BodySize, body) if err != nil { return err } diff --git a/gocache/cachers/disk.go b/gocache/cachers/disk.go index 3259b80..85c203b 100644 --- a/gocache/cachers/disk.go +++ b/gocache/cachers/disk.go @@ -50,22 +50,22 @@ func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPa return ie.OutputID, filepath.Join(dc.Dir, fmt.Sprintf("o-%v", ie.OutputID)), nil } -func (dc *DiskCache) OutputFilename(objectID string) string { - if len(objectID) < 4 || len(objectID) > 1000 { +func (dc *DiskCache) OutputFilename(outputID string) string { + if len(outputID) < 4 || len(outputID) > 1000 { return "" } - for i := range objectID { - b := objectID[i] + for i := range outputID { + b := outputID[i] if b >= '0' && b <= '9' || b >= 'a' && b <= 'f' { continue } return "" } - return filepath.Join(dc.Dir, fmt.Sprintf("o-%s", objectID)) + return filepath.Join(dc.Dir, fmt.Sprintf("o-%s", outputID)) } -func (dc *DiskCache) Put(ctx context.Context, actionID, objectID string, size int64, body io.Reader) (diskPath string, _ error) { - file := filepath.Join(dc.Dir, fmt.Sprintf("o-%s", objectID)) +func (dc *DiskCache) Put(ctx context.Context, actionID, outputID string, size int64, body io.Reader) (diskPath string, _ error) { + file := filepath.Join(dc.Dir, fmt.Sprintf("o-%s", outputID)) // Special case empty files; they're both common and easier to do race-free. if size == 0 { @@ -86,7 +86,7 @@ func (dc *DiskCache) Put(ctx context.Context, actionID, objectID string, size in ij, err := json.Marshal(indexEntry{ Version: 1, - OutputID: objectID, + OutputID: outputID, Size: size, TimeNanos: time.Now().UnixNano(), }) diff --git a/gocache/wire/wire.go b/gocache/wire/wire.go index 36f4810..a3acf3c 100644 --- a/gocache/wire/wire.go +++ b/gocache/wire/wire.go @@ -38,8 +38,12 @@ type Request struct { // ActionID is non-nil for get and puts. ActionID []byte `json:",omitempty"` // or nil if not used - // ObjectID is set for Type "put" and "output-file". - ObjectID []byte `json:",omitempty"` // or nil if not used + // OutputID is set for Type "put" and "output-file". + OutputID []byte `json:",omitempty"` // or nil if not used + + // ObjectID is the name of `OutputID` before Go1.24, it will be removed in Go1.25. + // It's used for backward compatibility. + ObjectID []byte `json:",omitempty"` // Body is the body for "put" requests. It's sent after the JSON object // as a base64-encoded JSON string when BodySize is non-zero. @@ -81,8 +85,8 @@ type Response struct { Size int64 `json:",omitempty"` TimeNanos int64 `json:",omitempty"` // TODO(bradfitz): document - // DiskPath is the absolute path on disk of the ObjectID corresponding + // DiskPath is the absolute path on disk of the OutputID corresponding // a "get" request's ActionID (on cache hit) or a "put" request's - // provided ObjectID. + // provided OutputID. DiskPath string `json:",omitempty"` } From e6088681323e0bfeb07764077571c4240c97a6e2 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sat, 9 Aug 2025 07:58:32 -0700 Subject: [PATCH 09/67] README.md: update Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@1a57bfe70def78b75a94a86385f5524862dbc4bf --- gocache/README.md | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/gocache/README.md b/gocache/README.md index e06e42a..2034a80 100644 --- a/gocache/README.md +++ b/gocache/README.md @@ -1,22 +1,25 @@ # go-tool-cache -Like Go's built-in build/test caching but wish it weren't purely stored on local disk in the `$GOCACHE` directory? +Do you like Go's built-in build & test caching but wish it weren't purely stored on local disk in the `$GOCACHE` directory? Want to share your cache over the network between your various machines, coworkers, and CI runs without all that GitHub actions/caches tarring and untarring? -Along with a [modification to Go's `cmd/go` tool](https://go-review.googlesource.com/c/go/+/486715) ([open proposal](https://github.com/golang/go/issues/59719)), this repo lets you write -custom `GOCACHE` implementations to handle the cache however you'd like. +Go's [GOCACHEPROG](https://pkg.go.dev/cmd/go/internal/cacheprog) lets you do that! + +This was a demonstration repro for when GOCACHEPROG was still a +[proposal](https://github.com/golang/go/issues/59719). Now it just contains some misc +examples. ## Status -Currently you need to build your own Go toolchain to use this. As of 2023-04-24 it's still an open proposal & work in progress. +GOCACHEPROG shipped as an experiment in Go 1.21. It became official in Go 1.24. ## Using First, build your cache child process. For example, ```sh -$ go install github.com/bradfitz/go-tool-cache/cmd/go-cacher +$ go install github.com/bradfitz/go-tool-cache/cmd/go-cacher@latest ``` Then tell Go to use it: From 6ecf2ab9ad1b28469abee69db820be93d417063d Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sat, 9 Aug 2025 11:06:36 -0700 Subject: [PATCH 10/67] cmd/gocached: add start of good server w/ sqlite indexes, etc Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@4c504e681d3aa944f134019603e5b10d07dad20d --- cmd/gocached/gocached.go | 329 ++++++++++++++++++++++++++++++++++ cmd/gocached/gocached_test.go | 104 +++++++++++ go.mod | 19 +- go.sum | 49 +++++ gocache/cachers/disk.go | 13 +- gocache/cachers/http.go | 71 +++++--- 6 files changed, 559 insertions(+), 26 deletions(-) create mode 100644 cmd/gocached/gocached.go create mode 100644 cmd/gocached/gocached_test.go diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go new file mode 100644 index 0000000..b4f6c04 --- /dev/null +++ b/cmd/gocached/gocached.go @@ -0,0 +1,329 @@ +// Copyright (c) Tailscale Inc & AUTHORS +// SPDX-License-Identifier: BSD-3-Clause + +// The gocached daemon is an HTTP server daemon that go-cacher can hit. It does +// cache tiering and evicts old large things from disk, and can fetch metadata +// and object contents from peer cache servers. +// +// It uses sqlite (the pure Go modernc.org/sqlite driver) to store metadata and +// indexes. +// +/* + +It speaks the same protocol as go-cacher-server, but requires +the "Want-Object: 1" header variant on the GET request. + + GET /action/ + Want-Object: 1 + + 200 OK + Content-Type: application/octet-stream + Content-Length: 1234 + Go-Output-Id: xxxxxxxxxxx + + + +And to insert an object: + + PUT // + Content-Length: 1234 + + + +*/ +package main + +import ( + "context" + "database/sql" + "errors" + "expvar" + "flag" + "fmt" + "io" + "log" + "net/http" + "os" + "path/filepath" + "strings" + "time" + + "github.com/tailscale/tb/gocache/cachers" + _ "modernc.org/sqlite" +) + +// smallObjectSize is the maximum size of an object that we store inline in the +// database, rather than on disk. Empirically, about half of objects are 1KB or +// smaller. +const smallObjectSize = 1 << 10 + +var ( + dir = flag.String("cache-dir", "", "cache directory, if empty defaults to /gocached") + verbose = flag.Bool("verbose", false, "be verbose") + listen = flag.String("listen", ":31364", "listen address") +) + +func main() { + flag.Parse() + if *dir == "" { + d, err := os.UserCacheDir() + if err != nil { + log.Fatal(err) + } + d = filepath.Join(d, "gocached") + log.Printf("Defaulting to cache dir %v ...", d) + *dir = d + } + if err := os.MkdirAll(*dir, 0755); err != nil { + log.Fatal(err) + } + + srv, err := newServer(*dir) + if err != nil { + log.Fatalf("newServer: %v", err) + } + srv.verbose = *verbose + + log.Fatal(http.ListenAndServe(*listen, srv)) +} + +func newServer(dir string) (*server, error) { + db, err := openDB(dir) + if err != nil { + return nil, fmt.Errorf("openDB: %w", err) + } + dc := &cachers.DiskCache{Dir: dir} + return &server{ + db: db, + disk: dc, + logf: log.Printf, + }, nil +} + +const schemaVersion = 1 + +const schema = ` +PRAGMA journal_mode=WAL; +CREATE TABLE IF NOT EXISTS Actions ( + ActionID TEXT NOT NULL PRIMARY KEY, + OutputID TEXT NOT NULL, + OutputSize INTEGER NOT NULL, -- bytes of the output (even if stored off-DB) + CreateTime INTEGER NOT NULL, -- unix sec when inserted (locally or on a peer) + AccessTime INTEGER NOT NULL, -- unix sec of last access + InlineOutput BLOB, -- optional inline output value (e.g. for small output) + + CHECK (ActionID = lower(ActionID)), + CHECK (OutputID = lower(OutputID)), + CHECK (ActionID GLOB '[0-9a-f]*'), + CHECK (OutputID GLOB '[0-9a-f]*'), + CHECK (OutputSize >= 0), + CHECK (CreateTime >= 0), + CHECK (AccessTime >= 0), + CHECK (InlineOutput IS NULL OR length(InlineOutput) = OutputSize) +) STRICT; + +CREATE INDEX IF NOT EXISTS idx_actions_access ON Actions(AccessTime, OutputSize); +` + +func openDB(dbDir string) (*sql.DB, error) { + dbPath := filepath.Join(dbDir, fmt.Sprintf("gocached-v%d.db", schemaVersion)) + db, err := sql.Open("sqlite", dbPath) + if err != nil { + return nil, err + } + if _, err := db.Exec(schema); err != nil { + return nil, err + } + return db, nil +} + +type server struct { + db *sql.DB + disk *cachers.DiskCache // for large outputs only + verbose bool + logf func(format string, args ...any) + clock func() time.Time // if non-nil, alternate time.Now for testing + + // Metrics + gets expvar.Int + getHits expvar.Int +} + +func (s *server) now() time.Time { + if s.clock != nil { + return s.clock() + } + return time.Now() +} + +func (s *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if s.verbose { + s.logf("ServeHTTP: %s %s", r.Method, r.RequestURI) + } + if r.Method == "PUT" { + s.handlePut(w, r) + return + } + if r.Method != "GET" { + http.Error(w, "bad method", http.StatusBadRequest) + return + } + switch { + case strings.HasPrefix(r.URL.Path, "/action/"): + s.handleGetAction(w, r) + case r.URL.Path == "/": + io.WriteString(w, "hi") + default: + http.Error(w, "not found", http.StatusNotFound) + } +} + +func getHexSuffix(r *http.Request, prefix string) (hexSuffix string, ok bool) { + hexSuffix, _ = strings.CutPrefix(r.RequestURI, prefix) + if !validHex(hexSuffix) { + return "", false + } + return hexSuffix, true +} + +func validHex(x string) bool { + if len(x) < 4 || len(x) > 1000 || len(x)%2 == 1 { + return false + } + for i := range x { + b := x[i] + if b >= '0' && b <= '9' || b >= 'a' && b <= 'f' { + continue + } + return false + } + return true +} + +func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { + s.gets.Add(1) + ctx := r.Context() + + actionID, ok := getHexSuffix(r, "/action/") + if !ok { + http.Error(w, "bad request", http.StatusBadRequest) + return + } + if r.Header.Get("Want-Object") != "1" { + http.Error(w, "bad request: missing Want-Object header", http.StatusBadRequest) + return + } + + var outputID string + var size int64 + var inlineOutput sql.NullString + err := s.db.QueryRow("SELECT OutputID, OutputSize, InlineOutput FROM Actions WHERE ActionID = ?", actionID).Scan(&outputID, &size, &inlineOutput) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + http.Error(w, "not found", http.StatusNotFound) + return + } + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + s.getHits.Add(1) + + w.Header().Set("Content-Type", "application/octet-stream") + w.Header().Set("Content-Length", fmt.Sprint(size)) + w.Header().Set("Go-Output-Id", outputID) + + if r.Method == "HEAD" || size == 0 { + return + } + + if inlineOutput.Valid { + // For small outputs stored inline in the database, we can return them directly. + io.WriteString(w, inlineOutput.String) + return + } + + // Otherwise, for large objects that we know about, we can try to get them + // from our local disk or a peer. + rc, err := s.getObjectFromDiskOrPeer(ctx, actionID) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + if rc == nil { + http.Error(w, "not found", http.StatusNotFound) + return + } + defer rc.Close() + io.Copy(w, rc) +} + +// getObjectFromDiskOrPeer retrieves the object for the given actionID, either +// from disk or a peer. This is used after a local DB lookup discovers the content +// exists but is not stored in SQLite. +// +// It returns (nil, nil) on miss. +func (s *server) getObjectFromDiskOrPeer(ctx context.Context, actionID string) (rc io.ReadCloser, err error) { + _, diskPath, diskErr := s.disk.Get(ctx, actionID) + if diskErr != nil { + return nil, diskErr + } + if diskPath != "" { + f, err := os.Open(diskPath) + if err != nil { + return nil, err + } + return f, nil + } + // TODO(bradfitz): search peers, S3, etc. + // For now, just return nil, nil on miss. + return nil, nil +} + +func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + if r.Method != "PUT" { + http.Error(w, "bad method", http.StatusMethodNotAllowed) + return + } + actionID, outputID, ok := strings.Cut(r.RequestURI[len("/"):], "/") + if !ok || !validHex(actionID) || !validHex(outputID) { + http.Error(w, "bad URI", http.StatusBadRequest) + return + } + if r.ContentLength == -1 { + http.Error(w, "missing Content-Length", http.StatusBadRequest) + return + } + + var inline []byte + if r.ContentLength <= smallObjectSize { + // Store small objects inline in the database. + var err error + inline, err = io.ReadAll(r.Body) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + } else { + // For larger objects, we store them on disk. + _, err := s.disk.Put(ctx, actionID, outputID, r.ContentLength, r.Body) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + } + + // Insert or update the action in the database. + nowUnix := s.now().Unix() + _, err := s.db.Exec(` +INSERT OR IGNORE INTO Actions (ActionID, OutputID, OutputSize, CreateTime, AccessTime, InlineOutput) +VALUES (?, ?, ?, ?, ?, ?)`, + actionID, outputID, r.ContentLength, nowUnix, nowUnix, inline) + if err != nil { + s.logf("INSERT error: %v", err) + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + + w.WriteHeader(http.StatusNoContent) +} diff --git a/cmd/gocached/gocached_test.go b/cmd/gocached/gocached_test.go new file mode 100644 index 0000000..5571d6a --- /dev/null +++ b/cmd/gocached/gocached_test.go @@ -0,0 +1,104 @@ +package main + +import ( + "context" + "net/http/httptest" + "os" + "strings" + "testing" + + "github.com/tailscale/tb/gocache/cachers" +) + +func TestServer(t *testing.T) { + ctx := context.Background() + + dir := t.TempDir() + srv, err := newServer(dir) + if err != nil { + t.Fatalf("newServer: %v", err) + } + srv.logf = t.Logf + srv.verbose = true + + hs := httptest.NewServer(srv) + defer hs.Close() + + mkClient := func() *cachers.HTTPClient { + clientCacheDir := t.TempDir() + return &cachers.HTTPClient{ + BaseURL: hs.URL, + Disk: &cachers.DiskCache{ + Dir: clientCacheDir, + Logf: func(format string, args ...any) { + t.Logf("client-disk: "+format, args...) + }, + }, + } + } + + // Make two clients (imagine: two different builder VMs) + c1 := mkClient() + c2 := mkClient() + + const testActionID = "0001" + const testActionIDMiss = "0002" // this one doesn't exist + const testOutputID = "9900" + const testObjectValue = "test data" + + // Populate from the first client. + clientDiskPath, err := c1.Put(ctx, testActionID, testOutputID, int64(len(testObjectValue)), strings.NewReader(testObjectValue)) + if err != nil { + t.Fatalf("Put: %v", err) + } + if clientDiskPath == "" { + t.Fatal("Put returned empty disk path") + } + wrote, err := os.ReadFile(clientDiskPath) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(wrote) != testObjectValue { + t.Errorf("ReadFile got %q, want %q", wrote, testObjectValue) + } + + // Read from the second client. + gotOutputID, diskPath, err := c2.Get(ctx, testActionID) + if err != nil { + t.Fatalf("Get: %v", err) + } + if gotOutputID != testOutputID { + t.Errorf("Get got outputID %q, want %q", gotOutputID, testOutputID) + } + if diskPath == "" { + t.Fatal("Get returned empty disk path") + } + wrote, err = os.ReadFile(diskPath) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(wrote) != testObjectValue { + t.Errorf("ReadFile got %q, want %q", wrote, testObjectValue) + } + // Check metrics + if got, want := srv.gets.Value(), int64(1); got != want { + t.Errorf("server metric gets = %d, want %d", got, want) + } + if got, want := srv.getHits.Value(), int64(1); got != want { + t.Errorf("server metric getHits = %d, want %d", got, want) + } + + // Do the same get again from the same client. This shouldn't hit the network. + if _, _, err = c2.Get(ctx, testActionID); err != nil { + t.Fatalf("Get: %v", err) + } else if srv.gets.Value() != 1 { + t.Errorf("server metric gets = %d, want 1", srv.gets.Value()) + } + + // Cache miss. This should hit the network and fail. + if _, _, err = c2.Get(ctx, testActionIDMiss); err != nil { + t.Fatalf("miss Get: %v", err) + } else if srv.gets.Value() != 2 { + t.Errorf("server metric gets = %d, want 1", srv.gets.Value()) + } +} diff --git a/go.mod b/go.mod index 60b4913..bdaf75d 100644 --- a/go.mod +++ b/go.mod @@ -1,2 +1,19 @@ module github.com/tailscale/tb -go 1.20 +go 1.23.0 + +toolchain go1.23.4 + +require modernc.org/sqlite v1.38.2 + +require ( + github.com/dustin/go-humanize v1.0.1 // indirect + github.com/google/uuid v1.6.0 // indirect + github.com/mattn/go-isatty v0.0.20 // indirect + github.com/ncruces/go-strftime v0.1.9 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect + golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b // indirect + golang.org/x/sys v0.34.0 // indirect + modernc.org/libc v1.66.3 // indirect + modernc.org/mathutil v1.7.1 // indirect + modernc.org/memory v1.11.0 // indirect +) diff --git a/go.sum b/go.sum index e69de29..aac187a 100644 --- a/go.sum +++ b/go.sum @@ -0,0 +1,49 @@ +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= +github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4= +github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b h1:M2rDM6z3Fhozi9O7NWsxAkg/yqS/lQJ6PmkyIV3YP+o= +golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8= +golang.org/x/mod v0.25.0 h1:n7a+ZbQKQA/Ysbyb0/6IbB1H/X41mKgbhfv7AfG/44w= +golang.org/x/mod v0.25.0/go.mod h1:IXM97Txy2VM4PJ3gI61r1YEk/gAj6zAHN3AdZt6S9Ww= +golang.org/x/sync v0.15.0 h1:KWH3jNZsfyT6xfAfKiz6MRNmd46ByHDYaZ7KSkCtdW8= +golang.org/x/sync v0.15.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA= +golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +golang.org/x/sys v0.34.0 h1:H5Y5sJ2L2JRdyv7ROF1he/lPdvFsd0mJHFw2ThKHxLA= +golang.org/x/sys v0.34.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= +golang.org/x/tools v0.34.0 h1:qIpSLOxeCYGg9TrcJokLBG4KFA6d795g0xkBkiESGlo= +golang.org/x/tools v0.34.0/go.mod h1:pAP9OwEaY1CAW3HOmg3hLZC5Z0CCmzjAF2UQMSqNARg= +modernc.org/cc/v4 v4.26.2 h1:991HMkLjJzYBIfha6ECZdjrIYz2/1ayr+FL8GN+CNzM= +modernc.org/cc/v4 v4.26.2/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0= +modernc.org/ccgo/v4 v4.28.0 h1:rjznn6WWehKq7dG4JtLRKxb52Ecv8OUGah8+Z/SfpNU= +modernc.org/ccgo/v4 v4.28.0/go.mod h1:JygV3+9AV6SmPhDasu4JgquwU81XAKLd3OKTUDNOiKE= +modernc.org/fileutil v1.3.8 h1:qtzNm7ED75pd1C7WgAGcK4edm4fvhtBsEiI/0NQ54YM= +modernc.org/fileutil v1.3.8/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc= +modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI= +modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito= +modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks= +modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI= +modernc.org/libc v1.66.3 h1:cfCbjTUcdsKyyZZfEUKfoHcP3S0Wkvz3jgSzByEWVCQ= +modernc.org/libc v1.66.3/go.mod h1:XD9zO8kt59cANKvHPXpx7yS2ELPheAey0vjIuZOhOU8= +modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= +modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= +modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= +modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= +modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8= +modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns= +modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w= +modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE= +modernc.org/sqlite v1.38.2 h1:Aclu7+tgjgcQVShZqim41Bbw9Cho0y/7WzYptXqkEek= +modernc.org/sqlite v1.38.2/go.mod h1:cPTJYSlgg3Sfg046yBShXENNtPrWrDX8bsbAQBzgQ5E= +modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0= +modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A= +modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y= +modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM= diff --git a/gocache/cachers/disk.go b/gocache/cachers/disk.go index 85c203b..0d015e8 100644 --- a/gocache/cachers/disk.go +++ b/gocache/cachers/disk.go @@ -24,6 +24,15 @@ type indexEntry struct { type DiskCache struct { Dir string Verbose bool + Logf func(format string, args ...any) // optional alt logger +} + +func (dc *DiskCache) logf(format string, args ...any) { + if dc.Logf != nil { + dc.Logf(format, args...) + } else if dc.Verbose { + log.Printf(format, args...) + } } func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPath string, err error) { @@ -33,14 +42,14 @@ func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPa if os.IsNotExist(err) { err = nil if dc.Verbose { - log.Printf("disk miss: %v", actionID) + dc.logf("disk miss: %v", actionID) } } return "", "", err } var ie indexEntry if err := json.Unmarshal(ij, &ie); err != nil { - log.Printf("Warning: JSON error for action %q: %v", actionID, err) + dc.logf("Warning: JSON error for action %q: %v", actionID, err) return "", "", nil } if _, err := hex.DecodeString(ie.OutputID); err != nil { diff --git a/gocache/cachers/http.go b/gocache/cachers/http.go index 1eed0bd..7ab9b99 100644 --- a/gocache/cachers/http.go +++ b/gocache/cachers/http.go @@ -47,6 +47,13 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa req, _ := http.NewRequestWithContext(ctx, "GET", c.BaseURL+"/action/"+actionID, nil) + // Set a header to indicate we want the object and metadata in one response. + // Prior to 2025-08-09, the protocol was two separate requests. Rather than + // change this repo's protocol and potentially break existing clients, + // we just add a header to indicate we want the object and then we support + // both the old and new response types. + req.Header.Set("Want-Object", "1") // opt in to new single roundtrip protocol + res, err := c.httpClient().Do(req) if err != nil { return "", "", err @@ -58,35 +65,53 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa if res.StatusCode != http.StatusOK { return "", "", fmt.Errorf("unexpected GET /action/%s status %v", actionID, res.Status) } - var av ActionValue - if err := json.NewDecoder(res.Body).Decode(&av); err != nil { - return "", "", err - } - outputID = av.OutputID - // If not on disk, download it to disk. - var putBody io.Reader - if av.Size == 0 { - putBody = bytes.NewReader(nil) - } else { - req, _ = http.NewRequestWithContext(ctx, "GET", c.BaseURL+"/output/"+outputID, nil) - res, err = c.httpClient().Do(req) - if err != nil { - return "", "", err - } - defer res.Body.Close() - if res.StatusCode == http.StatusNotFound { - return "", "", nil - } - if res.StatusCode != http.StatusOK { - return "", "", fmt.Errorf("unexpected GET /output/%s status %v", outputID, res.Status) + switch res.Header.Get("Content-Type") { + default: + return "", "", fmt.Errorf("unexpected Content-Type %q from server", res.Header.Get("Content-Type")) + + case "application/octet-stream": // new single roundtrip protocol + outputID = res.Header.Get("Go-Output-Id") + if outputID == "" { + return "", "", fmt.Errorf("missing Go-Output-Id header in response") } if res.ContentLength == -1 { return "", "", fmt.Errorf("no Content-Length from server") } - putBody = res.Body + diskPath, err = c.Disk.Put(ctx, actionID, outputID, res.ContentLength, res.Body) + + case "application/json": // old two-hop protocol + var av ActionValue + if err := json.NewDecoder(res.Body).Decode(&av); err != nil { + return "", "", err + } + outputID = av.OutputID + + // If not on disk, download it to disk. + var putBody io.Reader + if av.Size == 0 { + putBody = bytes.NewReader(nil) + } else { + req, _ = http.NewRequestWithContext(ctx, "GET", c.BaseURL+"/output/"+outputID, nil) + res, err = c.httpClient().Do(req) + if err != nil { + return "", "", err + } + defer res.Body.Close() + if res.StatusCode == http.StatusNotFound { + return "", "", nil + } + if res.StatusCode != http.StatusOK { + return "", "", fmt.Errorf("unexpected GET /output/%s status %v", outputID, res.Status) + } + if res.ContentLength == -1 { + return "", "", fmt.Errorf("no Content-Length from server") + } + putBody = res.Body + } + diskPath, err = c.Disk.Put(ctx, actionID, outputID, av.Size, putBody) } - diskPath, err = c.Disk.Put(ctx, actionID, outputID, av.Size, putBody) + return outputID, diskPath, err } From 063779f8e28dba20400e5177e99b0089ae98c750 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sat, 9 Aug 2025 15:22:24 -0700 Subject: [PATCH 11/67] cmd/gocached: add more metrics Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@2db3b0f7e072268b3dc076a12b190e8ed777dfbb --- cmd/gocached/gocached.go | 35 ++++++++--- cmd/gocached/gocached_test.go | 108 ++++++++++++++++++++-------------- 2 files changed, 93 insertions(+), 50 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index b4f6c04..7e18768 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -145,8 +145,14 @@ type server struct { clock func() time.Time // if non-nil, alternate time.Now for testing // Metrics - gets expvar.Int - getHits expvar.Int + gets expvar.Int // gets = getHits + getErrs + implicit misses + getBytes expvar.Int + getHits expvar.Int + getHitsInline expvar.Int // includes getHits; subset of getHits stored in SQLite + getErrs expvar.Int // errors from GET requests + puts expvar.Int + putsBytes expvar.Int + putsInline expvar.Int } func (s *server) now() time.Time { @@ -172,7 +178,7 @@ func (s *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { case strings.HasPrefix(r.URL.Path, "/action/"): s.handleGetAction(w, r) case r.URL.Path == "/": - io.WriteString(w, "hi") + io.WriteString(w, "gocached") default: http.Error(w, "not found", http.StatusNotFound) } @@ -204,13 +210,18 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { s.gets.Add(1) ctx := r.Context() + httpErr := func(msg string, code int) { + http.Error(w, msg, code) + s.getErrs.Add(1) + } + actionID, ok := getHexSuffix(r, "/action/") if !ok { - http.Error(w, "bad request", http.StatusBadRequest) + httpErr("bad request", http.StatusBadRequest) return } if r.Header.Get("Want-Object") != "1" { - http.Error(w, "bad request: missing Want-Object header", http.StatusBadRequest) + httpErr("bad request: missing Want-Object header", http.StatusBadRequest) return } @@ -223,9 +234,10 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { http.Error(w, "not found", http.StatusNotFound) return } - http.Error(w, err.Error(), http.StatusInternalServerError) + httpErr("bad request: missing Want-Object header", http.StatusBadRequest) return } + s.getHits.Add(1) w.Header().Set("Content-Type", "application/octet-stream") @@ -238,6 +250,8 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { if inlineOutput.Valid { // For small outputs stored inline in the database, we can return them directly. + s.getHitsInline.Add(1) + s.getBytes.Add(size) io.WriteString(w, inlineOutput.String) return } @@ -246,13 +260,14 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { // from our local disk or a peer. rc, err := s.getObjectFromDiskOrPeer(ctx, actionID) if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) + httpErr(err.Error(), http.StatusInternalServerError) return } if rc == nil { http.Error(w, "not found", http.StatusNotFound) return } + s.getBytes.Add(size) defer rc.Close() io.Copy(w, rc) } @@ -325,5 +340,11 @@ VALUES (?, ?, ?, ?, ?, ?)`, return } + s.puts.Add(1) + s.putsBytes.Add(r.ContentLength) + if inline != nil { + s.putsInline.Add(1) + } + w.WriteHeader(http.StatusNoContent) } diff --git a/cmd/gocached/gocached_test.go b/cmd/gocached/gocached_test.go index 5571d6a..38c949d 100644 --- a/cmd/gocached/gocached_test.go +++ b/cmd/gocached/gocached_test.go @@ -2,6 +2,7 @@ package main import ( "context" + "expvar" "net/http/httptest" "os" "strings" @@ -37,68 +38,89 @@ func TestServer(t *testing.T) { } } + // wantMetric is a helper to check an expvar.Int metric and reset it + // for future tests. + wantMetric := func(m *expvar.Int, want int64) { + t.Helper() + if got := m.Value(); got != want { + t.Errorf("metric = %d, want %d", got, want) + } + m.Set(0) + } + // Make two clients (imagine: two different builder VMs) c1 := mkClient() c2 := mkClient() const testActionID = "0001" const testActionIDMiss = "0002" // this one doesn't exist + const testActionIDBig = "0bbb" // non-inline object const testOutputID = "9900" + const testOutputIDBig = "9bbb" const testObjectValue = "test data" + testObjectValueBig := strings.Repeat("x", smallObjectSize+1) - // Populate from the first client. - clientDiskPath, err := c1.Put(ctx, testActionID, testOutputID, int64(len(testObjectValue)), strings.NewReader(testObjectValue)) - if err != nil { - t.Fatalf("Put: %v", err) - } - if clientDiskPath == "" { - t.Fatal("Put returned empty disk path") - } - wrote, err := os.ReadFile(clientDiskPath) - if err != nil { - t.Fatalf("ReadFile: %v", err) + wantPut := func(c *cachers.HTTPClient, actionID, outputID string, val string) { + t.Helper() + clientDiskPath, err := c.Put(ctx, actionID, outputID, int64(len(val)), strings.NewReader(val)) + if err != nil { + t.Fatalf("Put: %v", err) + } + if clientDiskPath == "" { + t.Fatal("Put returned empty disk path") + } + wantMetric(&srv.puts, 1) + wrote, err := os.ReadFile(clientDiskPath) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(wrote) != val { + t.Errorf("ReadFile got %q, want %q", wrote, val) + } } - if string(wrote) != testObjectValue { - t.Errorf("ReadFile got %q, want %q", wrote, testObjectValue) + + wantGet := func(c *cachers.HTTPClient, actionID, outputID, wantVal string) { + t.Helper() + gotOutputID, diskPath, err := c.Get(ctx, actionID) + if err != nil { + t.Fatalf("Get: %v", err) + } + if gotOutputID != outputID { + t.Errorf("Get got outputID %q, want %q", gotOutputID, outputID) + } + if diskPath == "" { + t.Fatal("Get returned empty disk path") + } + wrote, err := os.ReadFile(diskPath) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if string(wrote) != wantVal { + t.Errorf("ReadFile got %q, want %q", wrote, wantVal) + } } + // Populate from the first client. + wantPut(c1, testActionID, testOutputID, testObjectValue) + wantPut(c1, testActionIDBig, testOutputIDBig, testObjectValueBig) + // Read from the second client. - gotOutputID, diskPath, err := c2.Get(ctx, testActionID) - if err != nil { - t.Fatalf("Get: %v", err) - } - if gotOutputID != testOutputID { - t.Errorf("Get got outputID %q, want %q", gotOutputID, testOutputID) - } - if diskPath == "" { - t.Fatal("Get returned empty disk path") - } - wrote, err = os.ReadFile(diskPath) - if err != nil { - t.Fatalf("ReadFile: %v", err) - } - if string(wrote) != testObjectValue { - t.Errorf("ReadFile got %q, want %q", wrote, testObjectValue) - } + wantGet(c2, testActionID, testOutputID, testObjectValue) + wantGet(c2, testActionIDBig, testOutputIDBig, testObjectValueBig) + // Check metrics - if got, want := srv.gets.Value(), int64(1); got != want { - t.Errorf("server metric gets = %d, want %d", got, want) - } - if got, want := srv.getHits.Value(), int64(1); got != want { - t.Errorf("server metric getHits = %d, want %d", got, want) - } + wantMetric(&srv.gets, 2) + wantMetric(&srv.getHits, 2) + wantMetric(&srv.getHitsInline, 1) // Do the same get again from the same client. This shouldn't hit the network. - if _, _, err = c2.Get(ctx, testActionID); err != nil { - t.Fatalf("Get: %v", err) - } else if srv.gets.Value() != 1 { - t.Errorf("server metric gets = %d, want 1", srv.gets.Value()) - } + wantGet(c2, testActionID, testOutputID, testObjectValue) + wantMetric(&srv.gets, 0) // Cache miss. This should hit the network and fail. if _, _, err = c2.Get(ctx, testActionIDMiss); err != nil { t.Fatalf("miss Get: %v", err) - } else if srv.gets.Value() != 2 { - t.Errorf("server metric gets = %d, want 1", srv.gets.Value()) } + wantMetric(&srv.gets, 1) + wantMetric(&srv.getHits, 0) } From 0271c5bcb065eacd7ed51c8a38f22d7a277cd1ca Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sat, 9 Aug 2025 19:03:16 -0700 Subject: [PATCH 12/67] cmd/gocached: bump access time relatime-style on get Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@d997e5f3a456e9482c73cf13f70d6176233a35c3 --- cmd/gocached/gocached.go | 37 ++++++++++++++++++++++++++--------- cmd/gocached/gocached_test.go | 25 +++++++++++++++++++++++ 2 files changed, 53 insertions(+), 9 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index 7e18768..da4933c 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -145,14 +145,15 @@ type server struct { clock func() time.Time // if non-nil, alternate time.Now for testing // Metrics - gets expvar.Int // gets = getHits + getErrs + implicit misses - getBytes expvar.Int - getHits expvar.Int - getHitsInline expvar.Int // includes getHits; subset of getHits stored in SQLite - getErrs expvar.Int // errors from GET requests - puts expvar.Int - putsBytes expvar.Int - putsInline expvar.Int + gets expvar.Int // gets = getHits + getErrs + implicit misses + getBytes expvar.Int + getHits expvar.Int + getAccessBumps expvar.Int // number of times we updated the AccessTime + getHitsInline expvar.Int // includes getHits; subset of getHits stored in SQLite + getErrs expvar.Int // errors from GET requests + puts expvar.Int + putsBytes expvar.Int + putsInline expvar.Int } func (s *server) now() time.Time { @@ -206,6 +207,10 @@ func validHex(x string) bool { return true } +// relAtimeSeconds is how old an access time needs to be before +// we do a DB write to update it. +const relAtimeSeconds = 60 * 60 * 24 // 1 day + func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { s.gets.Add(1) ctx := r.Context() @@ -228,7 +233,8 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { var outputID string var size int64 var inlineOutput sql.NullString - err := s.db.QueryRow("SELECT OutputID, OutputSize, InlineOutput FROM Actions WHERE ActionID = ?", actionID).Scan(&outputID, &size, &inlineOutput) + var accessTime int64 + err := s.db.QueryRow("SELECT OutputID, OutputSize, InlineOutput, AccessTime FROM Actions WHERE ActionID = ?", actionID).Scan(&outputID, &size, &inlineOutput, &accessTime) if err != nil { if errors.Is(err, sql.ErrNoRows) { http.Error(w, "not found", http.StatusNotFound) @@ -238,6 +244,19 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { return } + // If it's been more than a day since the last access, update the access time. + // This is similar to the Linux "relatime" behavior. + now := s.now().Unix() + if accessTime < now-relAtimeSeconds { + _, err := s.db.Exec("UPDATE Actions SET AccessTime = ? WHERE ActionID = ?", now, actionID) + if err != nil { + s.logf("Update AccessTime error: %v", err) + httpErr("internal server error", http.StatusInternalServerError) + return + } + s.getAccessBumps.Add(1) + } + s.getHits.Add(1) w.Header().Set("Content-Type", "application/octet-stream") diff --git a/cmd/gocached/gocached_test.go b/cmd/gocached/gocached_test.go index 38c949d..077107d 100644 --- a/cmd/gocached/gocached_test.go +++ b/cmd/gocached/gocached_test.go @@ -6,7 +6,9 @@ import ( "net/http/httptest" "os" "strings" + "sync" "testing" + "time" "github.com/tailscale/tb/gocache/cachers" ) @@ -22,6 +24,21 @@ func TestServer(t *testing.T) { srv.logf = t.Logf srv.verbose = true + var ( + timeMu sync.Mutex + now = time.Unix(1234, 0) + ) + srv.clock = func() time.Time { + timeMu.Lock() + defer timeMu.Unlock() + return now + } + advanceClock := func(d time.Duration) { + timeMu.Lock() + defer timeMu.Unlock() + now = now.Add(d) + } + hs := httptest.NewServer(srv) defer hs.Close() @@ -123,4 +140,12 @@ func TestServer(t *testing.T) { } wantMetric(&srv.gets, 1) wantMetric(&srv.getHits, 0) + + // Check that access time gets updated. + // Do it from a fresh client without a disk cache. + wantMetric(&srv.getAccessBumps, 0) + advanceClock(relAtimeSeconds * 2 * time.Second) // advance clock by 2 days + c3 := mkClient() + wantGet(c3, testActionID, testOutputID, testObjectValue) + wantMetric(&srv.getAccessBumps, 1) } From ba6e502bdc7b379bf25c7e821b49be95b516ffea Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sun, 10 Aug 2025 07:47:29 -0700 Subject: [PATCH 13/67] cmd/gocached: add Prometheus /metrics handler Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@e221bff60ac0a2745839b85e566606b5bbe53493 --- cmd/gocached/gocached.go | 181 +++++++++++++++++++++++++++------- cmd/gocached/gocached_test.go | 18 ++-- go.mod | 12 ++- go.sum | 32 ++++++ 4 files changed, 195 insertions(+), 48 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index da4933c..959d969 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -45,9 +45,14 @@ import ( "net/http" "os" "path/filepath" + "reflect" "strings" "time" + "github.com/prometheus/client_golang/prometheus" + "github.com/prometheus/client_golang/prometheus/collectors" + "github.com/prometheus/client_golang/prometheus/promhttp" + dto "github.com/prometheus/client_model/go" "github.com/tailscale/tb/gocache/cachers" _ "modernc.org/sqlite" ) @@ -84,22 +89,10 @@ func main() { } srv.verbose = *verbose + log.Printf("gocached listening on %s ...", *listen) log.Fatal(http.ListenAndServe(*listen, srv)) } -func newServer(dir string) (*server, error) { - db, err := openDB(dir) - if err != nil { - return nil, fmt.Errorf("openDB: %w", err) - } - dc := &cachers.DiskCache{Dir: dir} - return &server{ - db: db, - disk: dc, - logf: log.Printf, - }, nil -} - const schemaVersion = 1 const schema = ` @@ -137,23 +130,88 @@ func openDB(dbDir string) (*sql.DB, error) { return db, nil } +func newServer(dir string) (*server, error) { + db, err := openDB(dir) + if err != nil { + return nil, fmt.Errorf("openDB: %w", err) + } + dc := &cachers.DiskCache{Dir: dir} + + reg := prometheus.NewRegistry() + reg.MustRegister( + collectors.NewGoCollector(), + collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}), + collectors.NewBuildInfoCollector(), + ) + + srv := &server{ + db: db, + disk: dc, + logf: log.Printf, + } + srv.registerMetrics(reg) + + srv.metricsHandler = promhttp.HandlerFor(reg, promhttp.HandlerOpts{ + ErrorLog: log.Default(), + }) + + return srv, nil +} + +func (s *server) registerMetrics(reg *prometheus.Registry) { + rv := reflect.ValueOf(s).Elem() + t := reflect.TypeOf(s).Elem() + for i := 0; i < t.NumField(); i++ { + sf := t.Field(i) + if sf.Type == reflect.TypeFor[expvar.Int]() { + expvarInt := rv.Field(i).Addr().Interface().(*expvar.Int) + typ := sf.Tag.Get("type") + name := sf.Tag.Get("name") + if typ == "" { + panic("missing type tag for " + sf.Name) + } + if name == "" { + panic("missing name tag for " + sf.Name) + } + help := sf.Tag.Get("help") + metricName := "gocached_" + name + + if tag := sf.Tag.Get("type"); tag != "" { + if tag == "gauge" { + reg.MustRegister(singleMetricCollector{&expvarGaugeMetric{ + desc: prometheus.NewDesc(metricName, help, nil, nil), + v: expvarInt, + }}) + } else if tag == "counter" { + reg.MustRegister(singleMetricCollector{&expvarCounterMetric{ + desc: prometheus.NewDesc(metricName, help, nil, nil), + v: expvarInt, + }}) + } + } + } + } +} + type server struct { - db *sql.DB - disk *cachers.DiskCache // for large outputs only - verbose bool - logf func(format string, args ...any) - clock func() time.Time // if non-nil, alternate time.Now for testing + db *sql.DB + disk *cachers.DiskCache // for large outputs only + verbose bool + logf func(format string, args ...any) + clock func() time.Time // if non-nil, alternate time.Now for testing + metricsHandler http.Handler // Metrics - gets expvar.Int // gets = getHits + getErrs + implicit misses - getBytes expvar.Int - getHits expvar.Int - getAccessBumps expvar.Int // number of times we updated the AccessTime - getHitsInline expvar.Int // includes getHits; subset of getHits stored in SQLite - getErrs expvar.Int // errors from GET requests - puts expvar.Int - putsBytes expvar.Int - putsInline expvar.Int + ActiveGets expvar.Int `type:"gauge" name:"active_gets" help:"currently pending get requests; should usually be zero"` + Gets expvar.Int `type:"counter" name:"gets" help:"total number of gocache get requests"` // gets = getHits + getErrs + implicit misses + GetBytes expvar.Int `type:"counter" name:"get_bytes" help:"total bytes fetched from gocache gets that were cache hits"` + GetHits expvar.Int `type:"counter" name:"get_hits" help:"total number of successful gocache get requests"` + GetAccessBumps expvar.Int `type:"counter" name:"get_access_bumps" help:"number of times a get request updated the access time of object"` + GetHitsInline expvar.Int `type:"counter" name:"get_hits_inline" help:"cache hits served from inline database storage (small objects)"` + GetErrs expvar.Int `type:"counter" name:"get_errs" help:"number of gocache get request errors"` + Puts expvar.Int `type:"counter" name:"puts" help:"total number of gocache put requests"` + PutsBytes expvar.Int `type:"counter" name:"put_bytes" help:"total bytes added from gocache puts"` + PutsInline expvar.Int `type:"counter" name:"put_inline" help:"subset of gocached_puts that were stored inline (small objects)"` } func (s *server) now() time.Time { @@ -171,7 +229,7 @@ func (s *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { s.handlePut(w, r) return } - if r.Method != "GET" { + if r.Method != "GET" && r.Method != "HEAD" { http.Error(w, "bad method", http.StatusBadRequest) return } @@ -180,6 +238,8 @@ func (s *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { s.handleGetAction(w, r) case r.URL.Path == "/": io.WriteString(w, "gocached") + case r.URL.Path == "/metrics": + s.metricsHandler.ServeHTTP(w, r) default: http.Error(w, "not found", http.StatusNotFound) } @@ -212,12 +272,15 @@ func validHex(x string) bool { const relAtimeSeconds = 60 * 60 * 24 // 1 day func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { - s.gets.Add(1) + s.ActiveGets.Add(1) + defer s.ActiveGets.Add(-1) + + s.Gets.Add(1) ctx := r.Context() httpErr := func(msg string, code int) { http.Error(w, msg, code) - s.getErrs.Add(1) + s.GetErrs.Add(1) } actionID, ok := getHexSuffix(r, "/action/") @@ -254,10 +317,10 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { httpErr("internal server error", http.StatusInternalServerError) return } - s.getAccessBumps.Add(1) + s.GetAccessBumps.Add(1) } - s.getHits.Add(1) + s.GetHits.Add(1) w.Header().Set("Content-Type", "application/octet-stream") w.Header().Set("Content-Length", fmt.Sprint(size)) @@ -269,8 +332,8 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { if inlineOutput.Valid { // For small outputs stored inline in the database, we can return them directly. - s.getHitsInline.Add(1) - s.getBytes.Add(size) + s.GetHitsInline.Add(1) + s.GetBytes.Add(size) io.WriteString(w, inlineOutput.String) return } @@ -286,7 +349,7 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { http.Error(w, "not found", http.StatusNotFound) return } - s.getBytes.Add(size) + s.GetBytes.Add(size) defer rc.Close() io.Copy(w, rc) } @@ -359,11 +422,53 @@ VALUES (?, ?, ?, ?, ?, ?)`, return } - s.puts.Add(1) - s.putsBytes.Add(r.ContentLength) + s.Puts.Add(1) + s.PutsBytes.Add(r.ContentLength) if inline != nil { - s.putsInline.Add(1) + s.PutsInline.Add(1) } w.WriteHeader(http.StatusNoContent) } + +// expvarCounterMetric is a Prometheus counter metric backed by an expvar.Int. +type expvarCounterMetric struct { + desc *prometheus.Desc + v *expvar.Int +} + +var _ prometheus.Metric = (*expvarCounterMetric)(nil) + +func (m *expvarCounterMetric) Desc() *prometheus.Desc { return m.desc } + +func (m *expvarCounterMetric) Write(out *dto.Metric) error { + val := float64(m.v.Value()) + out.Counter = &dto.Counter{Value: &val} + return nil +} + +// expvarGaugeMetric is a Prometheus gauge metric backed by an expvar.Int. +type expvarGaugeMetric struct { + desc *prometheus.Desc + v *expvar.Int +} + +var _ prometheus.Metric = (*expvarGaugeMetric)(nil) + +func (m *expvarGaugeMetric) Desc() *prometheus.Desc { return m.desc } + +func (m *expvarGaugeMetric) Write(out *dto.Metric) error { + val := float64(m.v.Value()) + out.Gauge = &dto.Gauge{Value: &val} + return nil +} + +// singleMetricCollector is a Prometheus collector that collects a single metric. +type singleMetricCollector struct { + metric prometheus.Metric +} + +var _ prometheus.Collector = singleMetricCollector{} + +func (c singleMetricCollector) Describe(ch chan<- *prometheus.Desc) { ch <- c.metric.Desc() } +func (c singleMetricCollector) Collect(ch chan<- prometheus.Metric) { ch <- c.metric } diff --git a/cmd/gocached/gocached_test.go b/cmd/gocached/gocached_test.go index 077107d..ba2ac6e 100644 --- a/cmd/gocached/gocached_test.go +++ b/cmd/gocached/gocached_test.go @@ -86,7 +86,7 @@ func TestServer(t *testing.T) { if clientDiskPath == "" { t.Fatal("Put returned empty disk path") } - wantMetric(&srv.puts, 1) + wantMetric(&srv.Puts, 1) wrote, err := os.ReadFile(clientDiskPath) if err != nil { t.Fatalf("ReadFile: %v", err) @@ -126,26 +126,26 @@ func TestServer(t *testing.T) { wantGet(c2, testActionIDBig, testOutputIDBig, testObjectValueBig) // Check metrics - wantMetric(&srv.gets, 2) - wantMetric(&srv.getHits, 2) - wantMetric(&srv.getHitsInline, 1) + wantMetric(&srv.Gets, 2) + wantMetric(&srv.GetHits, 2) + wantMetric(&srv.GetHitsInline, 1) // Do the same get again from the same client. This shouldn't hit the network. wantGet(c2, testActionID, testOutputID, testObjectValue) - wantMetric(&srv.gets, 0) + wantMetric(&srv.Gets, 0) // Cache miss. This should hit the network and fail. if _, _, err = c2.Get(ctx, testActionIDMiss); err != nil { t.Fatalf("miss Get: %v", err) } - wantMetric(&srv.gets, 1) - wantMetric(&srv.getHits, 0) + wantMetric(&srv.Gets, 1) + wantMetric(&srv.GetHits, 0) // Check that access time gets updated. // Do it from a fresh client without a disk cache. - wantMetric(&srv.getAccessBumps, 0) + wantMetric(&srv.GetAccessBumps, 0) advanceClock(relAtimeSeconds * 2 * time.Second) // advance clock by 2 days c3 := mkClient() wantGet(c3, testActionID, testOutputID, testObjectValue) - wantMetric(&srv.getAccessBumps, 1) + wantMetric(&srv.GetAccessBumps, 1) } diff --git a/go.mod b/go.mod index bdaf75d..2e9738d 100644 --- a/go.mod +++ b/go.mod @@ -3,16 +3,26 @@ go 1.23.0 toolchain go1.23.4 -require modernc.org/sqlite v1.38.2 +require ( + github.com/prometheus/client_golang v1.23.0 + github.com/prometheus/client_model v0.6.2 + modernc.org/sqlite v1.38.2 +) require ( + github.com/beorn7/perks v1.0.1 // indirect + github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/google/uuid v1.6.0 // indirect github.com/mattn/go-isatty v0.0.20 // indirect + github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect github.com/ncruces/go-strftime v0.1.9 // indirect + github.com/prometheus/common v0.65.0 // indirect + github.com/prometheus/procfs v0.16.1 // indirect github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b // indirect golang.org/x/sys v0.34.0 // indirect + google.golang.org/protobuf v1.36.6 // indirect modernc.org/libc v1.66.3 // indirect modernc.org/mathutil v1.7.1 // indirect modernc.org/memory v1.11.0 // indirect diff --git a/go.sum b/go.sum index aac187a..717c8c5 100644 --- a/go.sum +++ b/go.sum @@ -1,15 +1,43 @@ +github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= +github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo= +github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ= +github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc= +github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= +github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4= github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/prometheus/client_golang v1.23.0 h1:ust4zpdl9r4trLY/gSjlm07PuiBq2ynaXXlptpfy8Uc= +github.com/prometheus/client_golang v1.23.0/go.mod h1:i/o0R9ByOnHX0McrTMTyhYvKE4haaf2mW08I+jGAjEE= +github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk= +github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE= +github.com/prometheus/common v0.65.0 h1:QDwzd+G1twt//Kwj/Ww6E9FQq1iVMmODnILtW1t2VzE= +github.com/prometheus/common v0.65.0/go.mod h1:0gZns+BLRQ3V6NdaerOhMbwwRbNh9hkGINtQAsP5GS8= +github.com/prometheus/procfs v0.16.1 h1:hZ15bTNuirocR6u0JZ6BAHHmwS1p8B4P6MRqxtzMyRg= +github.com/prometheus/procfs v0.16.1/go.mod h1:teAbpZRB1iIAJYREa1LsoWUXykVXA1KlTmWl8x/U+Is= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA= +github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= +go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= +go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b h1:M2rDM6z3Fhozi9O7NWsxAkg/yqS/lQJ6PmkyIV3YP+o= golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8= golang.org/x/mod v0.25.0 h1:n7a+ZbQKQA/Ysbyb0/6IbB1H/X41mKgbhfv7AfG/44w= @@ -21,6 +49,10 @@ golang.org/x/sys v0.34.0 h1:H5Y5sJ2L2JRdyv7ROF1he/lPdvFsd0mJHFw2ThKHxLA= golang.org/x/sys v0.34.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= golang.org/x/tools v0.34.0 h1:qIpSLOxeCYGg9TrcJokLBG4KFA6d795g0xkBkiESGlo= golang.org/x/tools v0.34.0/go.mod h1:pAP9OwEaY1CAW3HOmg3hLZC5Z0CCmzjAF2UQMSqNARg= +google.golang.org/protobuf v1.36.6 h1:z1NpPI8ku2WgiWnf+t9wTPsn6eP1L7ksHUlkfLvd9xY= +google.golang.org/protobuf v1.36.6/go.mod h1:jduwjTPXsFjZGTmRluh+L6NjiWu7pchiJ2/5YcXBHnY= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= modernc.org/cc/v4 v4.26.2 h1:991HMkLjJzYBIfha6ECZdjrIYz2/1ayr+FL8GN+CNzM= modernc.org/cc/v4 v4.26.2/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0= modernc.org/ccgo/v4 v4.28.0 h1:rjznn6WWehKq7dG4JtLRKxb52Ecv8OUGah8+Z/SfpNU= From 8e09c0da06822e535f30503d79a5813333010d58 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sun, 10 Aug 2025 08:29:10 -0700 Subject: [PATCH 14/67] cmd/gocached: set busy timeout on sqlite Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@1471a41c1bb8a96bbd1b9f287eda43f1c6d67e04 --- cmd/gocached/gocached.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index 959d969..9af13fe 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -120,7 +120,7 @@ CREATE INDEX IF NOT EXISTS idx_actions_access ON Actions(AccessTime, OutputSize) func openDB(dbDir string) (*sql.DB, error) { dbPath := filepath.Join(dbDir, fmt.Sprintf("gocached-v%d.db", schemaVersion)) - db, err := sql.Open("sqlite", dbPath) + db, err := sql.Open("sqlite", "file:"+dbPath+"?_pragma=busy_timeout(5000)") if err != nil { return nil, err } From f7da204748eda374889b54924a6c1f334b229c0a Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sun, 10 Aug 2025 08:29:50 -0700 Subject: [PATCH 15/67] cachers: forest-ify disk cache like Go does Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@cf8026fb5682c63ce60c630b4f2b19739d4d181a --- gocache/cachers/disk.go | 49 ++++++++++++++++++++++++++++++++++++----- 1 file changed, 43 insertions(+), 6 deletions(-) diff --git a/gocache/cachers/disk.go b/gocache/cachers/disk.go index 0d015e8..790360d 100644 --- a/gocache/cachers/disk.go +++ b/gocache/cachers/disk.go @@ -3,7 +3,6 @@ package cachers import ( "bytes" "context" - "encoding/hex" "encoding/json" "fmt" "io" @@ -36,7 +35,12 @@ func (dc *DiskCache) logf(format string, args ...any) { } func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPath string, err error) { - actionFile := filepath.Join(dc.Dir, fmt.Sprintf("a-%s", actionID)) + if !validHex(actionID) { + return "", "", fmt.Errorf("actionID must be valid hex strings") + } + action2 := actionID[:2] + + actionFile := filepath.Join(dc.Dir, action2, fmt.Sprintf("a-%s", actionID)) ij, err := os.ReadFile(actionFile) if err != nil { if os.IsNotExist(err) { @@ -52,11 +56,12 @@ func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPa dc.logf("Warning: JSON error for action %q: %v", actionID, err) return "", "", nil } - if _, err := hex.DecodeString(ie.OutputID); err != nil { + if !validHex(ie.OutputID) { // Protect against malicious non-hex OutputID on disk return "", "", nil } - return ie.OutputID, filepath.Join(dc.Dir, fmt.Sprintf("o-%v", ie.OutputID)), nil + output2 := ie.OutputID[:2] + return ie.OutputID, filepath.Join(dc.Dir, output2, fmt.Sprintf("o-%v", ie.OutputID)), nil } func (dc *DiskCache) OutputFilename(outputID string) string { @@ -73,8 +78,40 @@ func (dc *DiskCache) OutputFilename(outputID string) string { return filepath.Join(dc.Dir, fmt.Sprintf("o-%s", outputID)) } +func validHex(x string) bool { + if len(x) < 4 || len(x) > 100 { + return false + } + for i := range len(x) { + b := x[i] + if b >= '0' && b <= '9' || b >= 'a' && b <= 'f' { + continue + } + return false + } + return true +} + func (dc *DiskCache) Put(ctx context.Context, actionID, outputID string, size int64, body io.Reader) (diskPath string, _ error) { - file := filepath.Join(dc.Dir, fmt.Sprintf("o-%s", outputID)) + if len(actionID) < 4 || len(outputID) < 4 { + return "", fmt.Errorf("actionID and outputID must be at least 4 characters long") + } + if !validHex(actionID) || !validHex(outputID) { + return "", fmt.Errorf("actionID and outputID must be valid hex strings") + } + + action2, output2 := actionID[:2], outputID[:2] + outputDir := filepath.Join(dc.Dir, output2) + actionDir := filepath.Join(dc.Dir, action2) + + if err := os.MkdirAll(actionDir, 0755); err != nil { + return "", fmt.Errorf("failed to create action directory: %w", err) + } + if err := os.MkdirAll(outputDir, 0755); err != nil { + return "", fmt.Errorf("failed to create output directory: %w", err) + } + + file := filepath.Join(outputDir, fmt.Sprintf("o-%s", outputID)) // Special case empty files; they're both common and easier to do race-free. if size == 0 { @@ -102,7 +139,7 @@ func (dc *DiskCache) Put(ctx context.Context, actionID, outputID string, size in if err != nil { return "", err } - actionFile := filepath.Join(dc.Dir, fmt.Sprintf("a-%s", actionID)) + actionFile := filepath.Join(actionDir, fmt.Sprintf("a-%s", actionID)) if _, err := writeAtomic(actionFile, bytes.NewReader(ij)); err != nil { return "", err } From 418bac071ec26ce56a2b05beb5bac809100d89fd Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sun, 10 Aug 2025 12:32:30 -0700 Subject: [PATCH 16/67] cmd/gocached: optimize storage of zero-sized objects Zero-sized objects are common and worth optimizing a bit. Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@77ecb3b4efda34ca8f06ab888e8ec2734d1fe23e --- cmd/gocached/gocached.go | 26 +++++++++++++++++++++++--- cmd/gocached/gocached_test.go | 8 ++++++-- 2 files changed, 29 insertions(+), 5 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index 9af13fe..da444c0 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -93,7 +93,7 @@ func main() { log.Fatal(http.ListenAndServe(*listen, srv)) } -const schemaVersion = 1 +const schemaVersion = 2 const schema = ` PRAGMA journal_mode=WAL; @@ -108,7 +108,7 @@ CREATE TABLE IF NOT EXISTS Actions ( CHECK (ActionID = lower(ActionID)), CHECK (OutputID = lower(OutputID)), CHECK (ActionID GLOB '[0-9a-f]*'), - CHECK (OutputID GLOB '[0-9a-f]*'), + CHECK (OutputID GLOB '[0-9a-f]*' OR OutputID GLOB 'wk*'), CHECK (OutputSize >= 0), CHECK (CreateTime >= 0), CHECK (AccessTime >= 0), @@ -306,6 +306,7 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { httpErr("bad request: missing Want-Object header", http.StatusBadRequest) return } + outputID = outputIDFromWellknown(outputID) // If it's been more than a day since the last access, update the access time. // This is similar to the Linux "relatime" behavior. @@ -415,7 +416,7 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { _, err := s.db.Exec(` INSERT OR IGNORE INTO Actions (ActionID, OutputID, OutputSize, CreateTime, AccessTime, InlineOutput) VALUES (?, ?, ?, ?, ?, ?)`, - actionID, outputID, r.ContentLength, nowUnix, nowUnix, inline) + actionID, outputIDOrWellKnown(outputID), r.ContentLength, nowUnix, nowUnix, inline) if err != nil { s.logf("INSERT error: %v", err) http.Error(w, err.Error(), http.StatusInternalServerError) @@ -431,6 +432,25 @@ VALUES (?, ?, ?, ?, ?, ?)`, w.WriteHeader(http.StatusNoContent) } +// sha256OfEmpty is the SHA-256 hash of an empty string, used as a well-known +// value in SQLite to store bytes, as it's common. We store it in SQLite +// as "wk0" (for "well known 0") to save space, as it's common. +const sha256OfEmpty = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" + +func outputIDOrWellKnown(id string) string { + if id == sha256OfEmpty { + return "wk0" + } + return id +} + +func outputIDFromWellknown(idStored string) string { + if idStored == "wk0" { + return sha256OfEmpty + } + return idStored +} + // expvarCounterMetric is a Prometheus counter metric backed by an expvar.Int. type expvarCounterMetric struct { desc *prometheus.Desc diff --git a/cmd/gocached/gocached_test.go b/cmd/gocached/gocached_test.go index ba2ac6e..fda6fe1 100644 --- a/cmd/gocached/gocached_test.go +++ b/cmd/gocached/gocached_test.go @@ -72,8 +72,10 @@ func TestServer(t *testing.T) { const testActionID = "0001" const testActionIDMiss = "0002" // this one doesn't exist const testActionIDBig = "0bbb" // non-inline object + const testActionIDEmpty = "0000" const testOutputID = "9900" const testOutputIDBig = "9bbb" + const testOutputIDEmpty = sha256OfEmpty const testObjectValue = "test data" testObjectValueBig := strings.Repeat("x", smallObjectSize+1) @@ -120,14 +122,16 @@ func TestServer(t *testing.T) { // Populate from the first client. wantPut(c1, testActionID, testOutputID, testObjectValue) wantPut(c1, testActionIDBig, testOutputIDBig, testObjectValueBig) + wantPut(c1, testActionIDEmpty, testOutputIDEmpty, "") // Read from the second client. wantGet(c2, testActionID, testOutputID, testObjectValue) wantGet(c2, testActionIDBig, testOutputIDBig, testObjectValueBig) + wantGet(c2, testActionIDEmpty, testOutputIDEmpty, "") // Check metrics - wantMetric(&srv.Gets, 2) - wantMetric(&srv.GetHits, 2) + wantMetric(&srv.Gets, 3) + wantMetric(&srv.GetHits, 3) wantMetric(&srv.GetHitsInline, 1) // Do the same get again from the same client. This shouldn't hit the network. From 91e5746ba521177bd4d896574201cd0ed988aaa4 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sun, 10 Aug 2025 15:36:14 -0700 Subject: [PATCH 17/67] cmd/gocached: redo the schema to store blobs separately ... for efficiency when multiple actions map to the same dup object ID. And to start laying the groundwork for namespaces. Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@fd8964a5ff2c21205d4438532814e718e00eacd8 --- cmd/gocached/gocached.go | 123 ++++++++++++++++++++++++++------------- gocache/cachers/disk.go | 10 +++- 2 files changed, 90 insertions(+), 43 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index da444c0..38aa8b7 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -34,7 +34,9 @@ And to insert an object: package main import ( + "cmp" "context" + "crypto/sha256" "database/sql" "errors" "expvar" @@ -93,29 +95,44 @@ func main() { log.Fatal(http.ListenAndServe(*listen, srv)) } -const schemaVersion = 2 +const schemaVersion = 3 const schema = ` PRAGMA journal_mode=WAL; CREATE TABLE IF NOT EXISTS Actions ( - ActionID TEXT NOT NULL PRIMARY KEY, - OutputID TEXT NOT NULL, - OutputSize INTEGER NOT NULL, -- bytes of the output (even if stored off-DB) + NamespaceID INTEGER NOT NULL, -- 0 for global trusted namespace + ActionID TEXT NOT NULL, + BlobID INTEGER NOT NULL, + AltOutputID TEXT NOT NULL DEFAULT '', -- if non-empty, the alternate object ID to use for this action; NULL means the blob's sha256 CreateTime INTEGER NOT NULL, -- unix sec when inserted (locally or on a peer) AccessTime INTEGER NOT NULL, -- unix sec of last access - InlineOutput BLOB, -- optional inline output value (e.g. for small output) + + PRIMARY KEY (NamespaceID, ActionID), CHECK (ActionID = lower(ActionID)), - CHECK (OutputID = lower(OutputID)), CHECK (ActionID GLOB '[0-9a-f]*'), - CHECK (OutputID GLOB '[0-9a-f]*' OR OutputID GLOB 'wk*'), - CHECK (OutputSize >= 0), CHECK (CreateTime >= 0), - CHECK (AccessTime >= 0), - CHECK (InlineOutput IS NULL OR length(InlineOutput) = OutputSize) + CHECK (AccessTime >= 0) +) STRICT; + +CREATE INDEX IF NOT EXISTS idx_actions_access ON Actions(AccessTime); +CREATE INDEX IF NOT EXISTS idx_actions_blobid ON Actions(BlobID); + +CREATE TABLE IF NOT EXISTS Blobs ( + BlobID INTEGER PRIMARY KEY AUTOINCREMENT, + SHA256 TEXT NOT NULL, + BlobSize INTEGER NOT NULL, -- size in bytes, either inline or on disk + SmallData BLOB, -- NULL if stored on disk + + CHECK (SmalLData IS NULL OR length(SmallData) = BlobSize) ) STRICT; -CREATE INDEX IF NOT EXISTS idx_actions_access ON Actions(AccessTime, OutputSize); +CREATE UNIQUE INDEX IF NOT EXISTS idx_blobs_sha256 ON Blobs(SHA256); + +CREATE TABLE IF NOT EXISTS Namespaces ( + NamespaceID INTEGER PRIMARY KEY AUTOINCREMENT, + Namespace TEXT NOT NULL UNIQUE CHECK (Namespace = lower(Namespace)) +) STRICT; ` func openDB(dbDir string) (*sql.DB, error) { @@ -293,20 +310,25 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { return } - var outputID string + var sha256hex string var size int64 - var inlineOutput sql.NullString + var smallData sql.NullString + var altObjectID string var accessTime int64 - err := s.db.QueryRow("SELECT OutputID, OutputSize, InlineOutput, AccessTime FROM Actions WHERE ActionID = ?", actionID).Scan(&outputID, &size, &inlineOutput, &accessTime) + namespaceID := 0 // global for now; TODO(bradfitz): support namespaces + err := s.db.QueryRow( + "SELECT b.SHA256, b.BlobSize, b.SmallData, a.AltOutputID, a.AccessTime FROM Actions a, Blobs b WHERE a.NameSpaceID = ? AND a.ActionID = ? AND a.BlobID = b.BlobID", + namespaceID, actionID).Scan( + &sha256hex, &size, &smallData, &altObjectID, &accessTime) if err != nil { if errors.Is(err, sql.ErrNoRows) { http.Error(w, "not found", http.StatusNotFound) return } - httpErr("bad request: missing Want-Object header", http.StatusBadRequest) + s.logf("QueryRow error: %v", err) + httpErr("Query: "+err.Error(), http.StatusInternalServerError) return } - outputID = outputIDFromWellknown(outputID) // If it's been more than a day since the last access, update the access time. // This is similar to the Linux "relatime" behavior. @@ -323,6 +345,8 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { s.GetHits.Add(1) + outputID := cmp.Or(altObjectID, sha256hex) + w.Header().Set("Content-Type", "application/octet-stream") w.Header().Set("Content-Length", fmt.Sprint(size)) w.Header().Set("Go-Output-Id", outputID) @@ -331,16 +355,19 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { return } - if inlineOutput.Valid { + if smallData.Valid { // For small outputs stored inline in the database, we can return them directly. s.GetHitsInline.Add(1) s.GetBytes.Add(size) - io.WriteString(w, inlineOutput.String) + io.WriteString(w, smallData.String) return } // Otherwise, for large objects that we know about, we can try to get them // from our local disk or a peer. + + // TODO(bradfitz): for namespace support (when that comes), this discovery + // should only be be sha256, not an actionID or namespace. rc, err := s.getObjectFromDiskOrPeer(ctx, actionID) if err != nil { httpErr(err.Error(), http.StatusInternalServerError) @@ -393,39 +420,67 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { return } - var inline []byte + hasher := sha256.New() + hashingBody := io.TeeReader(r.Body, hasher) + + var smallData []byte if r.ContentLength <= smallObjectSize { // Store small objects inline in the database. var err error - inline, err = io.ReadAll(r.Body) + smallData, err = io.ReadAll(hashingBody) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } } else { // For larger objects, we store them on disk. - _, err := s.disk.Put(ctx, actionID, outputID, r.ContentLength, r.Body) + _, err := s.disk.Put(ctx, actionID, outputID, r.ContentLength, hashingBody) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } } + sha256hex := fmt.Sprintf("%x", hasher.Sum(nil)) + blobSize := r.ContentLength + + var blobID int64 + err := s.db.QueryRow(`INSERT INTO Blobs (SHA256, BlobSize, SmallData) + VALUES (?, ?, ?) + ON CONFLICT(SHA256) DO UPDATE SET SHA256=excluded.SHA256 + RETURNING BlobID; +`, sha256hex, blobSize, smallData).Scan(&blobID) + if err != nil { + s.logf("Blobs insert error: %v", err) + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + // Insert or update the action in the database. nowUnix := s.now().Unix() - _, err := s.db.Exec(` -INSERT OR IGNORE INTO Actions (ActionID, OutputID, OutputSize, CreateTime, AccessTime, InlineOutput) -VALUES (?, ?, ?, ?, ?, ?)`, - actionID, outputIDOrWellKnown(outputID), r.ContentLength, nowUnix, nowUnix, inline) + altObjectID := "" + namespace := 0 // global for now; TODO(bradfitz): support namespaces + if sha256hex != outputID { + altObjectID = outputID + } + _, err = s.db.Exec(`INSERT OR IGNORE INTO Actions (NamespaceID, ActionID, BlobID, AltOutputID, CreateTime, AccessTime) + VALUES (?, ?, ?, ?, ?, ?)`, + namespace, + actionID, + blobID, + altObjectID, + nowUnix, + nowUnix, + ) if err != nil { - s.logf("INSERT error: %v", err) + s.logf("Actions insert error: %v", err) http.Error(w, err.Error(), http.StatusInternalServerError) return } s.Puts.Add(1) s.PutsBytes.Add(r.ContentLength) - if inline != nil { + if smallData != nil { s.PutsInline.Add(1) } @@ -437,20 +492,6 @@ VALUES (?, ?, ?, ?, ?, ?)`, // as "wk0" (for "well known 0") to save space, as it's common. const sha256OfEmpty = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" -func outputIDOrWellKnown(id string) string { - if id == sha256OfEmpty { - return "wk0" - } - return id -} - -func outputIDFromWellknown(idStored string) string { - if idStored == "wk0" { - return sha256OfEmpty - } - return idStored -} - // expvarCounterMetric is a Prometheus counter metric backed by an expvar.Int. type expvarCounterMetric struct { desc *prometheus.Desc diff --git a/gocache/cachers/disk.go b/gocache/cachers/disk.go index 790360d..3d64e67 100644 --- a/gocache/cachers/disk.go +++ b/gocache/cachers/disk.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "io" "log" @@ -96,8 +97,13 @@ func (dc *DiskCache) Put(ctx context.Context, actionID, outputID string, size in if len(actionID) < 4 || len(outputID) < 4 { return "", fmt.Errorf("actionID and outputID must be at least 4 characters long") } - if !validHex(actionID) || !validHex(outputID) { - return "", fmt.Errorf("actionID and outputID must be valid hex strings") + if !validHex(actionID) { + log.Printf("diskcache: got invalid actionID %q", actionID) + return "", errors.New("actionID must be hex") + } + if !validHex(outputID) { + log.Printf("diskcache: got invalid outputID %q", outputID) + return "", errors.New("outputID must be hex") } action2, output2 := actionID[:2], outputID[:2] From d50907c96847097cad5794664c62fb02bd709066 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sun, 10 Aug 2025 15:58:22 -0700 Subject: [PATCH 18/67] cmd/gocached: store only blobs on disk by sha256 Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@643853b17840bf8cecbd2530419ab473ee938ad3 --- cmd/gocached/gocached.go | 70 ++++++++++++++++++++++++++++++---------- 1 file changed, 53 insertions(+), 17 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index 38aa8b7..cb4c7d1 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -55,7 +55,6 @@ import ( "github.com/prometheus/client_golang/prometheus/collectors" "github.com/prometheus/client_golang/prometheus/promhttp" dto "github.com/prometheus/client_model/go" - "github.com/tailscale/tb/gocache/cachers" _ "modernc.org/sqlite" ) @@ -152,7 +151,6 @@ func newServer(dir string) (*server, error) { if err != nil { return nil, fmt.Errorf("openDB: %w", err) } - dc := &cachers.DiskCache{Dir: dir} reg := prometheus.NewRegistry() reg.MustRegister( @@ -163,7 +161,7 @@ func newServer(dir string) (*server, error) { srv := &server{ db: db, - disk: dc, + dir: dir, logf: log.Printf, } srv.registerMetrics(reg) @@ -212,7 +210,7 @@ func (s *server) registerMetrics(reg *prometheus.Registry) { type server struct { db *sql.DB - disk *cachers.DiskCache // for large outputs only + dir string // for SQLite DB + large blobs verbose bool logf func(format string, args ...any) clock func() time.Time // if non-nil, alternate time.Now for testing @@ -366,9 +364,7 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { // Otherwise, for large objects that we know about, we can try to get them // from our local disk or a peer. - // TODO(bradfitz): for namespace support (when that comes), this discovery - // should only be be sha256, not an actionID or namespace. - rc, err := s.getObjectFromDiskOrPeer(ctx, actionID) + rc, err := s.getObjectFromDiskOrPeer(ctx, sha256hex) if err != nil { httpErr(err.Error(), http.StatusInternalServerError) return @@ -387,25 +383,27 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { // exists but is not stored in SQLite. // // It returns (nil, nil) on miss. -func (s *server) getObjectFromDiskOrPeer(ctx context.Context, actionID string) (rc io.ReadCloser, err error) { - _, diskPath, diskErr := s.disk.Get(ctx, actionID) - if diskErr != nil { - return nil, diskErr +func (s *server) getObjectFromDiskOrPeer(ctx context.Context, sha256hex string) (rc io.ReadCloser, err error) { + if len(sha256hex) != sha256.Size*2 { + return nil, fmt.Errorf("invalid sha256hex %q", sha256hex) } - if diskPath != "" { - f, err := os.Open(diskPath) - if err != nil { + diskPath := filepath.Join(s.dir, sha256hex[:2], sha256hex) + f, err := os.Open(diskPath) + if err != nil { + if !os.IsNotExist(err) { return nil, err } + } + if err == nil { return f, nil } + // TODO(bradfitz): search peers, S3, etc. // For now, just return nil, nil on miss. return nil, nil } func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { - ctx := r.Context() if r.Method != "PUT" { http.Error(w, "bad method", http.StatusMethodNotAllowed) return @@ -432,10 +430,15 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { http.Error(w, err.Error(), http.StatusInternalServerError) return } + if int64(len(smallData)) != r.ContentLength { + // This check is redundant with net/http's validation, but + // for extra clarity. + http.Error(w, "bad content length", http.StatusInternalServerError) + return + } } else { // For larger objects, we store them on disk. - _, err := s.disk.Put(ctx, actionID, outputID, r.ContentLength, hashingBody) - if err != nil { + if err := s.writeDiskBlob(r.ContentLength, hashingBody); err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } @@ -487,6 +490,39 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusNoContent) } +func (s *server) writeDiskBlob(size int64, r io.Reader) (err error) { + nowUnix := s.now().Unix() + tf, err := os.CreateTemp(s.dir, fmt.Sprintf("upload-%d-*", nowUnix)) + if err != nil { + return err + } + defer func() { + if err == nil { + return + } + tf.Close() + os.Remove(tf.Name()) + }() + hasher := sha256.New() + n, err := io.Copy(tf, io.LimitReader(io.TeeReader(r, hasher), size+1)) + if err != nil { + return err + } + if n != size { + return fmt.Errorf("wrote %d bytes; wanted %d", n, size) + } + if err := tf.Close(); err != nil { + return err + } + hex := fmt.Sprintf("%02x", hasher.Sum(nil)) + dir := filepath.Join(s.dir, hex[:2]) + if err := os.MkdirAll(dir, 0755); err != nil { + return err + } + target := filepath.Join(dir, hex) + return os.Rename(tf.Name(), target) +} + // sha256OfEmpty is the SHA-256 hash of an empty string, used as a well-known // value in SQLite to store bytes, as it's common. We store it in SQLite // as "wk0" (for "well known 0") to save space, as it's common. From 3a129e38adb3a1f065a5e31075ee68419b7a31b2 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sun, 10 Aug 2025 20:48:47 -0700 Subject: [PATCH 19/67] cmd/gocached: start adding cleaning Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@5ffda40f35d0283c5c28814706deafe40eef02fc --- cmd/gocached/gocached.go | 161 +++++++++++++++-- cmd/gocached/gocached_test.go | 328 +++++++++++++++++++++++----------- go.mod | 1 + 3 files changed, 372 insertions(+), 118 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index cb4c7d1..2caef5a 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -44,11 +44,13 @@ import ( "fmt" "io" "log" + "math" "net/http" "os" "path/filepath" "reflect" "strings" + "sync" "time" "github.com/prometheus/client_golang/prometheus" @@ -216,6 +218,13 @@ type server struct { clock func() time.Time // if non-nil, alternate time.Now for testing metricsHandler http.Handler + // statsMu protects stats & cleanups to the DB, so that we don't have + // multiple goroutines writing to the DB at the same time. + // + // It's a read-write mutex so that writes can still happen concurrently (to + // the extent permitted by SQLite, which will still serialize its bit). + statsMu sync.RWMutex + // Metrics ActiveGets expvar.Int `type:"gauge" name:"active_gets" help:"currently pending get requests; should usually be zero"` Gets expvar.Int `type:"counter" name:"gets" help:"total number of gocache get requests"` // gets = getHits + getErrs + implicit misses @@ -332,6 +341,8 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { // This is similar to the Linux "relatime" behavior. now := s.now().Unix() if accessTime < now-relAtimeSeconds { + // TODO(bradfitz): do this async? not worth blocking the caller. + // But we need a mechanism for tests to wait on async work. _, err := s.db.Exec("UPDATE Actions SET AccessTime = ? WHERE ActionID = ?", now, actionID) if err != nil { s.logf("Update AccessTime error: %v", err) @@ -370,6 +381,11 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { return } if rc == nil { + // Our database suggested we should've had this object, + // but maybe somebody delete it by hand from the filesystem. + // Just treat it as a cache miss. The background cleanup + // will eventually remove the Action row from the DB + // after identifying it as a dangling reference. http.Error(w, "not found", http.StatusNotFound) return } @@ -390,20 +406,20 @@ func (s *server) getObjectFromDiskOrPeer(ctx context.Context, sha256hex string) diskPath := filepath.Join(s.dir, sha256hex[:2], sha256hex) f, err := os.Open(diskPath) if err != nil { - if !os.IsNotExist(err) { - return nil, err + if os.IsNotExist(err) { + // TODO(bradfitz): search peers, S3, etc. + // For now, just return nil, nil on miss. + return nil, nil } + return nil, err } - if err == nil { - return f, nil - } - - // TODO(bradfitz): search peers, S3, etc. - // For now, just return nil, nil on miss. - return nil, nil + return f, nil } func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { + s.statsMu.RLock() + defer s.statsMu.RUnlock() + if r.Method != "PUT" { http.Error(w, "bad method", http.StatusMethodNotAllowed) return @@ -523,10 +539,129 @@ func (s *server) writeDiskBlob(size int64, r io.Reader) (err error) { return os.Rename(tf.Name(), target) } -// sha256OfEmpty is the SHA-256 hash of an empty string, used as a well-known -// value in SQLite to store bytes, as it's common. We store it in SQLite -// as "wk0" (for "well known 0") to save space, as it's common. -const sha256OfEmpty = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" +type countAndSize struct { + Count int64 // number of actions + Size int64 // total size of all actions' blobs (even if shared by other actions) +} + +type usageStats struct { + // ActionsLE is a histogram of the actions in the DB by their access time. + // + // The key is a Prometheus-style histogram "less than" value. That is, if + // there are map keys for 24h and 48h, the latter includes the sum of the + // 24h values as well. + // + // The map keys are day-granularity, as the access time is only updated once + // it's over a day old. + // + // So the map keys are 24h, 48h, 72h, 96h, 168h (7d), 336h (14d), 720h + // (30d), and 2160h (90d) and math.MaxInt64 for infinity. + ActionsLE map[time.Duration]countAndSize + + // MissingBlobRows is the number of rows in the Actions table that + // reference a BlobID that doesn't exist in the Blobs table. + // This should always be zero in a healthy system. + MissingBlobRows int +} + +var durs = []time.Duration{ + 24 * time.Hour, 48 * time.Hour, 72 * time.Hour, + 96 * time.Hour, 168 * time.Hour, 336 * time.Hour, + 720 * time.Hour, 2160 * time.Hour, math.MaxInt64, +} + +func (s *server) usageStats() (_ *usageStats, err error) { + defer func() { + if err != nil { + s.logf("usageStats error: %v", err) + } + }() + + s.statsMu.Lock() + defer s.statsMu.Unlock() + + st := &usageStats{ + ActionsLE: make(map[time.Duration]countAndSize), + } + + now := s.now().Unix() + rows, err := s.db.Query( + "SELECT a.BlobID, a.AccessTime, b.BlobSize FROM Actions a LEFT JOIN Blobs b ON a.BlobID = b.BlobID") + if err != nil { + return nil, fmt.Errorf("query Actions: %w", err) + } + var blobID int64 + var accessTime int64 + var blobSize sql.NullInt64 + for rows.Next() { + if err := rows.Scan(&blobID, &accessTime, &blobSize); err != nil { + return nil, fmt.Errorf("rows.Scan: %w", err) + } + if !blobSize.Valid { + st.MissingBlobRows++ + continue + } + + dur := time.Duration(now-accessTime) * time.Second + if dur < 0 { + dur = 0 + } + for _, d := range durs { + if dur < d { + was := st.ActionsLE[d] + was.Count++ + was.Size += blobSize.Int64 + st.ActionsLE[d] = was + } + } + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("rows.Next: %w", err) + } + + return st, nil +} + +type cleanCandidate struct { + BlobID int64 + Age time.Duration + BlobSize int64 // size of the blob, in bytes +} + +func (s *server) cleanCandidates(olderThan time.Duration, limit int) ([]cleanCandidate, error) { + now := s.now() + nowUnix := now.Unix() + cutoff := now.Add(-olderThan).Unix() + + rows, err := s.db.Query(` + SELECT b.BlobID, MAX(a.AccessTime), b.BlobSize + FROM Blobs b LEFT JOIN Actions a ON b.BlobID = a.BlobID + GROUP BY b.BlobID + HAVING MAX(a.AccessTime) <= ? + ORDER BY MAX(a.AccessTime) + LIMIT ?`, cutoff, limit) + if err != nil { + return nil, fmt.Errorf("query clean candidates: %w", err) + } + defer rows.Close() + + var candidates []cleanCandidate + var accessTime int64 + for rows.Next() { + var c cleanCandidate + if err := rows.Scan(&c.BlobID, &accessTime, &c.BlobSize); err != nil { + return nil, fmt.Errorf("rows.Scan: %w", err) + } + c.Age = time.Duration(nowUnix-accessTime) * time.Second + candidates = append(candidates, c) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("rows.Next: %w", err) + } + + return candidates, nil + +} // expvarCounterMetric is a Prometheus counter metric backed by an expvar.Int. type expvarCounterMetric struct { diff --git a/cmd/gocached/gocached_test.go b/cmd/gocached/gocached_test.go index fda6fe1..e6385d9 100644 --- a/cmd/gocached/gocached_test.go +++ b/cmd/gocached/gocached_test.go @@ -3,6 +3,8 @@ package main import ( "context" "expvar" + "fmt" + "math" "net/http/httptest" "os" "strings" @@ -10,64 +12,134 @@ import ( "testing" "time" + "github.com/google/go-cmp/cmp" "github.com/tailscale/tb/gocache/cachers" ) -func TestServer(t *testing.T) { +// sha256OfEmpty is the SHA-256 hash of an empty string, used as a well-known +// value in SQLite to store bytes, as it's common. +const sha256OfEmpty = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" + +type tester struct { + t testing.TB + srv *server + hs *httptest.Server + + timeMu sync.Mutex + curTime time.Time +} + +func (t *tester) Logf(format string, args ...any) { + t.t.Logf(format, args...) +} + +func (t *tester) now() time.Time { + t.timeMu.Lock() + defer t.timeMu.Unlock() + return t.curTime +} + +func (t *tester) advanceClock(d time.Duration) { + t.timeMu.Lock() + defer t.timeMu.Unlock() + t.curTime = t.curTime.Add(d) +} + +func (t *tester) mkClient() *cachers.HTTPClient { + clientCacheDir := t.t.TempDir() + return &cachers.HTTPClient{ + BaseURL: t.hs.URL, + Disk: &cachers.DiskCache{ + Dir: clientCacheDir, + Logf: func(format string, args ...any) { + t.Logf("client-disk: "+format, args...) + }, + }, + } +} + +// wantMetric is a helper to check an expvar.Int metric and reset it +// for future tests. +func (st *tester) wantMetric(m *expvar.Int, want int64) { + st.t.Helper() + if got := m.Value(); got != want { + st.t.Errorf("metric = %d, want %d", got, want) + } + m.Set(0) +} + +func (st *tester) wantPut(c *cachers.HTTPClient, actionID, outputID string, val string) { ctx := context.Background() + st.t.Helper() + clientDiskPath, err := c.Put(ctx, actionID, outputID, int64(len(val)), strings.NewReader(val)) + if err != nil { + st.t.Fatalf("Put: %v", err) + } + if clientDiskPath == "" { + st.t.Fatal("Put returned empty disk path") + } + st.wantMetric(&st.srv.Puts, 1) + wrote, err := os.ReadFile(clientDiskPath) + if err != nil { + st.t.Fatalf("ReadFile: %v", err) + } + if string(wrote) != val { + st.t.Errorf("ReadFile got %q, want %q", wrote, val) + } +} - dir := t.TempDir() - srv, err := newServer(dir) +func (st *tester) wantGet(c *cachers.HTTPClient, actionID, outputID, wantVal string) { + ctx := context.Background() + st.t.Helper() + gotOutputID, diskPath, err := c.Get(ctx, actionID) if err != nil { - t.Fatalf("newServer: %v", err) + st.t.Fatalf("Get: %v", err) } - srv.logf = t.Logf - srv.verbose = true - - var ( - timeMu sync.Mutex - now = time.Unix(1234, 0) - ) - srv.clock = func() time.Time { - timeMu.Lock() - defer timeMu.Unlock() - return now - } - advanceClock := func(d time.Duration) { - timeMu.Lock() - defer timeMu.Unlock() - now = now.Add(d) - } - - hs := httptest.NewServer(srv) - defer hs.Close() - - mkClient := func() *cachers.HTTPClient { - clientCacheDir := t.TempDir() - return &cachers.HTTPClient{ - BaseURL: hs.URL, - Disk: &cachers.DiskCache{ - Dir: clientCacheDir, - Logf: func(format string, args ...any) { - t.Logf("client-disk: "+format, args...) - }, - }, - } + if gotOutputID != outputID { + st.t.Errorf("Get got outputID %q, want %q", gotOutputID, outputID) + } + if diskPath == "" { + st.t.Fatal("Get returned empty disk path") + } + wrote, err := os.ReadFile(diskPath) + if err != nil { + st.t.Fatalf("ReadFile: %v", err) + } + if string(wrote) != wantVal { + st.t.Errorf("ReadFile got %q, want %q", wrote, wantVal) + } +} + +func newServerTester(t testing.TB) *tester { + st := &tester{ + t: t, + curTime: time.Unix(1234, 0), } - // wantMetric is a helper to check an expvar.Int metric and reset it - // for future tests. - wantMetric := func(m *expvar.Int, want int64) { - t.Helper() - if got := m.Value(); got != want { - t.Errorf("metric = %d, want %d", got, want) - } - m.Set(0) + var err error + dir := t.TempDir() + st.srv, err = newServer(dir) + if err != nil { + t.Fatalf("newServer: %v", err) } + st.srv.logf = t.Logf + st.srv.verbose = true + st.srv.clock = st.now + + st.hs = httptest.NewServer(st.srv) + t.Cleanup(st.hs.Close) + + return st +} + +func TestServer(t *testing.T) { + st := newServerTester(t) + + ctx := context.Background() // Make two clients (imagine: two different builder VMs) - c1 := mkClient() - c2 := mkClient() + c1 := st.mkClient() + c2 := st.mkClient() const testActionID = "0001" const testActionIDMiss = "0002" // this one doesn't exist @@ -79,77 +151,123 @@ func TestServer(t *testing.T) { const testObjectValue = "test data" testObjectValueBig := strings.Repeat("x", smallObjectSize+1) - wantPut := func(c *cachers.HTTPClient, actionID, outputID string, val string) { - t.Helper() - clientDiskPath, err := c.Put(ctx, actionID, outputID, int64(len(val)), strings.NewReader(val)) - if err != nil { - t.Fatalf("Put: %v", err) - } - if clientDiskPath == "" { - t.Fatal("Put returned empty disk path") - } - wantMetric(&srv.Puts, 1) - wrote, err := os.ReadFile(clientDiskPath) - if err != nil { - t.Fatalf("ReadFile: %v", err) - } - if string(wrote) != val { - t.Errorf("ReadFile got %q, want %q", wrote, val) - } - } - - wantGet := func(c *cachers.HTTPClient, actionID, outputID, wantVal string) { - t.Helper() - gotOutputID, diskPath, err := c.Get(ctx, actionID) - if err != nil { - t.Fatalf("Get: %v", err) - } - if gotOutputID != outputID { - t.Errorf("Get got outputID %q, want %q", gotOutputID, outputID) - } - if diskPath == "" { - t.Fatal("Get returned empty disk path") - } - wrote, err := os.ReadFile(diskPath) - if err != nil { - t.Fatalf("ReadFile: %v", err) - } - if string(wrote) != wantVal { - t.Errorf("ReadFile got %q, want %q", wrote, wantVal) - } - } - // Populate from the first client. - wantPut(c1, testActionID, testOutputID, testObjectValue) - wantPut(c1, testActionIDBig, testOutputIDBig, testObjectValueBig) - wantPut(c1, testActionIDEmpty, testOutputIDEmpty, "") + st.wantPut(c1, testActionID, testOutputID, testObjectValue) + st.wantPut(c1, testActionIDBig, testOutputIDBig, testObjectValueBig) + st.wantPut(c1, testActionIDEmpty, testOutputIDEmpty, "") // Read from the second client. - wantGet(c2, testActionID, testOutputID, testObjectValue) - wantGet(c2, testActionIDBig, testOutputIDBig, testObjectValueBig) - wantGet(c2, testActionIDEmpty, testOutputIDEmpty, "") + st.wantGet(c2, testActionID, testOutputID, testObjectValue) + st.wantGet(c2, testActionIDBig, testOutputIDBig, testObjectValueBig) + st.wantGet(c2, testActionIDEmpty, testOutputIDEmpty, "") // Check metrics - wantMetric(&srv.Gets, 3) - wantMetric(&srv.GetHits, 3) - wantMetric(&srv.GetHitsInline, 1) + st.wantMetric(&st.srv.Gets, 3) + st.wantMetric(&st.srv.GetHits, 3) + st.wantMetric(&st.srv.GetHitsInline, 1) // Do the same get again from the same client. This shouldn't hit the network. - wantGet(c2, testActionID, testOutputID, testObjectValue) - wantMetric(&srv.Gets, 0) + st.wantGet(c2, testActionID, testOutputID, testObjectValue) + st.wantMetric(&st.srv.Gets, 0) // Cache miss. This should hit the network and fail. - if _, _, err = c2.Get(ctx, testActionIDMiss); err != nil { + if _, _, err := c2.Get(ctx, testActionIDMiss); err != nil { t.Fatalf("miss Get: %v", err) } - wantMetric(&srv.Gets, 1) - wantMetric(&srv.GetHits, 0) + st.wantMetric(&st.srv.Gets, 1) + st.wantMetric(&st.srv.GetHits, 0) // Check that access time gets updated. // Do it from a fresh client without a disk cache. - wantMetric(&srv.GetAccessBumps, 0) - advanceClock(relAtimeSeconds * 2 * time.Second) // advance clock by 2 days - c3 := mkClient() - wantGet(c3, testActionID, testOutputID, testObjectValue) - wantMetric(&srv.GetAccessBumps, 1) + st.wantMetric(&st.srv.GetAccessBumps, 0) + st.advanceClock(relAtimeSeconds * 2 * time.Second) // advance clock by 2 days + c3 := st.mkClient() + st.wantGet(c3, testActionID, testOutputID, testObjectValue) + st.wantMetric(&st.srv.GetAccessBumps, 1) + + // Get usage stats. + stats, err := st.srv.usageStats() + if err != nil { + t.Fatalf("usageStats: %v", err) + } + want := &usageStats{ + MissingBlobRows: 0, + ActionsLE: map[time.Duration]countAndSize{ + 24 * time.Hour: {Count: 1, Size: 9}, + 48 * time.Hour: {Count: 1, Size: 9}, + 72 * time.Hour: {Count: 3, Size: 1034}, + 96 * time.Hour: {Count: 3, Size: 1034}, + 168 * time.Hour: {Count: 3, Size: 1034}, + 336 * time.Hour: {Count: 3, Size: 1034}, + 720 * time.Hour: {Count: 3, Size: 1034}, + 2160 * time.Hour: {Count: 3, Size: 1034}, + math.MaxInt64: {Count: 3, Size: 1034}, + }, + } + if diff := cmp.Diff(stats, want); diff != "" { + t.Errorf("usageStats mismatch (-got +want):\n%s", diff) + } + + st.advanceClock(relAtimeSeconds * 2 * time.Second) // advance clock by 2 days +} + +func TestCleanCandidates(t *testing.T) { + st := newServerTester(t) + + // Populate some data. + c1 := st.mkClient() + st.wantPut(c1, "0001", "9901", "1") + st.advanceClock(24 * time.Hour) + st.wantPut(c1, "0002", "9902", "22") + st.advanceClock(24 * time.Hour) + st.wantPut(c1, "0003", "9903", "333") + st.advanceClock(24 * time.Hour) + st.wantPut(c1, "0004", "9904", strings.Repeat("x", smallObjectSize+1)) + + const day = 24 * time.Hour + + tests := []struct { + maxAge time.Duration + limit int + want []cleanCandidate + }{ + { + maxAge: 0, + limit: 100, + want: []cleanCandidate{ + {BlobID: 1, Age: 3 * day, BlobSize: 1}, + {BlobID: 2, Age: 2 * day, BlobSize: 2}, + {BlobID: 3, Age: 1 * day, BlobSize: 3}, + {BlobID: 4, Age: 0, BlobSize: smallObjectSize + 1}, + }, + }, + { + maxAge: 25 * time.Hour, + limit: 100, + want: []cleanCandidate{ + {BlobID: 1, Age: 3 * day, BlobSize: 1}, + {BlobID: 2, Age: 2 * day, BlobSize: 2}, + }, + }, + { + maxAge: 0, + limit: 2, + want: []cleanCandidate{ + {BlobID: 1, Age: 3 * day, BlobSize: 1}, + {BlobID: 2, Age: 2 * day, BlobSize: 2}, + }, + }, + } + for _, tt := range tests { + t.Run(fmt.Sprintf("maxAge=%v,limit=%d", tt.maxAge, tt.limit), func(t *testing.T) { + candidates, err := st.srv.cleanCandidates(tt.maxAge, tt.limit) + if err != nil { + t.Fatal(err) + } + if diff := cmp.Diff(candidates, tt.want); diff != "" { + t.Errorf("cleanCandidates mismatch (-got +want):\n%s", diff) + + } + }) + } } diff --git a/go.mod b/go.mod index 2e9738d..da652cd 100644 --- a/go.mod +++ b/go.mod @@ -4,6 +4,7 @@ go 1.23.0 toolchain go1.23.4 require ( + github.com/google/go-cmp v0.7.0 github.com/prometheus/client_golang v1.23.0 github.com/prometheus/client_model v0.6.2 modernc.org/sqlite v1.38.2 From 5adb8d8814cf57e81a4741ab9a922c8d15e1b9f4 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Mon, 11 Aug 2025 09:43:53 -0700 Subject: [PATCH 20/67] cmd/gocached: do background cleaning Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@13dbb96f480d79e27eb7cee577390ecd3abb61ba --- cmd/gocached/gocached.go | 374 +++++++++++++++++++++++++++++----- cmd/gocached/gocached_test.go | 116 ++++++++++- 2 files changed, 437 insertions(+), 53 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index 2caef5a..b94bbca 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -38,19 +38,22 @@ import ( "context" "crypto/sha256" "database/sql" + "encoding/hex" "errors" "expvar" "flag" "fmt" "io" "log" + "maps" "math" "net/http" "os" "path/filepath" "reflect" + "slices" "strings" - "sync" + "sync/atomic" "time" "github.com/prometheus/client_golang/prometheus" @@ -69,6 +72,9 @@ var ( dir = flag.String("cache-dir", "", "cache directory, if empty defaults to /gocached") verbose = flag.Bool("verbose", false, "be verbose") listen = flag.String("listen", ":31364", "listen address") + + maxSize = flag.Int("max-size-gb", 50, "maximum size of the cache in GiB; 0 means no limit") + maxAge = flag.Int("max-age-days", 60, "maximum age of objects in the cache in days; 0 means no limit") ) func main() { @@ -91,8 +97,25 @@ func main() { log.Fatalf("newServer: %v", err) } srv.verbose = *verbose + srv.maxSize = int64(*maxSize) << 30 + srv.maxAge = time.Duration(*maxAge) * 24 * time.Hour + + log.Printf("gocached: scanning usage & cleaning as needed...") + us, err := srv.usageStats() + if err != nil { + log.Fatalf("getting usage stats: %v", err) + } + + log.Printf("gocached: current usage: %v of limit %v", us.All(), bytesFmt(srv.maxSize)) + if res, err := srv.cleanOldObjects(us); err != nil { + log.Fatalf("clean old objects: %v", err) + } else if res.Count > 0 { + log.Printf("gocached: cleaned %v", res) + } - log.Printf("gocached listening on %s ...", *listen) + go srv.runCleanLoop() + + log.Printf("gocached: listening on %s ...", *listen) log.Fatal(http.ListenAndServe(*listen, srv)) } @@ -166,6 +189,7 @@ func newServer(dir string) (*server, error) { dir: dir, logf: log.Printf, } + srv.shutdownCtx, srv.shutdownCancel = context.WithCancel(context.Background()) srv.registerMetrics(reg) srv.metricsHandler = promhttp.HandlerFor(reg, promhttp.HandlerOpts{ @@ -217,16 +241,16 @@ type server struct { logf func(format string, args ...any) clock func() time.Time // if non-nil, alternate time.Now for testing metricsHandler http.Handler + maxSize int64 // maximum size of the cache in bytes; 0 means no limit + maxAge time.Duration // maximum age of objects; 0 means no limit + shutdownCtx context.Context + shutdownCancel context.CancelFunc - // statsMu protects stats & cleanups to the DB, so that we don't have - // multiple goroutines writing to the DB at the same time. - // - // It's a read-write mutex so that writes can still happen concurrently (to - // the extent permitted by SQLite, which will still serialize its bit). - statsMu sync.RWMutex + lastUsage atomic.Pointer[usageStats] // Metrics ActiveGets expvar.Int `type:"gauge" name:"active_gets" help:"currently pending get requests; should usually be zero"` + ActivePuts expvar.Int `type:"gauge" name:"active_puts" help:"currently pending put requests; should usually be zero"` Gets expvar.Int `type:"counter" name:"gets" help:"total number of gocache get requests"` // gets = getHits + getErrs + implicit misses GetBytes expvar.Int `type:"counter" name:"get_bytes" help:"total bytes fetched from gocache gets that were cache hits"` GetHits expvar.Int `type:"counter" name:"get_hits" help:"total number of successful gocache get requests"` @@ -236,21 +260,27 @@ type server struct { Puts expvar.Int `type:"counter" name:"puts" help:"total number of gocache put requests"` PutsBytes expvar.Int `type:"counter" name:"put_bytes" help:"total bytes added from gocache puts"` PutsInline expvar.Int `type:"counter" name:"put_inline" help:"subset of gocached_puts that were stored inline (small objects)"` + BlobCount expvar.Int `type:"gauge" name:"blob_count" help:"number of blobs currently stored in the cache"` + BlobBytes expvar.Int `type:"gauge" name:"blob_bytes" help:"sum of blob sizes currently stored in the cache"` } -func (s *server) now() time.Time { - if s.clock != nil { - return s.clock() +func (srv *server) now() time.Time { + if srv.clock != nil { + return srv.clock() } return time.Now() } -func (s *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { - if s.verbose { - s.logf("ServeHTTP: %s %s", r.Method, r.RequestURI) +func (srv *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if srv.verbose { + srv.logf("ServeHTTP: %s %s", r.Method, r.RequestURI) + } + if r.URL.Path == "/usage" { + srv.serveUsage(w, r) + return } if r.Method == "PUT" { - s.handlePut(w, r) + srv.handlePut(w, r) return } if r.Method != "GET" && r.Method != "HEAD" { @@ -259,11 +289,15 @@ func (s *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { } switch { case strings.HasPrefix(r.URL.Path, "/action/"): - s.handleGetAction(w, r) + srv.handleGetAction(w, r) case r.URL.Path == "/": - io.WriteString(w, "gocached") + w.Header().Set("Content-Type", "text/html; charset=utf-8") + io.WriteString(w, "

gocached

") + io.WriteString(w, "

This is a shared Go build cache server, hit by GOCACHEPROG clients.

") + io.WriteString(w, "

See /usage for usage stats.

") + io.WriteString(w, "

See /metrics for Prometheus metrics.

") case r.URL.Path == "/metrics": - s.metricsHandler.ServeHTTP(w, r) + srv.metricsHandler.ServeHTTP(w, r) default: http.Error(w, "not found", http.StatusNotFound) } @@ -295,16 +329,16 @@ func validHex(x string) bool { // we do a DB write to update it. const relAtimeSeconds = 60 * 60 * 24 // 1 day -func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { - s.ActiveGets.Add(1) - defer s.ActiveGets.Add(-1) +func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request) { + srv.ActiveGets.Add(1) + defer srv.ActiveGets.Add(-1) - s.Gets.Add(1) + srv.Gets.Add(1) ctx := r.Context() httpErr := func(msg string, code int) { http.Error(w, msg, code) - s.GetErrs.Add(1) + srv.GetErrs.Add(1) } actionID, ok := getHexSuffix(r, "/action/") @@ -323,7 +357,7 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { var altObjectID string var accessTime int64 namespaceID := 0 // global for now; TODO(bradfitz): support namespaces - err := s.db.QueryRow( + err := srv.db.QueryRow( "SELECT b.SHA256, b.BlobSize, b.SmallData, a.AltOutputID, a.AccessTime FROM Actions a, Blobs b WHERE a.NameSpaceID = ? AND a.ActionID = ? AND a.BlobID = b.BlobID", namespaceID, actionID).Scan( &sha256hex, &size, &smallData, &altObjectID, &accessTime) @@ -332,27 +366,27 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { http.Error(w, "not found", http.StatusNotFound) return } - s.logf("QueryRow error: %v", err) + srv.logf("QueryRow error: %v", err) httpErr("Query: "+err.Error(), http.StatusInternalServerError) return } // If it's been more than a day since the last access, update the access time. // This is similar to the Linux "relatime" behavior. - now := s.now().Unix() + now := srv.now().Unix() if accessTime < now-relAtimeSeconds { // TODO(bradfitz): do this async? not worth blocking the caller. // But we need a mechanism for tests to wait on async work. - _, err := s.db.Exec("UPDATE Actions SET AccessTime = ? WHERE ActionID = ?", now, actionID) + _, err := srv.db.Exec("UPDATE Actions SET AccessTime = ? WHERE ActionID = ?", now, actionID) if err != nil { - s.logf("Update AccessTime error: %v", err) + srv.logf("Update AccessTime error: %v", err) httpErr("internal server error", http.StatusInternalServerError) return } - s.GetAccessBumps.Add(1) + srv.GetAccessBumps.Add(1) } - s.GetHits.Add(1) + srv.GetHits.Add(1) outputID := cmp.Or(altObjectID, sha256hex) @@ -366,8 +400,8 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { if smallData.Valid { // For small outputs stored inline in the database, we can return them directly. - s.GetHitsInline.Add(1) - s.GetBytes.Add(size) + srv.GetHitsInline.Add(1) + srv.GetBytes.Add(size) io.WriteString(w, smallData.String) return } @@ -375,7 +409,7 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { // Otherwise, for large objects that we know about, we can try to get them // from our local disk or a peer. - rc, err := s.getObjectFromDiskOrPeer(ctx, sha256hex) + rc, err := srv.getObjectFromDiskOrPeer(ctx, sha256hex) if err != nil { httpErr(err.Error(), http.StatusInternalServerError) return @@ -389,7 +423,7 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { http.Error(w, "not found", http.StatusNotFound) return } - s.GetBytes.Add(size) + srv.GetBytes.Add(size) defer rc.Close() io.Copy(w, rc) } @@ -399,11 +433,11 @@ func (s *server) handleGetAction(w http.ResponseWriter, r *http.Request) { // exists but is not stored in SQLite. // // It returns (nil, nil) on miss. -func (s *server) getObjectFromDiskOrPeer(ctx context.Context, sha256hex string) (rc io.ReadCloser, err error) { +func (srv *server) getObjectFromDiskOrPeer(ctx context.Context, sha256hex string) (rc io.ReadCloser, err error) { if len(sha256hex) != sha256.Size*2 { return nil, fmt.Errorf("invalid sha256hex %q", sha256hex) } - diskPath := filepath.Join(s.dir, sha256hex[:2], sha256hex) + diskPath := filepath.Join(srv.dir, sha256hex[:2], sha256hex) f, err := os.Open(diskPath) if err != nil { if os.IsNotExist(err) { @@ -417,8 +451,8 @@ func (s *server) getObjectFromDiskOrPeer(ctx context.Context, sha256hex string) } func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { - s.statsMu.RLock() - defer s.statsMu.RUnlock() + s.ActivePuts.Add(1) + defer s.ActivePuts.Add(-1) if r.Method != "PUT" { http.Error(w, "bad method", http.StatusMethodNotAllowed) @@ -506,6 +540,11 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusNoContent) } +func (s *server) sha256Filepath(hash [sha256.Size]byte) string { + hex := fmt.Sprintf("%x", hash) + return filepath.Join(s.dir, hex[:2], hex) +} + func (s *server) writeDiskBlob(size int64, r io.Reader) (err error) { nowUnix := s.now().Unix() tf, err := os.CreateTemp(s.dir, fmt.Sprintf("upload-%d-*", nowUnix)) @@ -530,12 +569,13 @@ func (s *server) writeDiskBlob(size int64, r io.Reader) (err error) { if err := tf.Close(); err != nil { return err } - hex := fmt.Sprintf("%02x", hasher.Sum(nil)) - dir := filepath.Join(s.dir, hex[:2]) - if err := os.MkdirAll(dir, 0755); err != nil { + var hash [sha256.Size]byte + hasher.Sum(hash[:0]) + + target := s.sha256Filepath(hash) + if err := os.MkdirAll(filepath.Dir(target), 0755); err != nil { return err } - target := filepath.Join(dir, hex) return os.Rename(tf.Name(), target) } @@ -544,6 +584,13 @@ type countAndSize struct { Size int64 // total size of all actions' blobs (even if shared by other actions) } +func (cs countAndSize) String() string { + if cs.Count == 0 { + return "0 objects, 0 bytes" + } + return fmt.Sprintf("%d objects, %s", cs.Count, bytesFmt(cs.Size)) +} + type usageStats struct { // ActionsLE is a histogram of the actions in the DB by their access time. // @@ -554,7 +601,7 @@ type usageStats struct { // The map keys are day-granularity, as the access time is only updated once // it's over a day old. // - // So the map keys are 24h, 48h, 72h, 96h, 168h (7d), 336h (14d), 720h + // So the map keys are 24h, 48h, 96h, 168h (7d), 336h (14d), 720h // (30d), and 2160h (90d) and math.MaxInt64 for infinity. ActionsLE map[time.Duration]countAndSize @@ -564,10 +611,19 @@ type usageStats struct { MissingBlobRows int } -var durs = []time.Duration{ - 24 * time.Hour, 48 * time.Hour, 72 * time.Hour, - 96 * time.Hour, 168 * time.Hour, 336 * time.Hour, - 720 * time.Hour, 2160 * time.Hour, math.MaxInt64, +func (us *usageStats) All() countAndSize { return us.ActionsLE[math.MaxInt64] } + +const day = 24 * time.Hour + +var standardDurs = []time.Duration{ + 1 * day, + 2 * day, + 4 * day, + 7 * day, + 14 * day, + 30 * day, + 90 * day, + math.MaxInt64, } func (s *server) usageStats() (_ *usageStats, err error) { @@ -577,13 +633,28 @@ func (s *server) usageStats() (_ *usageStats, err error) { } }() - s.statsMu.Lock() - defer s.statsMu.Unlock() - st := &usageStats{ ActionsLE: make(map[time.Duration]countAndSize), } + // Build the durations to use for the histogram. + // The math.MaxInt64 value is always included. + // If s.maxAge is set, we ignore sizes above that, except + // for the math.MaxInt64 value. + var durs []time.Duration + if s.maxAge == 0 { + durs = standardDurs + } else { + durs = make([]time.Duration, 0, len(standardDurs)+1) + durs = append(durs, s.maxAge) + for _, d := range standardDurs { + if d < s.maxAge || d == math.MaxInt64 { + durs = append(durs, d) + } + } + slices.Sort(durs) + } + now := s.now().Unix() rows, err := s.db.Query( "SELECT a.BlobID, a.AccessTime, b.BlobSize FROM Actions a LEFT JOIN Blobs b ON a.BlobID = b.BlobID") @@ -619,6 +690,10 @@ func (s *server) usageStats() (_ *usageStats, err error) { return nil, fmt.Errorf("rows.Next: %w", err) } + s.lastUsage.Store(st) + all := st.All() + s.BlobCount.Set(all.Count) + s.BlobBytes.Set(all.Size) return st, nil } @@ -628,7 +703,7 @@ type cleanCandidate struct { BlobSize int64 // size of the blob, in bytes } -func (s *server) cleanCandidates(olderThan time.Duration, limit int) ([]cleanCandidate, error) { +func (s *server) cleanCandidates(olderThan time.Duration, limit int64) ([]cleanCandidate, error) { now := s.now() nowUnix := now.Unix() cutoff := now.Add(-olderThan).Unix() @@ -660,7 +735,204 @@ func (s *server) cleanCandidates(olderThan time.Duration, limit int) ([]cleanCan } return candidates, nil +} + +func (srv *server) deleteBlobs(blobIDs ...int64) error { + tx, err := srv.db.Begin() + if err != nil { + return fmt.Errorf("delete blob Begin: %w", err) + } + defer tx.Rollback() + for _, blobID := range blobIDs { + var sha256Hex string + if err := tx.QueryRow("SELECT SHA256 FROM Blobs WHERE BlobID = ?", blobID).Scan(&sha256Hex); err != nil && !errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("querying blob SHA256: %w", err) + } + if _, err := tx.Exec("DELETE FROM Blobs WHERE BlobID = ?", blobID); err != nil { + return fmt.Errorf("deleting blob: %w", err) + } + if _, err := tx.Exec("DELETE FROM Actions WHERE BlobID = ?", blobID); err != nil { + return fmt.Errorf("deleting actions: %w", err) + } + var hash [sha256.Size]byte + if _, err := hex.Decode(hash[:], []byte(sha256Hex)); err == nil { + if err := os.Remove(srv.sha256Filepath(hash)); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("removing disk file: %w", err) + } + } + } + return tx.Commit() +} + +func (srv *server) cleanOldObjects(us *usageStats) (countAndSize, error) { + var zero countAndSize + var ret countAndSize + + all := us.ActionsLE[math.MaxInt64] + if srv.verbose { + srv.logf("current usage stats: %v", all) + last := all + for _, d := range slices.Sorted(maps.Keys(us.ActionsLE)) { + if d == math.MaxInt64 { + continue // skip infinity + } + c := us.ActionsLE[d] + srv.logf(" <=%v: %v", durFmt(d), c) + if last == c { + break + } + last = c + } + } + + // First clean things that are just too old. + if srv.maxAge > 0 { + if toDelete := all.Count - us.ActionsLE[srv.maxAge].Count; toDelete > 0 { + srv.logf("Cleaning %d objects older than %v ...", toDelete, durFmt(srv.maxAge)) + candidates, err := srv.cleanCandidates(srv.maxAge, toDelete+1) + if err != nil { + return zero, fmt.Errorf("getting clean candidates: %v", err) + } + blobIDs := make([]int64, 0, len(candidates)) + var sumSize int64 + for _, c := range candidates { + blobIDs = append(blobIDs, c.BlobID) + sumSize += c.BlobSize + } + if err := srv.deleteBlobs(blobIDs...); err != nil { + return zero, fmt.Errorf("deleting old blobs: %v", err) + } + all.Count -= int64(len(candidates)) + all.Size -= sumSize + ret.Count += int64(len(candidates)) + ret.Size += sumSize + } + } + + for srv.maxSize > 0 && all.Size > srv.maxSize { + toClean := all.Size - srv.maxSize + if srv.verbose { + srv.logf("need to clean %v to get under max size of %v ...", + bytesFmt(toClean), bytesFmt(srv.maxSize)) + } + + var batchBytes int64 + var blobIDs []int64 + candidates, err := srv.cleanCandidates(0, 10000) + if err != nil { + return zero, fmt.Errorf("getting clean candidates: %v", err) + } + for _, c := range candidates { + blobIDs = append(blobIDs, c.BlobID) + batchBytes += c.BlobSize + if batchBytes >= toClean { + break + } + } + if err := srv.deleteBlobs(blobIDs...); err != nil { + return zero, fmt.Errorf("deleting old blobs: %v", err) + } + ret.Count += int64(len(blobIDs)) + ret.Size += batchBytes + all.Count -= int64(len(blobIDs)) + all.Size -= batchBytes + + if len(blobIDs) == len(candidates) { + // We didn't find enough candidates to delete. + // Just stop here. + srv.logf("[unexpected] didn't find enough candidates to delete") + break + } + } + + return ret, nil +} + +func (srv *server) runCleanLoop() { + for { + select { + case <-srv.shutdownCtx.Done(): + return + case <-time.After(5 * time.Minute): + } + + us, err := srv.usageStats() + if err != nil { + srv.logf("error getting usage stats: %v", err) + continue + } + + res, err := srv.cleanOldObjects(us) + if err != nil { + srv.logf("error cleaning old objects: %v", err) + continue + } + if res.Count > 0 { + srv.logf("cleaned %v", res) + srv.usageStats() // for side effect of updating lastUsage + } + + } +} + +func durFmt(d time.Duration) string { + days := int(d.Hours() / 24) + if days > 0 { + return fmt.Sprintf("%dd", days) + } + return d.String() +} + +func bytesFmt(n int64) string { + if n >= 1<<30 { + return fmt.Sprintf("%.1f GiB", float64(n)/(1<<30)) + } + if n >= 1<<20 { + return fmt.Sprintf("%.1f MiB", float64(n)/(1<<20)) + } + if n >= 1<<10 { + return fmt.Sprintf("%.1f KiB", float64(n)/(1<<10)) + } + return fmt.Sprintf("%d bytes", n) +} + +func (srv *server) serveUsage(w http.ResponseWriter, r *http.Request) { + if r.Method == "POST" { + // For side effect of updating lastUsage. + _, err := srv.usageStats() + if err != nil { + http.Error(w, "error getting usage stats: "+err.Error(), http.StatusInternalServerError) + return + } + } + + us := srv.lastUsage.Load() + if us == nil { + http.Error(w, "no usage stats available", http.StatusInternalServerError) + return + } + + // Print out an HTML table of the usage stats, sorted by age. + w.Header().Set("Content-Type", "text/html; charset=utf-8") + fmt.Fprintf(w, "

gocached usage stats

\n") + fmt.Fprintf(w, "

Current usage: %v of limit %v

\n", + us.All(), bytesFmt(srv.maxSize)) + + fmt.Fprintf(w, "\n") + fmt.Fprintf(w, "\n") + for _, d := range slices.Sorted(maps.Keys(us.ActionsLE)) { + var title string + if d == math.MaxInt64 { + title = "all" + } else { + title = "<= " + durFmt(d) + } + c := us.ActionsLE[d] + fmt.Fprintf(w, "\n", + title, c.Count, bytesFmt(c.Size)) + } + fmt.Fprintf(w, "
AgeCountSize
%s%d%s
\n") } // expvarCounterMetric is a Prometheus counter metric backed by an expvar.Int. diff --git a/cmd/gocached/gocached_test.go b/cmd/gocached/gocached_test.go index e6385d9..918cbd6 100644 --- a/cmd/gocached/gocached_test.go +++ b/cmd/gocached/gocached_test.go @@ -7,6 +7,8 @@ import ( "math" "net/http/httptest" "os" + "path/filepath" + "slices" "strings" "sync" "testing" @@ -58,6 +60,45 @@ func (t *tester) mkClient() *cachers.HTTPClient { } } +func (st *tester) usageStats() *usageStats { + st.t.Helper() + stats, err := st.srv.usageStats() + if err != nil { + st.t.Fatalf("usageStats: %v", err) + } + return stats +} + +func (st *tester) cleanOldObjects() countAndSize { + st.t.Helper() + stats, err := st.srv.cleanOldObjects(st.usageStats()) + if err != nil { + st.t.Fatalf("cleanOldObjects: %v", err) + } + return stats +} + +func (st *tester) diskFiles() []string { + st.t.Helper() + var ret []string + + err := filepath.Walk(st.srv.dir, func(path string, fi os.FileInfo, err error) error { + if err != nil { + return err + } + if !fi.Mode().IsRegular() || strings.HasPrefix(fi.Name(), ".") || strings.HasPrefix(fi.Name(), "gocached") { + return nil + } + ret = append(ret, fi.Name()) + return nil + }) + if err != nil { + st.t.Fatalf("Walk: %v", err) + } + slices.Sort(ret) + return ret +} + // wantMetric is a helper to check an expvar.Int metric and reset it // for future tests. func (st *tester) wantMetric(m *expvar.Int, want int64) { @@ -195,7 +236,6 @@ func TestServer(t *testing.T) { ActionsLE: map[time.Duration]countAndSize{ 24 * time.Hour: {Count: 1, Size: 9}, 48 * time.Hour: {Count: 1, Size: 9}, - 72 * time.Hour: {Count: 3, Size: 1034}, 96 * time.Hour: {Count: 3, Size: 1034}, 168 * time.Hour: {Count: 3, Size: 1034}, 336 * time.Hour: {Count: 3, Size: 1034}, @@ -228,7 +268,7 @@ func TestCleanCandidates(t *testing.T) { tests := []struct { maxAge time.Duration - limit int + limit int64 want []cleanCandidate }{ { @@ -271,3 +311,75 @@ func TestCleanCandidates(t *testing.T) { }) } } + +func TestCleanOldObjectsByAge(t *testing.T) { + st := newServerTester(t) + st.srv.maxAge = 24 * time.Hour + + // Populate some data. + c1 := st.mkClient() + st.wantPut(c1, "0001", "9901", strings.Repeat("x", smallObjectSize+1)) + st.advanceClock(25 * time.Hour) + st.wantPut(c1, "0002", "9902", strings.Repeat("x", smallObjectSize+2)) + st.wantPut(c1, "0003", "9903", "small") + smallLen := int64(len("small")) + + st1 := st.usageStats() + if all, want := st1.All(), (countAndSize{Count: 3, Size: smallObjectSize*2 + 3 + smallLen}); all != want { + t.Errorf("usageStats: %v; want %v", all, want) + } + if got, want := st.diskFiles(), []string{"333092a3daf718ed8f38a94e302df139edd4e3b5da4239a497995683942cf28c", "c6d8e9905300876046729949cc95c2385221270d389176f7234fe7ac00c4e430"}; !slices.Equal(got, want) { + t.Errorf("diskFiles: %v; want %v", got, want) + } + + clean1 := st.cleanOldObjects() + if clean1.Count != 1 || clean1.Size != smallObjectSize+1 { + t.Errorf("cleanOldObjects got %v, want {Count: 1, Size: %d}", clean1, smallObjectSize+1) + } + clean2 := st.cleanOldObjects() + if clean2.Count != 0 || clean2.Size != 0 { + t.Errorf("cleanOldObjects got %v, want {Count: 0, Size: 0}", clean2) + } + + st2 := st.usageStats() + if all, want := st2.All(), (countAndSize{Count: 2, Size: smallObjectSize + 2 + smallLen}); all != want { + t.Errorf("usageStats after clean: %v; want %v", all, want) + } + if got, want := st.diskFiles(), []string{"333092a3daf718ed8f38a94e302df139edd4e3b5da4239a497995683942cf28c"}; !slices.Equal(got, want) { + t.Errorf("diskFiles after clean: %v; want %v", got, want) + } +} + +func TestCleanOldObjectsBySize(t *testing.T) { + st := newServerTester(t) + + // Populate some data. + c1 := st.mkClient() + st.wantPut(c1, "0001", "9901", "1") + st.advanceClock(time.Second) + st.wantPut(c1, "0002", "9902", "22") + st.advanceClock(time.Second) + st.wantPut(c1, "0003", "9903", "333") + st.advanceClock(time.Second) + st.wantPut(c1, "0004", "9904", "4444") + st.advanceClock(time.Second) + + st1 := st.usageStats() + if all, want := st1.All(), (countAndSize{Count: 4, Size: 10}); all != want { + t.Errorf("usageStats: %v; want %v", all, want) + } + + clean1 := st.cleanOldObjects() + if clean1.Count != 0 || clean1.Size != 0 { + t.Errorf("cleanOldObjects got %v, want no clean", clean1) + } + + st.srv.maxSize = 8 // the only way get to 8 or under is by deleting "1" and "22" (3 bytes) + + if got, want := st.cleanOldObjects(), (countAndSize{Count: 2, Size: 3}); got != want { + t.Errorf("cleanOldObjects got %v, want %v", got, want) + } + if got, want := st.usageStats().All(), (countAndSize{Count: 2, Size: 7}); got != want { + t.Errorf("usageStats: %v; want %v", got, want) + } +} From 0aa66680ac193cf6118bef51155f0279835b5afe Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Wed, 13 Aug 2025 13:54:50 -0700 Subject: [PATCH 21/67] cmd/gocached: add dup put and eviction metrics Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@48dd315a01ace0b23a640db86f5a3cf68bf06af0 --- cmd/gocached/gocached.go | 31 ++++++++++++++++++++++++++++--- 1 file changed, 28 insertions(+), 3 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index b94bbca..657143a 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -258,10 +258,13 @@ type server struct { GetHitsInline expvar.Int `type:"counter" name:"get_hits_inline" help:"cache hits served from inline database storage (small objects)"` GetErrs expvar.Int `type:"counter" name:"get_errs" help:"number of gocache get request errors"` Puts expvar.Int `type:"counter" name:"puts" help:"total number of gocache put requests"` + PutsDup expvar.Int `type:"counter" name:"puts_dup" help:"total number of gocache put requests that are duplicates of a mapping we already had"` PutsBytes expvar.Int `type:"counter" name:"put_bytes" help:"total bytes added from gocache puts"` PutsInline expvar.Int `type:"counter" name:"put_inline" help:"subset of gocached_puts that were stored inline (small objects)"` BlobCount expvar.Int `type:"gauge" name:"blob_count" help:"number of blobs currently stored in the cache"` BlobBytes expvar.Int `type:"gauge" name:"blob_bytes" help:"sum of blob sizes currently stored in the cache"` + EvictedBlobs expvar.Int `type:"counter" name:"evicted_blobs" help:"number of blobs evicted from the cache"` + EvictedBytes expvar.Int `type:"counter" name:"evicted_bytes" help:"number of bytes evicted from the cache"` } func (srv *server) now() time.Time { @@ -516,7 +519,7 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { if sha256hex != outputID { altObjectID = outputID } - _, err = s.db.Exec(`INSERT OR IGNORE INTO Actions (NamespaceID, ActionID, BlobID, AltOutputID, CreateTime, AccessTime) + res, err := s.db.Exec(`INSERT OR IGNORE INTO Actions (NamespaceID, ActionID, BlobID, AltOutputID, CreateTime, AccessTime) VALUES (?, ?, ?, ?, ?, ?)`, namespace, actionID, @@ -531,6 +534,17 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { return } + affected, err := res.RowsAffected() + if err != nil { + s.logf("Actions rows affected error: %v", err) + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + + if affected == 0 { + s.PutsDup.Add(1) + } + s.Puts.Add(1) s.PutsBytes.Add(r.ContentLength) if smallData != nil { @@ -744,11 +758,14 @@ func (srv *server) deleteBlobs(blobIDs ...int64) error { } defer tx.Rollback() + var sumBytes int64 for _, blobID := range blobIDs { var sha256Hex string - if err := tx.QueryRow("SELECT SHA256 FROM Blobs WHERE BlobID = ?", blobID).Scan(&sha256Hex); err != nil && !errors.Is(err, sql.ErrNoRows) { + var blobSize int64 + if err := tx.QueryRow("SELECT SHA256, BlobSize FROM Blobs WHERE BlobID = ?", blobID).Scan(&sha256Hex, &blobSize); err != nil && !errors.Is(err, sql.ErrNoRows) { return fmt.Errorf("querying blob SHA256: %w", err) } + sumBytes += blobSize if _, err := tx.Exec("DELETE FROM Blobs WHERE BlobID = ?", blobID); err != nil { return fmt.Errorf("deleting blob: %w", err) } @@ -762,7 +779,14 @@ func (srv *server) deleteBlobs(blobIDs ...int64) error { } } } - return tx.Commit() + if err := tx.Commit(); err != nil { + return err + } + + srv.EvictedBlobs.Add(int64(len(blobIDs))) + srv.EvictedBytes.Add(sumBytes) + + return nil } func (srv *server) cleanOldObjects(us *usageStats) (countAndSize, error) { @@ -833,6 +857,7 @@ func (srv *server) cleanOldObjects(us *usageStats) (countAndSize, error) { if err := srv.deleteBlobs(blobIDs...); err != nil { return zero, fmt.Errorf("deleting old blobs: %v", err) } + ret.Count += int64(len(blobIDs)) ret.Size += batchBytes all.Count -= int64(len(blobIDs)) From 45703a872ec64992ab7255c31e54b2d4fc3d991b Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Fri, 15 Aug 2025 11:25:11 -0700 Subject: [PATCH 22/67] cmd/gocached: add PutErr metric, serialize SQLite writes Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@1346ba90474a41c0f3f9c01ac3284c733f0d1336 --- cmd/gocached/gocached.go | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index 657143a..c6119de 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -53,6 +53,7 @@ import ( "reflect" "slices" "strings" + "sync" "sync/atomic" "time" @@ -246,6 +247,12 @@ type server struct { shutdownCtx context.Context shutdownCancel context.CancelFunc + // sqliteWriteMu serializes access to SQLite. In theory the SQLite driver + // should serialize access with our 5000ms busy timeout, but empirically we + // sometimes seen DB busy errors. Just serialize it explicitly out of + // laziness for now. + sqliteWriteMu sync.Mutex + lastUsage atomic.Pointer[usageStats] // Metrics @@ -258,6 +265,7 @@ type server struct { GetHitsInline expvar.Int `type:"counter" name:"get_hits_inline" help:"cache hits served from inline database storage (small objects)"` GetErrs expvar.Int `type:"counter" name:"get_errs" help:"number of gocache get request errors"` Puts expvar.Int `type:"counter" name:"puts" help:"total number of gocache put requests"` + PutErrs expvar.Int `type:"counter" name:"put_errs" help:"number of gocache put request errors"` PutsDup expvar.Int `type:"counter" name:"puts_dup" help:"total number of gocache put requests that are duplicates of a mapping we already had"` PutsBytes expvar.Int `type:"counter" name:"put_bytes" help:"total bytes added from gocache puts"` PutsInline expvar.Int `type:"counter" name:"put_inline" help:"subset of gocached_puts that were stored inline (small objects)"` @@ -500,6 +508,9 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { sha256hex := fmt.Sprintf("%x", hasher.Sum(nil)) blobSize := r.ContentLength + s.sqliteWriteMu.Lock() + defer s.sqliteWriteMu.Unlock() + var blobID int64 err := s.db.QueryRow(`INSERT INTO Blobs (SHA256, BlobSize, SmallData) VALUES (?, ?, ?) @@ -508,6 +519,7 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { `, sha256hex, blobSize, smallData).Scan(&blobID) if err != nil { s.logf("Blobs insert error: %v", err) + s.PutErrs.Add(1) http.Error(w, err.Error(), http.StatusInternalServerError) return } @@ -530,6 +542,7 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { ) if err != nil { s.logf("Actions insert error: %v", err) + s.PutErrs.Add(1) http.Error(w, err.Error(), http.StatusInternalServerError) return } @@ -537,6 +550,7 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { affected, err := res.RowsAffected() if err != nil { s.logf("Actions rows affected error: %v", err) + s.PutErrs.Add(1) http.Error(w, err.Error(), http.StatusInternalServerError) return } @@ -752,6 +766,9 @@ func (s *server) cleanCandidates(olderThan time.Duration, limit int64) ([]cleanC } func (srv *server) deleteBlobs(blobIDs ...int64) error { + srv.sqliteWriteMu.Lock() + defer srv.sqliteWriteMu.Unlock() + tx, err := srv.db.Begin() if err != nil { return fmt.Errorf("delete blob Begin: %w", err) From 6f8051d80a428cc67e62d882df9412f0e5038ac6 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Wed, 20 Aug 2025 07:58:56 -0700 Subject: [PATCH 23/67] cmd/gocached: use sqliteWriteMu in another spot Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@01466c7bb897693b177bd24bd8ffbe2fe095330e --- cmd/gocached/gocached.go | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index c6119de..a051272 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -166,6 +166,9 @@ func openDB(dbDir string) (*sql.DB, error) { if err != nil { return nil, err } + db.SetMaxOpenConns(4) + db.SetMaxIdleConns(4) + db.SetConnMaxLifetime(0) // no limit if _, err := db.Exec(schema); err != nil { return nil, err } @@ -388,7 +391,9 @@ func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request) { if accessTime < now-relAtimeSeconds { // TODO(bradfitz): do this async? not worth blocking the caller. // But we need a mechanism for tests to wait on async work. + srv.sqliteWriteMu.Lock() _, err := srv.db.Exec("UPDATE Actions SET AccessTime = ? WHERE ActionID = ?", now, actionID) + srv.sqliteWriteMu.Unlock() if err != nil { srv.logf("Update AccessTime error: %v", err) httpErr("internal server error", http.StatusInternalServerError) From 897165f6bdd73fe9ecec1f17b92a697140f64062 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Tue, 23 Sep 2025 11:44:01 -0700 Subject: [PATCH 24/67] cmd/gocached: separate out debug port, add pprof info So we can configure it so guest VMs don't have access to debug info. Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@2513dd2d45325062cdee0fc7d9089734b5a5b24a --- cmd/gocached/gocached.go | 70 +++++++++++++++++++++++++++++----------- 1 file changed, 51 insertions(+), 19 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index a051272..b260f21 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -47,7 +47,9 @@ import ( "log" "maps" "math" + "net" "net/http" + "net/http/pprof" "os" "path/filepath" "reflect" @@ -70,9 +72,10 @@ import ( const smallObjectSize = 1 << 10 var ( - dir = flag.String("cache-dir", "", "cache directory, if empty defaults to /gocached") - verbose = flag.Bool("verbose", false, "be verbose") - listen = flag.String("listen", ":31364", "listen address") + dir = flag.String("cache-dir", "", "cache directory, if empty defaults to /gocached") + verbose = flag.Bool("verbose", false, "be verbose") + listen = flag.String("listen", ":31364", "listen address for the build-facing HTTP server") + debugListen = flag.String("debug-listen", "", "if non-empty, listen address for the debug HTTP server (pprof, metrics, etc)") maxSize = flag.Int("max-size-gb", 50, "maximum size of the cache in GiB; 0 means no limit") maxAge = flag.Int("max-age-days", 60, "maximum age of objects in the cache in days; 0 means no limit") @@ -114,6 +117,16 @@ func main() { log.Printf("gocached: cleaned %v", res) } + if *debugListen != "" { + debugLn, err := net.Listen("tcp", *debugListen) + if err != nil { + log.Fatalf("debug listen: %v", err) + } + go func() { + log.Fatal(http.Serve(debugLn, http.HandlerFunc(srv.ServeHTTPDebug))) + }() + } + go srv.runCleanLoop() log.Printf("gocached: listening on %s ...", *listen) @@ -285,38 +298,57 @@ func (srv *server) now() time.Time { return time.Now() } -func (srv *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { +func (srv *server) ServeHTTPDebug(w http.ResponseWriter, r *http.Request) { if srv.verbose { - srv.logf("ServeHTTP: %s %s", r.Method, r.RequestURI) - } - if r.URL.Path == "/usage" { - srv.serveUsage(w, r) - return - } - if r.Method == "PUT" { - srv.handlePut(w, r) - return - } - if r.Method != "GET" && r.Method != "HEAD" { - http.Error(w, "bad method", http.StatusBadRequest) - return + srv.logf("ServeHTTPDebug: %s %s", r.Method, r.RequestURI) } switch { - case strings.HasPrefix(r.URL.Path, "/action/"): - srv.handleGetAction(w, r) case r.URL.Path == "/": w.Header().Set("Content-Type", "text/html; charset=utf-8") io.WriteString(w, "

gocached

") io.WriteString(w, "

This is a shared Go build cache server, hit by GOCACHEPROG clients.

") io.WriteString(w, "

See /usage for usage stats.

") io.WriteString(w, "

See /metrics for Prometheus metrics.

") + io.WriteString(w, "

See /debug/pprof/ for pprof

") + io.WriteString(w, "

See /debug/pprof/goroutine?debug=2 - full goroutines

") + case r.URL.Path == "/usage": + srv.serveUsage(w, r) case r.URL.Path == "/metrics": srv.metricsHandler.ServeHTTP(w, r) + case strings.HasPrefix(r.URL.Path, "/debug/pprof/profile"): + pprof.Profile(w, r) + case strings.HasPrefix(r.URL.Path, "/debug/pprof/cmdline"): + pprof.Cmdline(w, r) + case strings.HasPrefix(r.URL.Path, "/debug/pprof/symbol"): + pprof.Symbol(w, r) + case strings.HasPrefix(r.URL.Path, "/debug/pprof/trace"): + pprof.Trace(w, r) + case strings.HasPrefix(r.URL.Path, "/debug/pprof/"): + pprof.Index(w, r) default: http.Error(w, "not found", http.StatusNotFound) } } +func (srv *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if srv.verbose { + srv.logf("ServeHTTP: %s %s", r.Method, r.RequestURI) + } + if r.Method == "PUT" { + srv.handlePut(w, r) + return + } + if r.Method != "GET" && r.Method != "HEAD" { + http.Error(w, "bad method", http.StatusBadRequest) + return + } + if strings.HasPrefix(r.URL.Path, "/action/") { + srv.handleGetAction(w, r) + return + } + http.Error(w, "not found", http.StatusNotFound) +} + func getHexSuffix(r *http.Request, prefix string) (hexSuffix string, ok bool) { hexSuffix, _ = strings.CutPrefix(r.RequestURI, prefix) if !validHex(hexSuffix) { From c6421676e233f2600f0858d1d2deddde316c283b Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Wed, 24 Sep 2025 07:41:20 -0700 Subject: [PATCH 25/67] cachers: consume 404 response bodies to allow HTTP conn reuse Fixes bradfitz/go-tool-cache#15 Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@e2eb857b506e7b24d7cdb65a66349ed43a34783f --- cmd/gocached/gocached_test.go | 44 +++++++++++++++++++++++++++++++++++ gocache/cachers/http.go | 24 +++++++++++++++++-- 2 files changed, 66 insertions(+), 2 deletions(-) diff --git a/cmd/gocached/gocached_test.go b/cmd/gocached/gocached_test.go index 918cbd6..be1da1d 100644 --- a/cmd/gocached/gocached_test.go +++ b/cmd/gocached/gocached_test.go @@ -5,12 +5,15 @@ import ( "expvar" "fmt" "math" + "net" + "net/http" "net/http/httptest" "os" "path/filepath" "slices" "strings" "sync" + "sync/atomic" "testing" "time" @@ -151,6 +154,21 @@ func (st *tester) wantGet(c *cachers.HTTPClient, actionID, outputID, wantVal str } } +func (st *tester) wantGetMiss(c *cachers.HTTPClient, actionID string) { + ctx := context.Background() + st.t.Helper() + gotOutputID, diskPath, err := c.Get(ctx, actionID) + if err != nil { + st.t.Fatalf("Get: %v", err) + } + if gotOutputID != "" { + st.t.Errorf("Get got outputID %q; want empty", gotOutputID) + } + if diskPath != "" { + st.t.Fatalf("Get returned disk path %q; want empty", diskPath) + } +} + func newServerTester(t testing.TB) *tester { st := &tester{ t: t, @@ -383,3 +401,29 @@ func TestCleanOldObjectsBySize(t *testing.T) { t.Errorf("usageStats: %v; want %v", got, want) } } + +func TestClientConnReuse(t *testing.T) { + st := newServerTester(t) + + var numDials atomic.Int32 + tr := http.DefaultTransport.(*http.Transport).Clone() + tr.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) { + num := numDials.Add(1) + t.Logf("DialContext #%d for %s %s", num, network, addr) + var std net.Dialer + return std.DialContext(ctx, network, addr) + } + t.Cleanup(func() { tr.CloseIdleConnections() }) + + c1 := st.mkClient() + c1.HTTPClient = &http.Client{Transport: tr} + const missAction = "0001" + st.wantGetMiss(c1, missAction) + st.wantGetMiss(c1, missAction) + st.wantGetMiss(c1, missAction) + st.wantPut(c1, "0001", "9901", "1") + st.wantGet(c1, "0001", "9901", "1") + if got := numDials.Load(); got != 1 { + t.Errorf("numDials = %d; want 1", got) + } +} diff --git a/gocache/cachers/http.go b/gocache/cachers/http.go index 7ab9b99..f245cff 100644 --- a/gocache/cachers/http.go +++ b/gocache/cachers/http.go @@ -39,6 +39,19 @@ func (c *HTTPClient) httpClient() *http.Client { return http.DefaultClient } +// tryDrainResponse reads and throws away a small bounded amount of data from +// res.Body. This is a best-effort attempt to allow connection reuse. (Go's +// HTTP/1 Transport won't reuse a TCP connection unless you fully consume HTTP +// responses) +func tryDrainResponse(res *http.Response) { + io.CopyN(io.Discard, res.Body, 4<<10) +} + +func tryReadErrorMessage(res *http.Response) []byte { + msg, _ := io.ReadAll(io.LimitReader(res.Body, 4<<10)) + return msg +} + func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPath string, err error) { outputID, diskPath, err = c.Disk.Get(ctx, actionID) if err == nil && outputID != "" { @@ -59,10 +72,13 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa return "", "", err } defer res.Body.Close() + defer tryDrainResponse(res) if res.StatusCode == http.StatusNotFound { return "", "", nil } if res.StatusCode != http.StatusOK { + msg := tryReadErrorMessage(res) + log.Printf("error GET /action/%s: %v, %s", actionID, res.Status, msg) return "", "", fmt.Errorf("unexpected GET /action/%s status %v", actionID, res.Status) } @@ -98,10 +114,13 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa return "", "", err } defer res.Body.Close() + defer tryDrainResponse(res) if res.StatusCode == http.StatusNotFound { return "", "", nil } if res.StatusCode != http.StatusOK { + msg := tryReadErrorMessage(res) + log.Printf("error GET /output/%s: %v, %s", outputID, res.Status, msg) return "", "", fmt.Errorf("unexpected GET /output/%s status %v", outputID, res.Status) } if res.ContentLength == -1 { @@ -152,8 +171,9 @@ func (c *HTTPClient) Put(ctx context.Context, actionID, outputID string, size in } defer res.Body.Close() if res.StatusCode != http.StatusNoContent { - all, _ := io.ReadAll(io.LimitReader(res.Body, 4<<10)) - return "", fmt.Errorf("unexpected PUT /%s/%s status %v: %s", actionID, outputID, res.Status, all) + msg := tryReadErrorMessage(res) + log.Printf("error PUT /%s/%s: %v, %s", actionID, outputID, res.Status, msg) + return "", fmt.Errorf("unexpected PUT /%s/%s status %v", actionID, outputID, res.Status) } v := <-diskPutCh if err, ok := v.(error); ok { From eae2c3b4d4fa2dcdd1e46446879537b632a0d58e Mon Sep 17 00:00:00 2001 From: Tom Proctor Date: Fri, 17 Oct 2025 10:09:31 +0100 Subject: [PATCH 26/67] cachers,cmd/go-cacher-server: fix unexpected http 404s (bradfitz/go-tool-cache#18) Running with a cold disk cache for go-cacher and a warm disk cache for go-cacher-server, I was getting no cache hits. This is due to a missing component in the output filename when serving up what go-cacher-server thinks is a cache hit; it was trying to serve from `/o-` instead of the real file location of `//o-` where o2 is the first 2 characters of the hex string. To help with consistency, I consolidated all places that need to calculate an action or output file path into 2 shared functions, but there is only one functional change. Tested via the benchmarks I included in bradfitz/go-tool-cache#17, but it's probably worth investing in some unit tests as well as we start to use this more seriously. Signed-off-by: Tom Proctor Migrated-from: bradfitz/go-tool-cache@c868ae47093aeb60feb2408548243b652cb06e46 --- gocache/cachers/disk.go | 39 +++++++++++++++++---------------------- 1 file changed, 17 insertions(+), 22 deletions(-) diff --git a/gocache/cachers/disk.go b/gocache/cachers/disk.go index 3d64e67..b78937f 100644 --- a/gocache/cachers/disk.go +++ b/gocache/cachers/disk.go @@ -39,9 +39,8 @@ func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPa if !validHex(actionID) { return "", "", fmt.Errorf("actionID must be valid hex strings") } - action2 := actionID[:2] - actionFile := filepath.Join(dc.Dir, action2, fmt.Sprintf("a-%s", actionID)) + actionFile := dc.ActionFilename(actionID) ij, err := os.ReadFile(actionFile) if err != nil { if os.IsNotExist(err) { @@ -61,30 +60,28 @@ func (dc *DiskCache) Get(ctx context.Context, actionID string) (outputID, diskPa // Protect against malicious non-hex OutputID on disk return "", "", nil } - output2 := ie.OutputID[:2] - return ie.OutputID, filepath.Join(dc.Dir, output2, fmt.Sprintf("o-%v", ie.OutputID)), nil + return ie.OutputID, dc.OutputFilename(ie.OutputID), nil } func (dc *DiskCache) OutputFilename(outputID string) string { - if len(outputID) < 4 || len(outputID) > 1000 { + if !validHex(outputID) { return "" } - for i := range outputID { - b := outputID[i] - if b >= '0' && b <= '9' || b >= 'a' && b <= 'f' { - continue - } + return filepath.Join(dc.Dir, outputID[:2], fmt.Sprintf("o-%s", outputID)) +} + +func (dc *DiskCache) ActionFilename(actionID string) string { + if !validHex(actionID) { return "" } - return filepath.Join(dc.Dir, fmt.Sprintf("o-%s", outputID)) + return filepath.Join(dc.Dir, actionID[:2], fmt.Sprintf("a-%s", actionID)) } func validHex(x string) bool { if len(x) < 4 || len(x) > 100 { return false } - for i := range len(x) { - b := x[i] + for _, b := range x { if b >= '0' && b <= '9' || b >= 'a' && b <= 'f' { continue } @@ -106,9 +103,10 @@ func (dc *DiskCache) Put(ctx context.Context, actionID, outputID string, size in return "", errors.New("outputID must be hex") } - action2, output2 := actionID[:2], outputID[:2] - outputDir := filepath.Join(dc.Dir, output2) - actionDir := filepath.Join(dc.Dir, action2) + actionFile := dc.ActionFilename(actionID) + outputFile := dc.OutputFilename(outputID) + actionDir := filepath.Dir(actionFile) + outputDir := filepath.Dir(outputFile) if err := os.MkdirAll(actionDir, 0755); err != nil { return "", fmt.Errorf("failed to create action directory: %w", err) @@ -117,17 +115,15 @@ func (dc *DiskCache) Put(ctx context.Context, actionID, outputID string, size in return "", fmt.Errorf("failed to create output directory: %w", err) } - file := filepath.Join(outputDir, fmt.Sprintf("o-%s", outputID)) - // Special case empty files; they're both common and easier to do race-free. if size == 0 { - zf, err := os.OpenFile(file, os.O_CREATE|os.O_RDWR, 0644) + zf, err := os.OpenFile(outputFile, os.O_CREATE|os.O_RDWR, 0644) if err != nil { return "", err } zf.Close() } else { - wrote, err := writeAtomic(file, body) + wrote, err := writeAtomic(outputFile, body) if err != nil { return "", err } @@ -145,11 +141,10 @@ func (dc *DiskCache) Put(ctx context.Context, actionID, outputID string, size in if err != nil { return "", err } - actionFile := filepath.Join(actionDir, fmt.Sprintf("a-%s", actionID)) if _, err := writeAtomic(actionFile, bytes.NewReader(ij)); err != nil { return "", err } - return file, nil + return outputFile, nil } func writeAtomic(dest string, r io.Reader) (int64, error) { From 504a80862cb8250ec76edbf54b8b146ea3ab3e28 Mon Sep 17 00:00:00 2001 From: Tom Proctor Date: Thu, 6 Nov 2025 04:04:01 +0000 Subject: [PATCH 27/67] cmd/gocached: tighten fs permissions and http errors (bradfitz/go-tool-cache#21) Some minor fix-ups to avoid unnecessarily wide fs permissions or HTTP error messages that could accidentally expose details that we expect to stay internal. Signed-off-by: Tom Proctor Migrated-from: bradfitz/go-tool-cache@38eb3410a6ae4f2c2336f5b6e309734d0033ee04 --- cmd/gocached/gocached.go | 23 +++++++++++++---------- 1 file changed, 13 insertions(+), 10 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index b260f21..cda8539 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -92,7 +92,7 @@ func main() { log.Printf("Defaulting to cache dir %v ...", d) *dir = d } - if err := os.MkdirAll(*dir, 0755); err != nil { + if err := os.MkdirAll(*dir, 0750); err != nil { log.Fatal(err) } @@ -413,7 +413,7 @@ func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request) { return } srv.logf("QueryRow error: %v", err) - httpErr("Query: "+err.Error(), http.StatusInternalServerError) + httpErr("QueryRow error", http.StatusInternalServerError) return } @@ -459,7 +459,8 @@ func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request) { rc, err := srv.getObjectFromDiskOrPeer(ctx, sha256hex) if err != nil { - httpErr(err.Error(), http.StatusInternalServerError) + srv.logf("Get object error: %v", actionID, err) + httpErr("Get object error", http.StatusInternalServerError) return } if rc == nil { @@ -481,7 +482,7 @@ func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request) { // exists but is not stored in SQLite. // // It returns (nil, nil) on miss. -func (srv *server) getObjectFromDiskOrPeer(ctx context.Context, sha256hex string) (rc io.ReadCloser, err error) { +func (srv *server) getObjectFromDiskOrPeer(_ context.Context, sha256hex string) (rc io.ReadCloser, err error) { if len(sha256hex) != sha256.Size*2 { return nil, fmt.Errorf("invalid sha256hex %q", sha256hex) } @@ -525,7 +526,8 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { var err error smallData, err = io.ReadAll(hashingBody) if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) + s.logf("Read content error: %v", err) + http.Error(w, "Read content error", http.StatusInternalServerError) return } if int64(len(smallData)) != r.ContentLength { @@ -537,7 +539,8 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { } else { // For larger objects, we store them on disk. if err := s.writeDiskBlob(r.ContentLength, hashingBody); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) + s.logf("Write disk blob error: %v", err) + http.Error(w, "Write disk blob error", http.StatusInternalServerError) return } } @@ -557,7 +560,7 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { if err != nil { s.logf("Blobs insert error: %v", err) s.PutErrs.Add(1) - http.Error(w, err.Error(), http.StatusInternalServerError) + http.Error(w, "Blobs insert error", http.StatusInternalServerError) return } @@ -580,7 +583,7 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { if err != nil { s.logf("Actions insert error: %v", err) s.PutErrs.Add(1) - http.Error(w, err.Error(), http.StatusInternalServerError) + http.Error(w, "Actions insert error", http.StatusInternalServerError) return } @@ -588,7 +591,7 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { if err != nil { s.logf("Actions rows affected error: %v", err) s.PutErrs.Add(1) - http.Error(w, err.Error(), http.StatusInternalServerError) + http.Error(w, "Actions rows affected error", http.StatusInternalServerError) return } @@ -638,7 +641,7 @@ func (s *server) writeDiskBlob(size int64, r io.Reader) (err error) { hasher.Sum(hash[:0]) target := s.sha256Filepath(hash) - if err := os.MkdirAll(filepath.Dir(target), 0755); err != nil { + if err := os.MkdirAll(filepath.Dir(target), 0750); err != nil { return err } return os.Rename(tf.Name(), target) From a4563019d3944b66860283aa62f3de4b5f49e4b8 Mon Sep 17 00:00:00 2001 From: Tom Proctor Date: Tue, 28 Oct 2025 12:35:33 +0000 Subject: [PATCH 28/67] cachers,cmd/{go-cacher,gocached}: add JWT-based auth sessions Adds the ability to run gocached with a token exchange endpoint that starts a new session each time it receives a valid JWT. Each session is associated with an opaque access token and has session-based cache stats that can be retrieved out of band at the end of the build/test process. JWT issuer and claim support is intentionally narrow initially, only supporting a single issuer and exact-match string claims. The next planned follow-up is to support read/write for _all_ tokens, but only JWTs that satisfy the -global-jwt-claim requirements will be able to open a session that can write to the global namespace 0. All other sessions will write to their own namespace based on some other criteria. Signed-off-by: Tom Proctor Migrated-from: bradfitz/go-tool-cache@f8ee03e77fc8a7ff06bf1756a1835ec8353f0d46 --- cmd/gocached/gocached.go | 412 +++++++++++++++++++++++++++++-- cmd/gocached/gocached_test.go | 315 +++++++++++++++++++++++ cmd/gocached/internal/jwt/jwt.go | 201 +++++++++++++++ go.mod | 6 +- go.sum | 4 + gocache/cachers/http.go | 10 + 6 files changed, 920 insertions(+), 28 deletions(-) create mode 100644 cmd/gocached/internal/jwt/jwt.go diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index cda8539..4f46405 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -36,9 +36,11 @@ package main import ( "cmp" "context" + "crypto/rand" "crypto/sha256" "database/sql" "encoding/hex" + "encoding/json" "errors" "expvar" "flag" @@ -63,13 +65,23 @@ import ( "github.com/prometheus/client_golang/prometheus/collectors" "github.com/prometheus/client_golang/prometheus/promhttp" dto "github.com/prometheus/client_model/go" + "github.com/tailscale/tb/cmd/gocached/internal/jwt" _ "modernc.org/sqlite" ) -// smallObjectSize is the maximum size of an object that we store inline in the -// database, rather than on disk. Empirically, about half of objects are 1KB or -// smaller. -const smallObjectSize = 1 << 10 +const ( + // smallObjectSize is the maximum size of an object that we store inline in the + // database, rather than on disk. Empirically, about half of objects are 1KB or + // smaller. + smallObjectSize = 1 << 10 + + // tokenPrefix is the prefix for all gocached access tokens. + tokenPrefix = "gocached-token-" + + // gocachedAudience is the audience we require JWTs to have. Could be + // configurable in future, but for now just needs to be specific to gocached. + gocachedAudience = "gocached" +) var ( dir = flag.String("cache-dir", "", "cache directory, if empty defaults to /gocached") @@ -79,9 +91,32 @@ var ( maxSize = flag.Int("max-size-gb", 50, "maximum size of the cache in GiB; 0 means no limit") maxAge = flag.Int("max-age-days", 60, "maximum age of objects in the cache in days; 0 means no limit") + + jwtIssuer = flag.String("jwt-issuer", "", "the issuer to trust JWTs from; if set, all requests will require auth, and must set at least one -jwt-claim") + // See example GitHub token claims for what can be available: + // https://docs.github.com/en/actions/concepts/security/openid-connect + jwtClaims = make(jwtClaimValue) + globalJWTClaims = make(jwtClaimValue) ) +type jwtClaimValue map[string]string + +func (v jwtClaimValue) String() string { + return fmt.Sprintf("%v", map[string]string(v)) +} + +func (v jwtClaimValue) Set(s string) error { + claim, value, ok := strings.Cut(s, "=") + if !ok || claim == "" || value == "" { + return fmt.Errorf("bad claim %q, want x=y", s) + } + v[claim] = value + return nil +} + func main() { + flag.Var(&jwtClaims, "jwt-claim", "a claim in the form x=y that any JWT presented must have to start a session; may be specified more than once") + flag.Var(&globalJWTClaims, "global-jwt-claim", "an additional claim in the form x=y that a JWT must have to allow writing to the cache's global namespace; may be specified more than once") flag.Parse() if *dir == "" { d, err := os.UserCacheDir() @@ -127,6 +162,26 @@ func main() { }() } + if *jwtIssuer != "" { + if len(jwtClaims) == 0 { + log.Fatal("must specify --jwt-claim at least once when --jwt-issuer is set") + } + srv.jwtValidator = jwt.NewJWTValidator(*jwtIssuer, gocachedAudience) + if err := srv.jwtValidator.RunUpdateJWKSLoop(srv.shutdownCtx); err != nil { + log.Fatalf("failed to fetch JWKS for JWT validator: %v", err) + } + srv.jwtClaims = jwtClaims + + globalClaims := map[string]string{} + maps.Copy(globalClaims, jwtClaims) + maps.Copy(globalClaims, globalJWTClaims) + srv.globalJWTClaims = globalClaims + + log.Printf("gocached: using JWT issuer %q with claims %v, global claims %v", *jwtIssuer, srv.jwtClaims, srv.globalJWTClaims) + + go srv.runCleanSessionsLoop() + } + go srv.runCleanLoop() log.Printf("gocached: listening on %s ...", *listen) @@ -202,9 +257,10 @@ func newServer(dir string) (*server, error) { ) srv := &server{ - db: db, - dir: dir, - logf: log.Printf, + db: db, + dir: dir, + logf: log.Printf, + sessions: make(map[string]*sessionData), } srv.shutdownCtx, srv.shutdownCancel = context.WithCancel(context.Background()) srv.registerMetrics(reg) @@ -263,6 +319,13 @@ type server struct { shutdownCtx context.Context shutdownCancel context.CancelFunc + jwtValidator *jwt.Validator // nil unless -jwt-issuer flag is set + jwtClaims map[string]string // claims required for any JWT to start a session + globalJWTClaims map[string]string // additional claims required to write to global namespace + + sessionsMu sync.RWMutex // guards sessions + sessions map[string]*sessionData // maps access token -> session data. + // sqliteWriteMu serializes access to SQLite. In theory the SQLite driver // should serialize access with our 5000ms busy timeout, but empirically we // sometimes seen DB busy errors. Just serialize it explicitly out of @@ -289,6 +352,39 @@ type server struct { BlobBytes expvar.Int `type:"gauge" name:"blob_bytes" help:"sum of blob sizes currently stored in the cache"` EvictedBlobs expvar.Int `type:"counter" name:"evicted_blobs" help:"number of blobs evicted from the cache"` EvictedBytes expvar.Int `type:"counter" name:"evicted_bytes" help:"number of bytes evicted from the cache"` + Sessions expvar.Int `type:"gauge" name:"sessions" help:"number of active authenticated sessions"` + Auths expvar.Int `type:"counter" name:"auth_attempts" help:"number of successful token exchanges"` + AuthErrs expvar.Int `type:"counter" name:"auth_errs" help:"number of failed token exchanges"` +} + +// sessionData corresponds to a specific access token, and is only used if JWT +// auth is enabled. +type sessionData struct { + expiry time.Time // Session valid until. + globalNSWrite bool // Whether this session can write to the cache's global namespace. + claims map[string]any // Claims from the JWT used to create this session, stored for debug. + + mu sync.Mutex // Guards stats. + stats stats +} + +// stats holds per-request or per-session stats which get rolled up into server +// stats. See [server] struct for detailed definitions. +type stats struct { + LastUsed time.Time // Only applies to session stats. Last time the access token for this session was used. + Gets int64 + GetBytes int64 + GetHits int64 + GetAccessBumps int64 + GetHitsInline int64 + GetNanos int64 + GetErrs int64 + Puts int64 + PutErrs int64 + PutsDup int64 + PutsBytes int64 + PutsInline int64 + PutsNanos int64 } func (srv *server) now() time.Time { @@ -308,11 +404,14 @@ func (srv *server) ServeHTTPDebug(w http.ResponseWriter, r *http.Request) { io.WriteString(w, "

gocached

") io.WriteString(w, "

This is a shared Go build cache server, hit by GOCACHEPROG clients.

") io.WriteString(w, "

See /usage for usage stats.

") + io.WriteString(w, "

See /sessions for session data

") io.WriteString(w, "

See /metrics for Prometheus metrics.

") io.WriteString(w, "

See /debug/pprof/ for pprof

") io.WriteString(w, "

See /debug/pprof/goroutine?debug=2 - full goroutines

") case r.URL.Path == "/usage": srv.serveUsage(w, r) + case r.URL.Path == "/sessions": + srv.serveSessions(w, r) case r.URL.Path == "/metrics": srv.metricsHandler.ServeHTTP(w, r) case strings.HasPrefix(r.URL.Path, "/debug/pprof/profile"): @@ -334,8 +433,51 @@ func (srv *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { if srv.verbose { srv.logf("ServeHTTP: %s %s", r.Method, r.RequestURI) } + + var sessionData *sessionData // remains nil for unauthenticated requests. + reqStats := &stats{} + defer func() { + // Call inside func to capture maybe-updated sessionData pointer. + srv.processRequestStats(reqStats, sessionData) + }() + + // Handle session auth first if enabled. + if srv.jwtValidator != nil { + // If JWT auth enabled, this is the only unauthenticated (non-debug) endpoint. + if r.Method == "POST" && r.URL.Path == "/auth/exchange-token" { + srv.handleTokenExchange(w, r) + return + } + + // Check for session data and error if none. + token := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ") + if !strings.HasPrefix(token, tokenPrefix) { + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + + var ok bool + sessionData, ok = srv.getSessionData(token) + if !ok || srv.now().After(sessionData.expiry) { + if srv.verbose { + reason := fmt.Sprintf("exists: %v", ok) + if sessionData != nil { + reason += fmt.Sprintf(", expiry: %v", sessionData.expiry) + } + srv.logf("unauthorized; %s", reason) + } + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + } + if r.Method == "PUT" { - srv.handlePut(w, r) + if sessionData != nil && !sessionData.globalNSWrite { + // TODO(tomhjp): support per-namespace writes. + http.Error(w, "forbidden", http.StatusForbidden) + return + } + srv.handlePut(w, r, reqStats) return } if r.Method != "GET" && r.Method != "HEAD" { @@ -343,12 +485,64 @@ func (srv *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { return } if strings.HasPrefix(r.URL.Path, "/action/") { - srv.handleGetAction(w, r) + srv.handleGetAction(w, r, reqStats) + return + } + if sessionData != nil && r.URL.Path == "/session/stats" { + srv.handleSessionStats(w, sessionData) return } http.Error(w, "not found", http.StatusNotFound) } +func (srv *server) getSessionData(token string) (*sessionData, bool) { + srv.sessionsMu.Lock() + defer srv.sessionsMu.Unlock() + sessionData, ok := srv.sessions[token] + return sessionData, ok +} + +func (srv *server) addSessionData(token string, sessionData *sessionData) { + srv.sessionsMu.Lock() + defer srv.sessionsMu.Unlock() + srv.sessions[token] = sessionData + srv.Sessions.Add(1) +} + +func (srv *server) processRequestStats(req *stats, sessionData *sessionData) { + srv.Gets.Add(req.Gets) + srv.GetBytes.Add(req.GetBytes) + srv.GetHits.Add(req.GetHits) + srv.GetAccessBumps.Add(req.GetAccessBumps) + srv.GetHitsInline.Add(req.GetHitsInline) + srv.GetErrs.Add(req.GetErrs) + srv.Puts.Add(req.Puts) + srv.PutErrs.Add(req.PutErrs) + srv.PutsDup.Add(req.PutsDup) + srv.PutsBytes.Add(req.PutsBytes) + srv.PutsInline.Add(req.PutsInline) + + if sessionData != nil { + sessionData.mu.Lock() + defer sessionData.mu.Unlock() + + sessionData.stats.LastUsed = srv.now().UTC() + sessionData.stats.Gets += req.Gets + sessionData.stats.GetBytes += req.GetBytes + sessionData.stats.GetHits += req.GetHits + sessionData.stats.GetAccessBumps += req.GetAccessBumps + sessionData.stats.GetHitsInline += req.GetHitsInline + sessionData.stats.GetErrs += req.GetErrs + sessionData.stats.GetNanos += req.GetNanos + sessionData.stats.Puts += req.Puts + sessionData.stats.PutErrs += req.PutErrs + sessionData.stats.PutsDup += req.PutsDup + sessionData.stats.PutsBytes += req.PutsBytes + sessionData.stats.PutsInline += req.PutsInline + sessionData.stats.PutsNanos += req.PutsNanos + } +} + func getHexSuffix(r *http.Request, prefix string) (hexSuffix string, ok bool) { hexSuffix, _ = strings.CutPrefix(r.RequestURI, prefix) if !validHex(hexSuffix) { @@ -375,16 +569,20 @@ func validHex(x string) bool { // we do a DB write to update it. const relAtimeSeconds = 60 * 60 * 24 // 1 day -func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request) { +func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request, stats *stats) { srv.ActiveGets.Add(1) defer srv.ActiveGets.Add(-1) - srv.Gets.Add(1) + start := srv.now() + defer func() { + stats.GetNanos += srv.now().Sub(start).Nanoseconds() + }() + stats.Gets++ ctx := r.Context() httpErr := func(msg string, code int) { http.Error(w, msg, code) - srv.GetErrs.Add(1) + stats.GetErrs++ } actionID, ok := getHexSuffix(r, "/action/") @@ -431,10 +629,10 @@ func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request) { httpErr("internal server error", http.StatusInternalServerError) return } - srv.GetAccessBumps.Add(1) + stats.GetAccessBumps++ } - srv.GetHits.Add(1) + stats.GetHits++ outputID := cmp.Or(altObjectID, sha256hex) @@ -448,8 +646,8 @@ func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request) { if smallData.Valid { // For small outputs stored inline in the database, we can return them directly. - srv.GetHitsInline.Add(1) - srv.GetBytes.Add(size) + stats.GetHitsInline++ + stats.GetBytes += size io.WriteString(w, smallData.String) return } @@ -472,7 +670,7 @@ func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request) { http.Error(w, "not found", http.StatusNotFound) return } - srv.GetBytes.Add(size) + stats.GetBytes += size defer rc.Close() io.Copy(w, rc) } @@ -499,10 +697,15 @@ func (srv *server) getObjectFromDiskOrPeer(_ context.Context, sha256hex string) return f, nil } -func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { +func (s *server) handlePut(w http.ResponseWriter, r *http.Request, stats *stats) { s.ActivePuts.Add(1) defer s.ActivePuts.Add(-1) + start := s.now() + defer func() { + stats.PutsNanos += s.now().Sub(start).Nanoseconds() + }() + if r.Method != "PUT" { http.Error(w, "bad method", http.StatusMethodNotAllowed) return @@ -559,7 +762,7 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { `, sha256hex, blobSize, smallData).Scan(&blobID) if err != nil { s.logf("Blobs insert error: %v", err) - s.PutErrs.Add(1) + stats.PutErrs++ http.Error(w, "Blobs insert error", http.StatusInternalServerError) return } @@ -582,7 +785,7 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { ) if err != nil { s.logf("Actions insert error: %v", err) - s.PutErrs.Add(1) + stats.PutErrs++ http.Error(w, "Actions insert error", http.StatusInternalServerError) return } @@ -590,24 +793,114 @@ func (s *server) handlePut(w http.ResponseWriter, r *http.Request) { affected, err := res.RowsAffected() if err != nil { s.logf("Actions rows affected error: %v", err) - s.PutErrs.Add(1) + stats.PutErrs++ http.Error(w, "Actions rows affected error", http.StatusInternalServerError) return } if affected == 0 { - s.PutsDup.Add(1) + stats.PutsDup++ } - s.Puts.Add(1) - s.PutsBytes.Add(r.ContentLength) + stats.Puts++ + stats.PutsBytes += r.ContentLength if smallData != nil { - s.PutsInline.Add(1) + stats.PutsInline++ } w.WriteHeader(http.StatusNoContent) } +// handleTokenExchange handles POST /auth/exchange-token requests to exchange +// a JWT for an access token. Each access token represents a session that will +// last for one hour and have cache stats associated with it. +func (srv *server) handleTokenExchange(w http.ResponseWriter, r *http.Request) { + var req struct { + JWT string `json:"jwt"` + } + // JWTs are often sent in HTTP headers, so 4KiB should ~always be enough. + if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, 4<<10)).Decode(&req); err != nil { + srv.AuthErrs.Add(1) + http.Error(w, "bad request: "+err.Error(), http.StatusBadRequest) + return + } + + jwtClaims, err := srv.jwtValidator.Validate(r.Context(), req.JWT) + if err != nil { + srv.AuthErrs.Add(1) + if srv.verbose { + srv.logf("token exchange: JWT validation error: %v", err) + } + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + + globalNSWrite, err := srv.evaluateClaims(jwtClaims) + if err != nil { + srv.AuthErrs.Add(1) + if srv.verbose { + srv.logf("token exchange: %v", err) + } + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + + const ttl = time.Hour + // 52 base32 characters, 256 bits of entropy. + accessToken := tokenPrefix + strings.ToLower(rand.Text()+rand.Text()) + srv.addSessionData(accessToken, &sessionData{ + expiry: srv.now().UTC().Add(ttl), + globalNSWrite: globalNSWrite, + claims: jwtClaims, + }) + + resp := map[string]any{ + "access_token": accessToken, + "token_type": "Bearer", + "expires_in": ttl.Seconds(), + } + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(resp); err != nil { + srv.AuthErrs.Add(1) + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + + srv.Auths.Add(1) +} + +func (srv *server) evaluateClaims(claims map[string]any) (globalNSWrite bool, _ error) { + if missing := findMissingClaims(srv.jwtClaims, claims); len(missing) > 0 { + return false, fmt.Errorf("got claims %v; missing required claims: %v", claims, missing) + } + + if missing := findMissingClaims(srv.globalJWTClaims, claims); len(missing) == 0 { + return true, nil + } else if srv.verbose { + srv.logf("token exchange: missing global namespace write claims: %v", missing) + } + + return false, nil +} + +func findMissingClaims(wantClaims map[string]string, gotClaims map[string]any) map[string]string { + missing := make(map[string]string) + for k, want := range wantClaims { + if got, ok := gotClaims[k]; !ok || got != want { + missing[k] = want + } + } + return missing +} + +func (srv *server) handleSessionStats(w http.ResponseWriter, sessionData *sessionData) { + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(sessionData.stats); err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } +} + func (s *server) sha256Filepath(hash [sha256.Size]byte) string { hex := fmt.Sprintf("%x", hash) return filepath.Join(s.dir, hex[:2], hex) @@ -958,6 +1251,29 @@ func (srv *server) runCleanLoop() { } } +func (srv *server) runCleanSessionsLoop() { + for { + select { + case <-srv.shutdownCtx.Done(): + return + case <-time.After(time.Hour): + } + + srv.sessionsMu.Lock() + count := len(srv.sessions) + var deleted int + for token, metadata := range srv.sessions { + if time.Now().After(metadata.expiry) { + delete(srv.sessions, token) + deleted++ + srv.Sessions.Add(-1) + } + } + srv.sessionsMu.Unlock() + srv.logf("cleaned up %d/%d access tokens", deleted, count) + } +} + func durFmt(d time.Duration) string { days := int(d.Hours() / 24) if days > 0 { @@ -1017,6 +1333,52 @@ func (srv *server) serveUsage(w http.ResponseWriter, r *http.Request) { fmt.Fprintf(w, "\n") } +func (srv *server) serveSessions(w http.ResponseWriter, r *http.Request) { + if r.Method != "GET" { + http.Error(w, "bad method", http.StatusMethodNotAllowed) + return + } + + srv.sessionsMu.RLock() + // Make a copy of all session data (excluding the mutex). + sessions := make([]*sessionData, 0, len(srv.sessions)) + for _, v := range srv.sessions { + v.mu.Lock() + sessions = append(sessions, &sessionData{ + expiry: v.expiry, + globalNSWrite: v.globalNSWrite, + claims: v.claims, + stats: v.stats, + }) + v.mu.Unlock() + } + srv.sessionsMu.RUnlock() + + w.Header().Set("Content-Type", "text/html; charset=utf-8") + fmt.Fprintf(w, "

gocached sessions

\n") + fmt.Fprintf(w, "

JWT issuer: %s

\n", *jwtIssuer) + fmt.Fprintf(w, "

JWT claims required: %v

\n", srv.jwtClaims) + fmt.Fprintf(w, "

JWT global write claims required: %v

\n", srv.globalJWTClaims) + fmt.Fprintf(w, "

Number of sessions: %d

\n", len(sessions)) + + fmt.Fprintf(w, "\n") + fmt.Fprintf(w, "\n") + slices.SortFunc(sessions, func(a, b *sessionData) int { + return a.stats.LastUsed.Compare(b.stats.LastUsed) + }) + for _, d := range slices.Backward(sessions) { + lastUsed := "never" + if !d.stats.LastUsed.IsZero() { + lastUsed = durFmt(time.Since(d.stats.LastUsed)) + " ago" + } + statsJSON, _ := json.MarshalIndent(d.stats, "", " ") + claimsJSON, _ := json.MarshalIndent(d.claims, "", " ") + fmt.Fprintf(w, "\n", + lastUsed, d.expiry.Format(time.RFC3339), d.globalNSWrite, statsJSON, claimsJSON) + } + fmt.Fprintf(w, "
Last usedExpiry timeGlobal writeStatsClaims
%s%s%v
%s
%s
\n") +} + // expvarCounterMetric is a Prometheus counter metric backed by an expvar.Int. type expvarCounterMetric struct { desc *prometheus.Desc diff --git a/cmd/gocached/gocached_test.go b/cmd/gocached/gocached_test.go index be1da1d..f28a0ba 100644 --- a/cmd/gocached/gocached_test.go +++ b/cmd/gocached/gocached_test.go @@ -1,9 +1,16 @@ package main import ( + "bytes" "context" + "crypto" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "encoding/json" "expvar" "fmt" + "io" "math" "net" "net/http" @@ -17,7 +24,10 @@ import ( "testing" "time" + "github.com/go-jose/go-jose/v4" + "github.com/golang-jwt/jwt/v5" "github.com/google/go-cmp/cmp" + ijwt "github.com/tailscale/tb/cmd/gocached/internal/jwt" "github.com/tailscale/tb/gocache/cachers" ) @@ -30,6 +40,10 @@ type tester struct { srv *server hs *httptest.Server + issuer string + audience string + createJWT func(claims jwt.MapClaims, signingKey *ecdsa.PrivateKey) string + timeMu sync.Mutex curTime time.Time } @@ -191,6 +205,65 @@ func newServerTester(t testing.TB) *tester { return st } +// startOIDCServer starts a mock OIDC server and configures gocached to use it +// for JWT validation. The provided publicKey is what JWT signatures will be +// validated against. Use st.createJWT to create signed JWTs. +func (st *tester) startOIDCServer(publicKey crypto.PublicKey) { + st.t.Helper() + mux := http.NewServeMux() + srv := httptest.NewServer(mux) + st.t.Cleanup(srv.Close) + + st.issuer = fmt.Sprintf("http://%s", srv.Listener.Addr().String()) + st.audience = gocachedAudience + + mux.HandleFunc("/.well-known/openid-configuration", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "issuer": st.issuer, + "jwks_uri": fmt.Sprintf("%s/jwks", st.issuer), + }) + }) + mux.HandleFunc("/jwks", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "keys": []jose.JSONWebKey{ + { + Key: publicKey, + KeyID: "test-key", + Algorithm: "ES256", + Use: "sig", + }, + }, + }) + }) + + st.srv.jwtValidator = ijwt.NewJWTValidator(st.issuer, st.audience) + st.srv.jwtValidator.Logf = st.Logf + if err := st.srv.jwtValidator.RunUpdateJWKSLoop(st.srv.shutdownCtx); err != nil { + st.t.Fatalf("failed to start JWKS loop for JWT validator: %v", err) + } + + st.createJWT = func(claims jwt.MapClaims, signingKey *ecdsa.PrivateKey) string { + st.t.Helper() + unsignedTk := &jwt.Token{ + Header: map[string]any{ + "typ": "JWT", + "alg": jwt.SigningMethodES256.Alg(), + "kid": "test-key", + }, + Claims: claims, + Method: jwt.SigningMethodES256, + } + tk, err := unsignedTk.SignedString(signingKey) + if err != nil { + st.t.Fatalf("error signing token: %v", err) + } + + return tk + } +} + func TestServer(t *testing.T) { st := newServerTester(t) @@ -427,3 +500,245 @@ func TestClientConnReuse(t *testing.T) { t.Errorf("numDials = %d; want 1", got) } } + +func TestExchangeToken(t *testing.T) { + // Generate private keys outside of the loop for speed. + privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("error generating OIDC server private key: %v", err) + } + otherPrivateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("error generating OIDC server private key: %v", err) + } + + for name, tc := range map[string]struct { + mutateClaims func(jwt.MapClaims) + signingKey *ecdsa.PrivateKey + wantStatusCode int + wantWrite bool + }{ + // Base case: no mutation. + "valid_read": { + wantStatusCode: http.StatusOK, + wantWrite: false, + }, + // Additional claim needed for write scope. + "valid_write": { + mutateClaims: func(cl jwt.MapClaims) { + cl["ref"] = "refs/heads/main" + }, + wantStatusCode: http.StatusOK, + wantWrite: true, + }, + // Every other test makes one mutation from the base case that should cause failure. + "missing_sub": { + mutateClaims: func(cl jwt.MapClaims) { + delete(cl, "sub") + }, + wantStatusCode: http.StatusUnauthorized, + }, + "invalid_sub": { + mutateClaims: func(cl jwt.MapClaims) { + cl["sub"] = "user456" + }, + wantStatusCode: http.StatusUnauthorized, + }, + "invalid_iss": { + mutateClaims: func(cl jwt.MapClaims) { + cl["iss"] = "invalid_issuer" + }, + wantStatusCode: http.StatusUnauthorized, + }, + "invalid_aud": { + mutateClaims: func(cl jwt.MapClaims) { + cl["aud"] = "invalid_audience" + }, + wantStatusCode: http.StatusUnauthorized, + }, + "not_yet_valid": { + mutateClaims: func(cl jwt.MapClaims) { + cl["nbf"] = jwt.NewNumericDate(time.Now().Add(10 * time.Minute)) + }, + wantStatusCode: http.StatusUnauthorized, + }, + "expired": { + mutateClaims: func(cl jwt.MapClaims) { + cl["exp"] = jwt.NewNumericDate(time.Now().Add(-time.Minute)) + }, + wantStatusCode: http.StatusUnauthorized, + }, + "invalid_signature": { + signingKey: otherPrivateKey, + wantStatusCode: http.StatusUnauthorized, + }, + } { + t.Run(name, func(t *testing.T) { + st := newServerTester(t) + st.startOIDCServer(privateKey.Public()) + st.srv.jwtClaims = map[string]string{ + "sub": "user123", + } + st.srv.globalJWTClaims = map[string]string{ + "sub": "user123", + "ref": "refs/heads/main", + } + + // Generate JWT. + tokenClaims := jwt.MapClaims{ + "sub": "user123", + "iss": st.issuer, + "aud": st.audience, + "nbf": jwt.NewNumericDate(time.Now().Add(-time.Minute)), + "exp": jwt.NewNumericDate(time.Now().Add(time.Hour)), + } + if tc.mutateClaims != nil { + tc.mutateClaims(tokenClaims) + } + signingKey := privateKey + if tc.signingKey != nil { + signingKey = tc.signingKey + } + body, err := json.Marshal(map[string]any{ + "jwt": st.createJWT(tokenClaims, signingKey), + }) + if err != nil { + t.Fatalf("error marshaling request body: %v", err) + } + + // Exchange JWT for access token. + req, err := http.NewRequest("POST", st.hs.URL+"/auth/exchange-token", bytes.NewReader(body)) + if err != nil { + t.Fatalf("error creating request: %v", err) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("error making request: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != tc.wantStatusCode { + t.Fatalf("unexpected status code: want %d, got %d", tc.wantStatusCode, resp.StatusCode) + } + body, err = io.ReadAll(resp.Body) + if err != nil { + t.Fatalf("error reading response body: %v", err) + } + + if tc.wantStatusCode != http.StatusOK { + if string(body) != "unauthorized\n" { + t.Fatalf("unexpected error body: %s", string(body)) + } + + // No access token to do further checks with; test finished. + return + } + + // Check returned access token. + var d struct { + AccessToken string `json:"access_token"` + } + if err := json.Unmarshal(body, &d); err != nil { + t.Fatalf("error decoding response body: %v", err) + } + if d.AccessToken == "" { + t.Fatalf("expected access_token in response, got %s", string(body)) + } + + cl := st.mkClient() + if _, _, err := cl.Get(t.Context(), "abc123"); err == nil { + t.Fatalf("Get without access token succeeded unexpectedly") + } + + cl.AccessToken = d.AccessToken + st.wantGetMiss(cl, "abc123") + + if tc.wantWrite { + st.wantPut(cl, "abc123", "def456", "data789") + st.wantGet(cl, "abc123", "def456", "data789") + } else { + if _, err := cl.Put(t.Context(), "abc123", "def456", 0, nil); err == nil { + t.Fatalf("Put without write scope succeeded unexpectedly") + } + } + + // Check session stats. + reqStats, err := http.NewRequest("GET", st.hs.URL+"/session/stats", nil) + if err != nil { + t.Fatalf("error creating stats request: %v", err) + } + reqStats.Header.Set("Authorization", "Bearer "+d.AccessToken) + respStats, err := http.DefaultClient.Do(reqStats) + if err != nil { + t.Fatalf("error making stats request: %v", err) + } + defer respStats.Body.Close() + if respStats.StatusCode != http.StatusOK { + t.Fatalf("unexpected stats status code: want %d, got %d", http.StatusOK, respStats.StatusCode) + } + bodyStats, err := io.ReadAll(respStats.Body) + if err != nil { + t.Fatalf("error reading stats response body: %v", err) + } + var stats stats + if err := json.Unmarshal(bodyStats, &stats); err != nil { + t.Fatalf("error decoding stats response body: %v", err) + } + t.Logf("stats: %v", stats) + if stats.Gets == 0 { + t.Errorf("expected non-zero gets in session stats") + } + if stats.Puts == 0 && tc.wantWrite { + t.Errorf("expected non-zero puts in session stats") + } + }) + } +} + +func TestJWTClaimFlag(t *testing.T) { + for name, tc := range map[string]struct { + input []string + want jwtClaimValue + }{ + "single_claim": { + input: []string{"role=builder"}, + want: jwtClaimValue{"role": "builder"}, + }, + "multiple_claims": { + input: []string{"env=prod", "team=devops"}, + want: jwtClaimValue{"env": "prod", "team": "devops"}, + }, + "duplicate_keys": { + input: []string{"key=value1", "key=value2"}, + want: jwtClaimValue{"key": "value2"}, + }, + "value_with_equals": { + input: []string{"data=a=b=c"}, + want: jwtClaimValue{"data": "a=b=c"}, + }, + } { + t.Run(name, func(t *testing.T) { + claim := make(jwtClaimValue) + for _, v := range tc.input { + if err := claim.Set(v); err != nil { + t.Fatalf("Set failed: %v", err) + } + } + if diff := cmp.Diff(claim, tc.want); diff != "" { + t.Errorf("jwtClaimValue mismatch (-got +want):\n%s", diff) + } + }) + } +} + +func TestInvalidJWTClaimFlags(t *testing.T) { + for _, s := range []string{ + "invalidclaim", + "=nokey", + "novalue=", + } { + claim := make(jwtClaimValue) + if err := claim.Set(s); err == nil { + t.Fatalf("Set with invalid claim %q did not return error", s) + } + } +} diff --git a/cmd/gocached/internal/jwt/jwt.go b/cmd/gocached/internal/jwt/jwt.go new file mode 100644 index 0000000..1344b73 --- /dev/null +++ b/cmd/gocached/internal/jwt/jwt.go @@ -0,0 +1,201 @@ +package jwt + +import ( + "context" + "encoding/json" + "fmt" + "io" + "log" + "net/http" + "net/url" + "path" + "slices" + "sync/atomic" + "time" + + "github.com/go-jose/go-jose/v4" + "github.com/golang-jwt/jwt/v5" +) + +const oidcConfigWellKnownPath string = "/.well-known/openid-configuration" + +var ( + // Required + recommended algorithms from https://datatracker.ietf.org/doc/html/rfc7518#section-3.1 + supportedAlgorithms = []string{"HS256", "RS256", "ES256"} +) + +// NewJWTValidator constructs a [Validator] for validating JWTs. Must call +// [RunUpdateJWKSLoop] before validating any JWTs. Every JWT must exactly match +// the provided issuer and audience values in its "iss" and "aud" claims +// respectively. The issuer must be a reachable HTTP server that serves the JWT +// public signing keys via the path defined by [oidcConfigWellKnownPath], and +// the audience should be a value specific to the trust boundary that gocached +// resides within. +func NewJWTValidator(issuer, audience string) *Validator { + return &Validator{ + Logf: log.Printf, + issuer: issuer, + parser: jwt.NewParser( + jwt.WithValidMethods(supportedAlgorithms), + jwt.WithIssuer(issuer), + jwt.WithAudience(audience), + jwt.WithLeeway(10*time.Second), + jwt.WithIssuedAt(), + ), + } +} + +// RunUpdateJWKSLoop fetches the JWKS synchronously once to surface any config +// errors early, and then starts a background goroutine that periodically fetches +// the JWKS from the issuer to keep the signing keys up to date. Must be called +// before validating any JWTs. +func (v *Validator) RunUpdateJWKSLoop(ctx context.Context) error { + // Initial fetch to error early on misconfiguration. + if err := v.updateJWKS(ctx); err != nil { + return fmt.Errorf("failed to initialize JWT validator: %w", err) + } + + go v.runUpdateJWKSLoop(ctx) + + return nil +} + +// Validator provides methods for validating JWTs. Use [NewJWTValidator] to +// construct a working Validator. +type Validator struct { + Logf func(format string, args ...any) + issuer string + parser *jwt.Parser + + signingKeys atomic.Value // []jose.JSONWebKey + + // TODO(tomhjp): metrics +} + +// Validate returns an error if the provided JWT fails validation for an invalid +// signature or standard claim (iss, aud, iat, nbf, exp). It returns the token's +// verified claims if validation succeeds. The caller should then make policy +// decisions based on other claims such as "sub" or other custom claims. +func (v *Validator) Validate(ctx context.Context, jwtString string) (map[string]any, error) { + tk, err := v.parser.Parse(jwtString, v.keyFunc) + if err != nil { + return nil, fmt.Errorf("failed to parse token: %w", err) + } + + if !tk.Valid { + return nil, fmt.Errorf("invalid token") + } + + gotClaims, ok := tk.Claims.(jwt.MapClaims) + if !ok { + return nil, fmt.Errorf("unexpected claims type: %T", tk.Claims) + } + + return gotClaims, nil +} + +// keyFunc is how github.com/golang-jwt/jwt gets the public key it needs to +// verify a JWT signature. +func (v *Validator) keyFunc(t *jwt.Token) (any, error) { + var kid string + if v, ok := t.Header["kid"]; ok { + kid, _ = v.(string) + } + if kid == "" { + return nil, fmt.Errorf("no kid found in token header") + } + + signingKeys := v.signingKeys.Load().([]jose.JSONWebKey) + for _, k := range signingKeys { + if k.KeyID == kid { + return k.Key, nil + } + } + + return nil, fmt.Errorf("unknown key ID: %s", kid) +} + +func (v *Validator) runUpdateJWKSLoop(ctx context.Context) { + for { + select { + case <-ctx.Done(): + return + case <-time.After(12 * time.Hour): + } + + if err := v.updateJWKS(ctx); err != nil { + // Non-fatal; in practice, the most recent keys will normally + // still be valid for a long time, but JWT validation will start + // erroring more loudly than this if not. + v.Logf("jwt: failed to update JWKS: %v", err) + } + } +} + +func (v *Validator) updateJWKS(ctx context.Context) error { + v.Logf("jwt: fetching JWKS from issuer %q", v.issuer) + u, err := url.Parse(v.issuer) + if err != nil { + return fmt.Errorf("failed to parse issuer URL %q: %w", v.issuer, err) + } + + u.Path = path.Join(u.Path, oidcConfigWellKnownPath) + + ctx, cancel := context.WithTimeout(ctx, time.Minute) + defer cancel() + + // Discover the JWKS endpoint. + var config struct { + JSONWebKeySetURI string `json:"jwks_uri"` + } + if err := get(ctx, u.String(), &config); err != nil { + return fmt.Errorf("failed to fetch OIDC configuration from %q: %w", u, err) + } + if config.JSONWebKeySetURI == "" { + return fmt.Errorf("jwks_uri not found in OIDC configuration") + } + + // Fetch JWKS. + var keySet jose.JSONWebKeySet + if err := get(ctx, config.JSONWebKeySetURI, &keySet); err != nil { + return fmt.Errorf("failed to fetch JWKS from %q: %w", config.JSONWebKeySetURI, err) + } + + var signingKeys []jose.JSONWebKey + for _, k := range keySet.Keys { + if k.Use != "sig" { + continue + } + if !slices.Contains(supportedAlgorithms, k.Algorithm) { + continue + } + signingKeys = append(signingKeys, k) + } + + v.signingKeys.Store(signingKeys) + return nil +} + +func get(ctx context.Context, url string, out any) error { + req, err := http.NewRequestWithContext(ctx, "GET", url, nil) + if err != nil { + return fmt.Errorf("failed to create HTTP request: %w", err) + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + return fmt.Errorf("failed to perform HTTP request: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + return fmt.Errorf("received non-200 response: %d, body: %q", resp.StatusCode, body) + } + + if err := json.NewDecoder(resp.Body).Decode(out); err != nil { + return fmt.Errorf("failed to decode response body: %w", err) + } + + return nil +} diff --git a/go.mod b/go.mod index da652cd..779fdf3 100644 --- a/go.mod +++ b/go.mod @@ -1,9 +1,9 @@ module github.com/tailscale/tb -go 1.23.0 - -toolchain go1.23.4 +go 1.24.0 require ( + github.com/go-jose/go-jose/v4 v4.1.3 + github.com/golang-jwt/jwt/v5 v5.3.0 github.com/google/go-cmp v0.7.0 github.com/prometheus/client_golang v1.23.0 github.com/prometheus/client_model v0.6.2 diff --git a/go.sum b/go.sum index 717c8c5..fee6bb1 100644 --- a/go.sum +++ b/go.sum @@ -6,6 +6,10 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= +github.com/go-jose/go-jose/v4 v4.1.3 h1:CVLmWDhDVRa6Mi/IgCgaopNosCaHz7zrMeF9MlZRkrs= +github.com/go-jose/go-jose/v4 v4.1.3/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08= +github.com/golang-jwt/jwt/v5 v5.3.0 h1:pv4AsKCKKZuqlgs5sUmn4x8UlGa0kEVt/puTpKx9vvo= +github.com/golang-jwt/jwt/v5 v5.3.0/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= diff --git a/gocache/cachers/http.go b/gocache/cachers/http.go index f245cff..1b809e2 100644 --- a/gocache/cachers/http.go +++ b/gocache/cachers/http.go @@ -30,6 +30,10 @@ type HTTPClient struct { // Verbose optionally specifies whether to log verbose messages. Verbose bool + + // AccessToken optionally specifies a Bearer access token to include + // in requests to the server. + AccessToken string } func (c *HTTPClient) httpClient() *http.Client { @@ -59,6 +63,9 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa } req, _ := http.NewRequestWithContext(ctx, "GET", c.BaseURL+"/action/"+actionID, nil) + if c.AccessToken != "" { + req.Header.Set("Authorization", "Bearer "+c.AccessToken) + } // Set a header to indicate we want the object and metadata in one response. // Prior to 2025-08-09, the protocol was two separate requests. Rather than @@ -163,6 +170,9 @@ func (c *HTTPClient) Put(ctx context.Context, actionID, outputID string, size in } req, _ := http.NewRequestWithContext(ctx, "PUT", c.BaseURL+"/"+actionID+"/"+outputID, putBody) req.ContentLength = size + if c.AccessToken != "" { + req.Header.Set("Authorization", "Bearer "+c.AccessToken) + } res, err := c.httpClient().Do(req) pw.Close() if err != nil { From 7ab488e1318982e300c8e1289f617fc2116066d4 Mon Sep 17 00:00:00 2001 From: Tom Proctor Date: Thu, 13 Nov 2025 22:35:07 +0000 Subject: [PATCH 29/67] {cmd/,}gocached: make gocached library package (bradfitz/go-tool-cache#22) Pulls the bulk of gocached into a library package to more easily reuse it with custom config from other main packages. There are still likely some more changes to come to the JWT configuration API surface, but this makes sense as a first chunk of work to land. Signed-off-by: Tom Proctor Migrated-from: bradfitz/go-tool-cache@0124e698e0bde17d7c50f32fea625a1aab2a24c7 --- cmd/gocached/gocached.go | 1369 +--------------- cmd/gocached/gocached_test.go | 691 +------- gocache/gocached/gocached.go | 1453 +++++++++++++++++ gocache/gocached/gocached_test.go | 699 ++++++++ {cmd/gocached => gocache}/internal/jwt/jwt.go | 14 +- 5 files changed, 2185 insertions(+), 2041 deletions(-) create mode 100644 gocache/gocached/gocached.go create mode 100644 gocache/gocached/gocached_test.go rename {cmd/gocached => gocache}/internal/jwt/jwt.go (94%) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index 4f46405..d516622 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -1,88 +1,24 @@ // Copyright (c) Tailscale Inc & AUTHORS // SPDX-License-Identifier: BSD-3-Clause -// The gocached daemon is an HTTP server daemon that go-cacher can hit. It does -// cache tiering and evicts old large things from disk, and can fetch metadata -// and object contents from peer cache servers. -// -// It uses sqlite (the pure Go modernc.org/sqlite driver) to store metadata and -// indexes. -// -/* - -It speaks the same protocol as go-cacher-server, but requires -the "Want-Object: 1" header variant on the GET request. - - GET /action/ - Want-Object: 1 - - 200 OK - Content-Type: application/octet-stream - Content-Length: 1234 - Go-Output-Id: xxxxxxxxxxx - - - -And to insert an object: - - PUT // - Content-Length: 1234 - - - -*/ +// The gocached command is a small example binary showing how to use +// the gocached library package. package main import ( - "cmp" - "context" - "crypto/rand" - "crypto/sha256" - "database/sql" - "encoding/hex" - "encoding/json" - "errors" - "expvar" "flag" "fmt" - "io" "log" "maps" - "math" "net" "net/http" - "net/http/pprof" - "os" - "path/filepath" - "reflect" - "slices" "strings" - "sync" - "sync/atomic" "time" - "github.com/prometheus/client_golang/prometheus" - "github.com/prometheus/client_golang/prometheus/collectors" - "github.com/prometheus/client_golang/prometheus/promhttp" - dto "github.com/prometheus/client_model/go" - "github.com/tailscale/tb/cmd/gocached/internal/jwt" + "github.com/tailscale/tb/gocache/gocached" _ "modernc.org/sqlite" ) -const ( - // smallObjectSize is the maximum size of an object that we store inline in the - // database, rather than on disk. Empirically, about half of objects are 1KB or - // smaller. - smallObjectSize = 1 << 10 - - // tokenPrefix is the prefix for all gocached access tokens. - tokenPrefix = "gocached-token-" - - // gocachedAudience is the audience we require JWTs to have. Could be - // configurable in future, but for now just needs to be specific to gocached. - gocachedAudience = "gocached" -) - var ( dir = flag.String("cache-dir", "", "cache directory, if empty defaults to /gocached") verbose = flag.Bool("verbose", false, "be verbose") @@ -118,1305 +54,44 @@ func main() { flag.Var(&jwtClaims, "jwt-claim", "a claim in the form x=y that any JWT presented must have to start a session; may be specified more than once") flag.Var(&globalJWTClaims, "global-jwt-claim", "an additional claim in the form x=y that a JWT must have to allow writing to the cache's global namespace; may be specified more than once") flag.Parse() - if *dir == "" { - d, err := os.UserCacheDir() - if err != nil { - log.Fatal(err) - } - d = filepath.Join(d, "gocached") - log.Printf("Defaulting to cache dir %v ...", d) - *dir = d - } - if err := os.MkdirAll(*dir, 0750); err != nil { - log.Fatal(err) - } - - srv, err := newServer(*dir) - if err != nil { - log.Fatalf("newServer: %v", err) - } - srv.verbose = *verbose - srv.maxSize = int64(*maxSize) << 30 - srv.maxAge = time.Duration(*maxAge) * 24 * time.Hour - - log.Printf("gocached: scanning usage & cleaning as needed...") - us, err := srv.usageStats() - if err != nil { - log.Fatalf("getting usage stats: %v", err) - } - log.Printf("gocached: current usage: %v of limit %v", us.All(), bytesFmt(srv.maxSize)) - if res, err := srv.cleanOldObjects(us); err != nil { - log.Fatalf("clean old objects: %v", err) - } else if res.Count > 0 { - log.Printf("gocached: cleaned %v", res) - } - - if *debugListen != "" { - debugLn, err := net.Listen("tcp", *debugListen) - if err != nil { - log.Fatalf("debug listen: %v", err) - } - go func() { - log.Fatal(http.Serve(debugLn, http.HandlerFunc(srv.ServeHTTPDebug))) - }() + opts := []gocached.ServerOption{ + gocached.WithDir(*dir), + gocached.WithVerbose(*verbose), + gocached.WithMaxSize(int64(*maxSize) << 30), + gocached.WithMaxAge(time.Duration(*maxAge) * 24 * time.Hour), } if *jwtIssuer != "" { if len(jwtClaims) == 0 { log.Fatal("must specify --jwt-claim at least once when --jwt-issuer is set") } - srv.jwtValidator = jwt.NewJWTValidator(*jwtIssuer, gocachedAudience) - if err := srv.jwtValidator.RunUpdateJWKSLoop(srv.shutdownCtx); err != nil { - log.Fatalf("failed to fetch JWKS for JWT validator: %v", err) - } - srv.jwtClaims = jwtClaims globalClaims := map[string]string{} maps.Copy(globalClaims, jwtClaims) maps.Copy(globalClaims, globalJWTClaims) - srv.globalJWTClaims = globalClaims - - log.Printf("gocached: using JWT issuer %q with claims %v, global claims %v", *jwtIssuer, srv.jwtClaims, srv.globalJWTClaims) - - go srv.runCleanSessionsLoop() - } - - go srv.runCleanLoop() - - log.Printf("gocached: listening on %s ...", *listen) - log.Fatal(http.ListenAndServe(*listen, srv)) -} - -const schemaVersion = 3 - -const schema = ` -PRAGMA journal_mode=WAL; -CREATE TABLE IF NOT EXISTS Actions ( - NamespaceID INTEGER NOT NULL, -- 0 for global trusted namespace - ActionID TEXT NOT NULL, - BlobID INTEGER NOT NULL, - AltOutputID TEXT NOT NULL DEFAULT '', -- if non-empty, the alternate object ID to use for this action; NULL means the blob's sha256 - CreateTime INTEGER NOT NULL, -- unix sec when inserted (locally or on a peer) - AccessTime INTEGER NOT NULL, -- unix sec of last access - - PRIMARY KEY (NamespaceID, ActionID), - - CHECK (ActionID = lower(ActionID)), - CHECK (ActionID GLOB '[0-9a-f]*'), - CHECK (CreateTime >= 0), - CHECK (AccessTime >= 0) -) STRICT; - -CREATE INDEX IF NOT EXISTS idx_actions_access ON Actions(AccessTime); -CREATE INDEX IF NOT EXISTS idx_actions_blobid ON Actions(BlobID); - -CREATE TABLE IF NOT EXISTS Blobs ( - BlobID INTEGER PRIMARY KEY AUTOINCREMENT, - SHA256 TEXT NOT NULL, - BlobSize INTEGER NOT NULL, -- size in bytes, either inline or on disk - SmallData BLOB, -- NULL if stored on disk - - CHECK (SmalLData IS NULL OR length(SmallData) = BlobSize) -) STRICT; - -CREATE UNIQUE INDEX IF NOT EXISTS idx_blobs_sha256 ON Blobs(SHA256); - -CREATE TABLE IF NOT EXISTS Namespaces ( - NamespaceID INTEGER PRIMARY KEY AUTOINCREMENT, - Namespace TEXT NOT NULL UNIQUE CHECK (Namespace = lower(Namespace)) -) STRICT; -` - -func openDB(dbDir string) (*sql.DB, error) { - dbPath := filepath.Join(dbDir, fmt.Sprintf("gocached-v%d.db", schemaVersion)) - db, err := sql.Open("sqlite", "file:"+dbPath+"?_pragma=busy_timeout(5000)") - if err != nil { - return nil, err - } - db.SetMaxOpenConns(4) - db.SetMaxIdleConns(4) - db.SetConnMaxLifetime(0) // no limit - if _, err := db.Exec(schema); err != nil { - return nil, err - } - return db, nil -} - -func newServer(dir string) (*server, error) { - db, err := openDB(dir) - if err != nil { - return nil, fmt.Errorf("openDB: %w", err) - } - - reg := prometheus.NewRegistry() - reg.MustRegister( - collectors.NewGoCollector(), - collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}), - collectors.NewBuildInfoCollector(), - ) - - srv := &server{ - db: db, - dir: dir, - logf: log.Printf, - sessions: make(map[string]*sessionData), - } - srv.shutdownCtx, srv.shutdownCancel = context.WithCancel(context.Background()) - srv.registerMetrics(reg) - - srv.metricsHandler = promhttp.HandlerFor(reg, promhttp.HandlerOpts{ - ErrorLog: log.Default(), - }) - - return srv, nil -} - -func (s *server) registerMetrics(reg *prometheus.Registry) { - rv := reflect.ValueOf(s).Elem() - t := reflect.TypeOf(s).Elem() - for i := 0; i < t.NumField(); i++ { - sf := t.Field(i) - if sf.Type == reflect.TypeFor[expvar.Int]() { - expvarInt := rv.Field(i).Addr().Interface().(*expvar.Int) - typ := sf.Tag.Get("type") - name := sf.Tag.Get("name") - if typ == "" { - panic("missing type tag for " + sf.Name) - } - if name == "" { - panic("missing name tag for " + sf.Name) - } - help := sf.Tag.Get("help") - metricName := "gocached_" + name - - if tag := sf.Tag.Get("type"); tag != "" { - if tag == "gauge" { - reg.MustRegister(singleMetricCollector{&expvarGaugeMetric{ - desc: prometheus.NewDesc(metricName, help, nil, nil), - v: expvarInt, - }}) - } else if tag == "counter" { - reg.MustRegister(singleMetricCollector{&expvarCounterMetric{ - desc: prometheus.NewDesc(metricName, help, nil, nil), - v: expvarInt, - }}) - } - } - } - } -} - -type server struct { - db *sql.DB - dir string // for SQLite DB + large blobs - verbose bool - logf func(format string, args ...any) - clock func() time.Time // if non-nil, alternate time.Now for testing - metricsHandler http.Handler - maxSize int64 // maximum size of the cache in bytes; 0 means no limit - maxAge time.Duration // maximum age of objects; 0 means no limit - shutdownCtx context.Context - shutdownCancel context.CancelFunc - - jwtValidator *jwt.Validator // nil unless -jwt-issuer flag is set - jwtClaims map[string]string // claims required for any JWT to start a session - globalJWTClaims map[string]string // additional claims required to write to global namespace - - sessionsMu sync.RWMutex // guards sessions - sessions map[string]*sessionData // maps access token -> session data. - - // sqliteWriteMu serializes access to SQLite. In theory the SQLite driver - // should serialize access with our 5000ms busy timeout, but empirically we - // sometimes seen DB busy errors. Just serialize it explicitly out of - // laziness for now. - sqliteWriteMu sync.Mutex - - lastUsage atomic.Pointer[usageStats] - - // Metrics - ActiveGets expvar.Int `type:"gauge" name:"active_gets" help:"currently pending get requests; should usually be zero"` - ActivePuts expvar.Int `type:"gauge" name:"active_puts" help:"currently pending put requests; should usually be zero"` - Gets expvar.Int `type:"counter" name:"gets" help:"total number of gocache get requests"` // gets = getHits + getErrs + implicit misses - GetBytes expvar.Int `type:"counter" name:"get_bytes" help:"total bytes fetched from gocache gets that were cache hits"` - GetHits expvar.Int `type:"counter" name:"get_hits" help:"total number of successful gocache get requests"` - GetAccessBumps expvar.Int `type:"counter" name:"get_access_bumps" help:"number of times a get request updated the access time of object"` - GetHitsInline expvar.Int `type:"counter" name:"get_hits_inline" help:"cache hits served from inline database storage (small objects)"` - GetErrs expvar.Int `type:"counter" name:"get_errs" help:"number of gocache get request errors"` - Puts expvar.Int `type:"counter" name:"puts" help:"total number of gocache put requests"` - PutErrs expvar.Int `type:"counter" name:"put_errs" help:"number of gocache put request errors"` - PutsDup expvar.Int `type:"counter" name:"puts_dup" help:"total number of gocache put requests that are duplicates of a mapping we already had"` - PutsBytes expvar.Int `type:"counter" name:"put_bytes" help:"total bytes added from gocache puts"` - PutsInline expvar.Int `type:"counter" name:"put_inline" help:"subset of gocached_puts that were stored inline (small objects)"` - BlobCount expvar.Int `type:"gauge" name:"blob_count" help:"number of blobs currently stored in the cache"` - BlobBytes expvar.Int `type:"gauge" name:"blob_bytes" help:"sum of blob sizes currently stored in the cache"` - EvictedBlobs expvar.Int `type:"counter" name:"evicted_blobs" help:"number of blobs evicted from the cache"` - EvictedBytes expvar.Int `type:"counter" name:"evicted_bytes" help:"number of bytes evicted from the cache"` - Sessions expvar.Int `type:"gauge" name:"sessions" help:"number of active authenticated sessions"` - Auths expvar.Int `type:"counter" name:"auth_attempts" help:"number of successful token exchanges"` - AuthErrs expvar.Int `type:"counter" name:"auth_errs" help:"number of failed token exchanges"` -} - -// sessionData corresponds to a specific access token, and is only used if JWT -// auth is enabled. -type sessionData struct { - expiry time.Time // Session valid until. - globalNSWrite bool // Whether this session can write to the cache's global namespace. - claims map[string]any // Claims from the JWT used to create this session, stored for debug. - - mu sync.Mutex // Guards stats. - stats stats -} - -// stats holds per-request or per-session stats which get rolled up into server -// stats. See [server] struct for detailed definitions. -type stats struct { - LastUsed time.Time // Only applies to session stats. Last time the access token for this session was used. - Gets int64 - GetBytes int64 - GetHits int64 - GetAccessBumps int64 - GetHitsInline int64 - GetNanos int64 - GetErrs int64 - Puts int64 - PutErrs int64 - PutsDup int64 - PutsBytes int64 - PutsInline int64 - PutsNanos int64 -} - -func (srv *server) now() time.Time { - if srv.clock != nil { - return srv.clock() - } - return time.Now() -} - -func (srv *server) ServeHTTPDebug(w http.ResponseWriter, r *http.Request) { - if srv.verbose { - srv.logf("ServeHTTPDebug: %s %s", r.Method, r.RequestURI) - } - switch { - case r.URL.Path == "/": - w.Header().Set("Content-Type", "text/html; charset=utf-8") - io.WriteString(w, "

gocached

") - io.WriteString(w, "

This is a shared Go build cache server, hit by GOCACHEPROG clients.

") - io.WriteString(w, "

See /usage for usage stats.

") - io.WriteString(w, "

See /sessions for session data

") - io.WriteString(w, "

See /metrics for Prometheus metrics.

") - io.WriteString(w, "

See /debug/pprof/ for pprof

") - io.WriteString(w, "

See /debug/pprof/goroutine?debug=2 - full goroutines

") - case r.URL.Path == "/usage": - srv.serveUsage(w, r) - case r.URL.Path == "/sessions": - srv.serveSessions(w, r) - case r.URL.Path == "/metrics": - srv.metricsHandler.ServeHTTP(w, r) - case strings.HasPrefix(r.URL.Path, "/debug/pprof/profile"): - pprof.Profile(w, r) - case strings.HasPrefix(r.URL.Path, "/debug/pprof/cmdline"): - pprof.Cmdline(w, r) - case strings.HasPrefix(r.URL.Path, "/debug/pprof/symbol"): - pprof.Symbol(w, r) - case strings.HasPrefix(r.URL.Path, "/debug/pprof/trace"): - pprof.Trace(w, r) - case strings.HasPrefix(r.URL.Path, "/debug/pprof/"): - pprof.Index(w, r) - default: - http.Error(w, "not found", http.StatusNotFound) - } -} - -func (srv *server) ServeHTTP(w http.ResponseWriter, r *http.Request) { - if srv.verbose { - srv.logf("ServeHTTP: %s %s", r.Method, r.RequestURI) - } - - var sessionData *sessionData // remains nil for unauthenticated requests. - reqStats := &stats{} - defer func() { - // Call inside func to capture maybe-updated sessionData pointer. - srv.processRequestStats(reqStats, sessionData) - }() - - // Handle session auth first if enabled. - if srv.jwtValidator != nil { - // If JWT auth enabled, this is the only unauthenticated (non-debug) endpoint. - if r.Method == "POST" && r.URL.Path == "/auth/exchange-token" { - srv.handleTokenExchange(w, r) - return - } - // Check for session data and error if none. - token := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ") - if !strings.HasPrefix(token, tokenPrefix) { - http.Error(w, "unauthorized", http.StatusUnauthorized) - return - } - - var ok bool - sessionData, ok = srv.getSessionData(token) - if !ok || srv.now().After(sessionData.expiry) { - if srv.verbose { - reason := fmt.Sprintf("exists: %v", ok) - if sessionData != nil { - reason += fmt.Sprintf(", expiry: %v", sessionData.expiry) - } - srv.logf("unauthorized; %s", reason) - } - http.Error(w, "unauthorized", http.StatusUnauthorized) - return - } - } - - if r.Method == "PUT" { - if sessionData != nil && !sessionData.globalNSWrite { - // TODO(tomhjp): support per-namespace writes. - http.Error(w, "forbidden", http.StatusForbidden) - return - } - srv.handlePut(w, r, reqStats) - return - } - if r.Method != "GET" && r.Method != "HEAD" { - http.Error(w, "bad method", http.StatusBadRequest) - return - } - if strings.HasPrefix(r.URL.Path, "/action/") { - srv.handleGetAction(w, r, reqStats) - return - } - if sessionData != nil && r.URL.Path == "/session/stats" { - srv.handleSessionStats(w, sessionData) - return - } - http.Error(w, "not found", http.StatusNotFound) -} - -func (srv *server) getSessionData(token string) (*sessionData, bool) { - srv.sessionsMu.Lock() - defer srv.sessionsMu.Unlock() - sessionData, ok := srv.sessions[token] - return sessionData, ok -} - -func (srv *server) addSessionData(token string, sessionData *sessionData) { - srv.sessionsMu.Lock() - defer srv.sessionsMu.Unlock() - srv.sessions[token] = sessionData - srv.Sessions.Add(1) -} - -func (srv *server) processRequestStats(req *stats, sessionData *sessionData) { - srv.Gets.Add(req.Gets) - srv.GetBytes.Add(req.GetBytes) - srv.GetHits.Add(req.GetHits) - srv.GetAccessBumps.Add(req.GetAccessBumps) - srv.GetHitsInline.Add(req.GetHitsInline) - srv.GetErrs.Add(req.GetErrs) - srv.Puts.Add(req.Puts) - srv.PutErrs.Add(req.PutErrs) - srv.PutsDup.Add(req.PutsDup) - srv.PutsBytes.Add(req.PutsBytes) - srv.PutsInline.Add(req.PutsInline) - - if sessionData != nil { - sessionData.mu.Lock() - defer sessionData.mu.Unlock() - - sessionData.stats.LastUsed = srv.now().UTC() - sessionData.stats.Gets += req.Gets - sessionData.stats.GetBytes += req.GetBytes - sessionData.stats.GetHits += req.GetHits - sessionData.stats.GetAccessBumps += req.GetAccessBumps - sessionData.stats.GetHitsInline += req.GetHitsInline - sessionData.stats.GetErrs += req.GetErrs - sessionData.stats.GetNanos += req.GetNanos - sessionData.stats.Puts += req.Puts - sessionData.stats.PutErrs += req.PutErrs - sessionData.stats.PutsDup += req.PutsDup - sessionData.stats.PutsBytes += req.PutsBytes - sessionData.stats.PutsInline += req.PutsInline - sessionData.stats.PutsNanos += req.PutsNanos + opts = append(opts, + gocached.WithJWTAuth(*jwtIssuer, jwtClaims), + gocached.WithGlobalNamespaceJWTClaims(globalClaims), + ) } -} - -func getHexSuffix(r *http.Request, prefix string) (hexSuffix string, ok bool) { - hexSuffix, _ = strings.CutPrefix(r.RequestURI, prefix) - if !validHex(hexSuffix) { - return "", false - } - return hexSuffix, true -} -func validHex(x string) bool { - if len(x) < 4 || len(x) > 1000 || len(x)%2 == 1 { - return false - } - for i := range x { - b := x[i] - if b >= '0' && b <= '9' || b >= 'a' && b <= 'f' { - continue - } - return false - } - return true -} - -// relAtimeSeconds is how old an access time needs to be before -// we do a DB write to update it. -const relAtimeSeconds = 60 * 60 * 24 // 1 day - -func (srv *server) handleGetAction(w http.ResponseWriter, r *http.Request, stats *stats) { - srv.ActiveGets.Add(1) - defer srv.ActiveGets.Add(-1) - - start := srv.now() - defer func() { - stats.GetNanos += srv.now().Sub(start).Nanoseconds() - }() - stats.Gets++ - ctx := r.Context() - - httpErr := func(msg string, code int) { - http.Error(w, msg, code) - stats.GetErrs++ - } - - actionID, ok := getHexSuffix(r, "/action/") - if !ok { - httpErr("bad request", http.StatusBadRequest) - return - } - if r.Header.Get("Want-Object") != "1" { - httpErr("bad request: missing Want-Object header", http.StatusBadRequest) - return - } - - var sha256hex string - var size int64 - var smallData sql.NullString - var altObjectID string - var accessTime int64 - namespaceID := 0 // global for now; TODO(bradfitz): support namespaces - err := srv.db.QueryRow( - "SELECT b.SHA256, b.BlobSize, b.SmallData, a.AltOutputID, a.AccessTime FROM Actions a, Blobs b WHERE a.NameSpaceID = ? AND a.ActionID = ? AND a.BlobID = b.BlobID", - namespaceID, actionID).Scan( - &sha256hex, &size, &smallData, &altObjectID, &accessTime) + srv, err := gocached.NewServer(opts...) if err != nil { - if errors.Is(err, sql.ErrNoRows) { - http.Error(w, "not found", http.StatusNotFound) - return - } - srv.logf("QueryRow error: %v", err) - httpErr("QueryRow error", http.StatusInternalServerError) - return - } - - // If it's been more than a day since the last access, update the access time. - // This is similar to the Linux "relatime" behavior. - now := srv.now().Unix() - if accessTime < now-relAtimeSeconds { - // TODO(bradfitz): do this async? not worth blocking the caller. - // But we need a mechanism for tests to wait on async work. - srv.sqliteWriteMu.Lock() - _, err := srv.db.Exec("UPDATE Actions SET AccessTime = ? WHERE ActionID = ?", now, actionID) - srv.sqliteWriteMu.Unlock() - if err != nil { - srv.logf("Update AccessTime error: %v", err) - httpErr("internal server error", http.StatusInternalServerError) - return - } - stats.GetAccessBumps++ - } - - stats.GetHits++ - - outputID := cmp.Or(altObjectID, sha256hex) - - w.Header().Set("Content-Type", "application/octet-stream") - w.Header().Set("Content-Length", fmt.Sprint(size)) - w.Header().Set("Go-Output-Id", outputID) - - if r.Method == "HEAD" || size == 0 { - return - } - - if smallData.Valid { - // For small outputs stored inline in the database, we can return them directly. - stats.GetHitsInline++ - stats.GetBytes += size - io.WriteString(w, smallData.String) - return - } - - // Otherwise, for large objects that we know about, we can try to get them - // from our local disk or a peer. - - rc, err := srv.getObjectFromDiskOrPeer(ctx, sha256hex) - if err != nil { - srv.logf("Get object error: %v", actionID, err) - httpErr("Get object error", http.StatusInternalServerError) - return - } - if rc == nil { - // Our database suggested we should've had this object, - // but maybe somebody delete it by hand from the filesystem. - // Just treat it as a cache miss. The background cleanup - // will eventually remove the Action row from the DB - // after identifying it as a dangling reference. - http.Error(w, "not found", http.StatusNotFound) - return - } - stats.GetBytes += size - defer rc.Close() - io.Copy(w, rc) -} - -// getObjectFromDiskOrPeer retrieves the object for the given actionID, either -// from disk or a peer. This is used after a local DB lookup discovers the content -// exists but is not stored in SQLite. -// -// It returns (nil, nil) on miss. -func (srv *server) getObjectFromDiskOrPeer(_ context.Context, sha256hex string) (rc io.ReadCloser, err error) { - if len(sha256hex) != sha256.Size*2 { - return nil, fmt.Errorf("invalid sha256hex %q", sha256hex) - } - diskPath := filepath.Join(srv.dir, sha256hex[:2], sha256hex) - f, err := os.Open(diskPath) - if err != nil { - if os.IsNotExist(err) { - // TODO(bradfitz): search peers, S3, etc. - // For now, just return nil, nil on miss. - return nil, nil - } - return nil, err - } - return f, nil -} - -func (s *server) handlePut(w http.ResponseWriter, r *http.Request, stats *stats) { - s.ActivePuts.Add(1) - defer s.ActivePuts.Add(-1) - - start := s.now() - defer func() { - stats.PutsNanos += s.now().Sub(start).Nanoseconds() - }() - - if r.Method != "PUT" { - http.Error(w, "bad method", http.StatusMethodNotAllowed) - return - } - actionID, outputID, ok := strings.Cut(r.RequestURI[len("/"):], "/") - if !ok || !validHex(actionID) || !validHex(outputID) { - http.Error(w, "bad URI", http.StatusBadRequest) - return - } - if r.ContentLength == -1 { - http.Error(w, "missing Content-Length", http.StatusBadRequest) - return - } - - hasher := sha256.New() - hashingBody := io.TeeReader(r.Body, hasher) - - var smallData []byte - if r.ContentLength <= smallObjectSize { - // Store small objects inline in the database. - var err error - smallData, err = io.ReadAll(hashingBody) - if err != nil { - s.logf("Read content error: %v", err) - http.Error(w, "Read content error", http.StatusInternalServerError) - return - } - if int64(len(smallData)) != r.ContentLength { - // This check is redundant with net/http's validation, but - // for extra clarity. - http.Error(w, "bad content length", http.StatusInternalServerError) - return - } - } else { - // For larger objects, we store them on disk. - if err := s.writeDiskBlob(r.ContentLength, hashingBody); err != nil { - s.logf("Write disk blob error: %v", err) - http.Error(w, "Write disk blob error", http.StatusInternalServerError) - return - } - } - - sha256hex := fmt.Sprintf("%x", hasher.Sum(nil)) - blobSize := r.ContentLength - - s.sqliteWriteMu.Lock() - defer s.sqliteWriteMu.Unlock() - - var blobID int64 - err := s.db.QueryRow(`INSERT INTO Blobs (SHA256, BlobSize, SmallData) - VALUES (?, ?, ?) - ON CONFLICT(SHA256) DO UPDATE SET SHA256=excluded.SHA256 - RETURNING BlobID; -`, sha256hex, blobSize, smallData).Scan(&blobID) - if err != nil { - s.logf("Blobs insert error: %v", err) - stats.PutErrs++ - http.Error(w, "Blobs insert error", http.StatusInternalServerError) - return - } - - // Insert or update the action in the database. - nowUnix := s.now().Unix() - altObjectID := "" - namespace := 0 // global for now; TODO(bradfitz): support namespaces - if sha256hex != outputID { - altObjectID = outputID - } - res, err := s.db.Exec(`INSERT OR IGNORE INTO Actions (NamespaceID, ActionID, BlobID, AltOutputID, CreateTime, AccessTime) - VALUES (?, ?, ?, ?, ?, ?)`, - namespace, - actionID, - blobID, - altObjectID, - nowUnix, - nowUnix, - ) - if err != nil { - s.logf("Actions insert error: %v", err) - stats.PutErrs++ - http.Error(w, "Actions insert error", http.StatusInternalServerError) - return - } - - affected, err := res.RowsAffected() - if err != nil { - s.logf("Actions rows affected error: %v", err) - stats.PutErrs++ - http.Error(w, "Actions rows affected error", http.StatusInternalServerError) - return - } - - if affected == 0 { - stats.PutsDup++ - } - - stats.Puts++ - stats.PutsBytes += r.ContentLength - if smallData != nil { - stats.PutsInline++ - } - - w.WriteHeader(http.StatusNoContent) -} - -// handleTokenExchange handles POST /auth/exchange-token requests to exchange -// a JWT for an access token. Each access token represents a session that will -// last for one hour and have cache stats associated with it. -func (srv *server) handleTokenExchange(w http.ResponseWriter, r *http.Request) { - var req struct { - JWT string `json:"jwt"` - } - // JWTs are often sent in HTTP headers, so 4KiB should ~always be enough. - if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, 4<<10)).Decode(&req); err != nil { - srv.AuthErrs.Add(1) - http.Error(w, "bad request: "+err.Error(), http.StatusBadRequest) - return + log.Fatalf("starting gocached: %v", err) } - jwtClaims, err := srv.jwtValidator.Validate(r.Context(), req.JWT) - if err != nil { - srv.AuthErrs.Add(1) - if srv.verbose { - srv.logf("token exchange: JWT validation error: %v", err) - } - http.Error(w, "unauthorized", http.StatusUnauthorized) - return - } - - globalNSWrite, err := srv.evaluateClaims(jwtClaims) - if err != nil { - srv.AuthErrs.Add(1) - if srv.verbose { - srv.logf("token exchange: %v", err) - } - http.Error(w, "unauthorized", http.StatusUnauthorized) - return - } - - const ttl = time.Hour - // 52 base32 characters, 256 bits of entropy. - accessToken := tokenPrefix + strings.ToLower(rand.Text()+rand.Text()) - srv.addSessionData(accessToken, &sessionData{ - expiry: srv.now().UTC().Add(ttl), - globalNSWrite: globalNSWrite, - claims: jwtClaims, - }) - - resp := map[string]any{ - "access_token": accessToken, - "token_type": "Bearer", - "expires_in": ttl.Seconds(), - } - w.Header().Set("Content-Type", "application/json") - if err := json.NewEncoder(w).Encode(resp); err != nil { - srv.AuthErrs.Add(1) - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - - srv.Auths.Add(1) -} - -func (srv *server) evaluateClaims(claims map[string]any) (globalNSWrite bool, _ error) { - if missing := findMissingClaims(srv.jwtClaims, claims); len(missing) > 0 { - return false, fmt.Errorf("got claims %v; missing required claims: %v", claims, missing) - } - - if missing := findMissingClaims(srv.globalJWTClaims, claims); len(missing) == 0 { - return true, nil - } else if srv.verbose { - srv.logf("token exchange: missing global namespace write claims: %v", missing) - } - - return false, nil -} - -func findMissingClaims(wantClaims map[string]string, gotClaims map[string]any) map[string]string { - missing := make(map[string]string) - for k, want := range wantClaims { - if got, ok := gotClaims[k]; !ok || got != want { - missing[k] = want - } - } - return missing -} - -func (srv *server) handleSessionStats(w http.ResponseWriter, sessionData *sessionData) { - w.Header().Set("Content-Type", "application/json") - if err := json.NewEncoder(w).Encode(sessionData.stats); err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } -} - -func (s *server) sha256Filepath(hash [sha256.Size]byte) string { - hex := fmt.Sprintf("%x", hash) - return filepath.Join(s.dir, hex[:2], hex) -} - -func (s *server) writeDiskBlob(size int64, r io.Reader) (err error) { - nowUnix := s.now().Unix() - tf, err := os.CreateTemp(s.dir, fmt.Sprintf("upload-%d-*", nowUnix)) - if err != nil { - return err - } - defer func() { - if err == nil { - return - } - tf.Close() - os.Remove(tf.Name()) - }() - hasher := sha256.New() - n, err := io.Copy(tf, io.LimitReader(io.TeeReader(r, hasher), size+1)) - if err != nil { - return err - } - if n != size { - return fmt.Errorf("wrote %d bytes; wanted %d", n, size) - } - if err := tf.Close(); err != nil { - return err - } - var hash [sha256.Size]byte - hasher.Sum(hash[:0]) - - target := s.sha256Filepath(hash) - if err := os.MkdirAll(filepath.Dir(target), 0750); err != nil { - return err - } - return os.Rename(tf.Name(), target) -} - -type countAndSize struct { - Count int64 // number of actions - Size int64 // total size of all actions' blobs (even if shared by other actions) -} - -func (cs countAndSize) String() string { - if cs.Count == 0 { - return "0 objects, 0 bytes" - } - return fmt.Sprintf("%d objects, %s", cs.Count, bytesFmt(cs.Size)) -} - -type usageStats struct { - // ActionsLE is a histogram of the actions in the DB by their access time. - // - // The key is a Prometheus-style histogram "less than" value. That is, if - // there are map keys for 24h and 48h, the latter includes the sum of the - // 24h values as well. - // - // The map keys are day-granularity, as the access time is only updated once - // it's over a day old. - // - // So the map keys are 24h, 48h, 96h, 168h (7d), 336h (14d), 720h - // (30d), and 2160h (90d) and math.MaxInt64 for infinity. - ActionsLE map[time.Duration]countAndSize - - // MissingBlobRows is the number of rows in the Actions table that - // reference a BlobID that doesn't exist in the Blobs table. - // This should always be zero in a healthy system. - MissingBlobRows int -} - -func (us *usageStats) All() countAndSize { return us.ActionsLE[math.MaxInt64] } - -const day = 24 * time.Hour - -var standardDurs = []time.Duration{ - 1 * day, - 2 * day, - 4 * day, - 7 * day, - 14 * day, - 30 * day, - 90 * day, - math.MaxInt64, -} - -func (s *server) usageStats() (_ *usageStats, err error) { - defer func() { - if err != nil { - s.logf("usageStats error: %v", err) - } - }() - - st := &usageStats{ - ActionsLE: make(map[time.Duration]countAndSize), - } - - // Build the durations to use for the histogram. - // The math.MaxInt64 value is always included. - // If s.maxAge is set, we ignore sizes above that, except - // for the math.MaxInt64 value. - var durs []time.Duration - if s.maxAge == 0 { - durs = standardDurs - } else { - durs = make([]time.Duration, 0, len(standardDurs)+1) - durs = append(durs, s.maxAge) - for _, d := range standardDurs { - if d < s.maxAge || d == math.MaxInt64 { - durs = append(durs, d) - } - } - slices.Sort(durs) - } - - now := s.now().Unix() - rows, err := s.db.Query( - "SELECT a.BlobID, a.AccessTime, b.BlobSize FROM Actions a LEFT JOIN Blobs b ON a.BlobID = b.BlobID") - if err != nil { - return nil, fmt.Errorf("query Actions: %w", err) - } - var blobID int64 - var accessTime int64 - var blobSize sql.NullInt64 - for rows.Next() { - if err := rows.Scan(&blobID, &accessTime, &blobSize); err != nil { - return nil, fmt.Errorf("rows.Scan: %w", err) - } - if !blobSize.Valid { - st.MissingBlobRows++ - continue - } - - dur := time.Duration(now-accessTime) * time.Second - if dur < 0 { - dur = 0 - } - for _, d := range durs { - if dur < d { - was := st.ActionsLE[d] - was.Count++ - was.Size += blobSize.Int64 - st.ActionsLE[d] = was - } - } - } - if err := rows.Err(); err != nil { - return nil, fmt.Errorf("rows.Next: %w", err) - } - - s.lastUsage.Store(st) - all := st.All() - s.BlobCount.Set(all.Count) - s.BlobBytes.Set(all.Size) - return st, nil -} - -type cleanCandidate struct { - BlobID int64 - Age time.Duration - BlobSize int64 // size of the blob, in bytes -} - -func (s *server) cleanCandidates(olderThan time.Duration, limit int64) ([]cleanCandidate, error) { - now := s.now() - nowUnix := now.Unix() - cutoff := now.Add(-olderThan).Unix() - - rows, err := s.db.Query(` - SELECT b.BlobID, MAX(a.AccessTime), b.BlobSize - FROM Blobs b LEFT JOIN Actions a ON b.BlobID = a.BlobID - GROUP BY b.BlobID - HAVING MAX(a.AccessTime) <= ? - ORDER BY MAX(a.AccessTime) - LIMIT ?`, cutoff, limit) - if err != nil { - return nil, fmt.Errorf("query clean candidates: %w", err) - } - defer rows.Close() - - var candidates []cleanCandidate - var accessTime int64 - for rows.Next() { - var c cleanCandidate - if err := rows.Scan(&c.BlobID, &accessTime, &c.BlobSize); err != nil { - return nil, fmt.Errorf("rows.Scan: %w", err) - } - c.Age = time.Duration(nowUnix-accessTime) * time.Second - candidates = append(candidates, c) - } - if err := rows.Err(); err != nil { - return nil, fmt.Errorf("rows.Next: %w", err) - } - - return candidates, nil -} - -func (srv *server) deleteBlobs(blobIDs ...int64) error { - srv.sqliteWriteMu.Lock() - defer srv.sqliteWriteMu.Unlock() - - tx, err := srv.db.Begin() - if err != nil { - return fmt.Errorf("delete blob Begin: %w", err) - } - defer tx.Rollback() - - var sumBytes int64 - for _, blobID := range blobIDs { - var sha256Hex string - var blobSize int64 - if err := tx.QueryRow("SELECT SHA256, BlobSize FROM Blobs WHERE BlobID = ?", blobID).Scan(&sha256Hex, &blobSize); err != nil && !errors.Is(err, sql.ErrNoRows) { - return fmt.Errorf("querying blob SHA256: %w", err) - } - sumBytes += blobSize - if _, err := tx.Exec("DELETE FROM Blobs WHERE BlobID = ?", blobID); err != nil { - return fmt.Errorf("deleting blob: %w", err) - } - if _, err := tx.Exec("DELETE FROM Actions WHERE BlobID = ?", blobID); err != nil { - return fmt.Errorf("deleting actions: %w", err) - } - var hash [sha256.Size]byte - if _, err := hex.Decode(hash[:], []byte(sha256Hex)); err == nil { - if err := os.Remove(srv.sha256Filepath(hash)); err != nil && !os.IsNotExist(err) { - return fmt.Errorf("removing disk file: %w", err) - } - } - } - if err := tx.Commit(); err != nil { - return err - } - - srv.EvictedBlobs.Add(int64(len(blobIDs))) - srv.EvictedBytes.Add(sumBytes) - - return nil -} - -func (srv *server) cleanOldObjects(us *usageStats) (countAndSize, error) { - var zero countAndSize - var ret countAndSize - - all := us.ActionsLE[math.MaxInt64] - if srv.verbose { - srv.logf("current usage stats: %v", all) - last := all - for _, d := range slices.Sorted(maps.Keys(us.ActionsLE)) { - if d == math.MaxInt64 { - continue // skip infinity - } - c := us.ActionsLE[d] - srv.logf(" <=%v: %v", durFmt(d), c) - if last == c { - break - } - last = c - } - } - - // First clean things that are just too old. - if srv.maxAge > 0 { - if toDelete := all.Count - us.ActionsLE[srv.maxAge].Count; toDelete > 0 { - srv.logf("Cleaning %d objects older than %v ...", toDelete, durFmt(srv.maxAge)) - candidates, err := srv.cleanCandidates(srv.maxAge, toDelete+1) - if err != nil { - return zero, fmt.Errorf("getting clean candidates: %v", err) - } - blobIDs := make([]int64, 0, len(candidates)) - var sumSize int64 - for _, c := range candidates { - blobIDs = append(blobIDs, c.BlobID) - sumSize += c.BlobSize - } - if err := srv.deleteBlobs(blobIDs...); err != nil { - return zero, fmt.Errorf("deleting old blobs: %v", err) - } - all.Count -= int64(len(candidates)) - all.Size -= sumSize - ret.Count += int64(len(candidates)) - ret.Size += sumSize - } - } - - for srv.maxSize > 0 && all.Size > srv.maxSize { - toClean := all.Size - srv.maxSize - if srv.verbose { - srv.logf("need to clean %v to get under max size of %v ...", - bytesFmt(toClean), bytesFmt(srv.maxSize)) - } - - var batchBytes int64 - var blobIDs []int64 - candidates, err := srv.cleanCandidates(0, 10000) - if err != nil { - return zero, fmt.Errorf("getting clean candidates: %v", err) - } - for _, c := range candidates { - blobIDs = append(blobIDs, c.BlobID) - batchBytes += c.BlobSize - if batchBytes >= toClean { - break - } - } - if err := srv.deleteBlobs(blobIDs...); err != nil { - return zero, fmt.Errorf("deleting old blobs: %v", err) - } - - ret.Count += int64(len(blobIDs)) - ret.Size += batchBytes - all.Count -= int64(len(blobIDs)) - all.Size -= batchBytes - - if len(blobIDs) == len(candidates) { - // We didn't find enough candidates to delete. - // Just stop here. - srv.logf("[unexpected] didn't find enough candidates to delete") - break - } - } - - return ret, nil -} - -func (srv *server) runCleanLoop() { - for { - select { - case <-srv.shutdownCtx.Done(): - return - case <-time.After(5 * time.Minute): - } - - us, err := srv.usageStats() - if err != nil { - srv.logf("error getting usage stats: %v", err) - continue - } - - res, err := srv.cleanOldObjects(us) - if err != nil { - srv.logf("error cleaning old objects: %v", err) - continue - } - if res.Count > 0 { - srv.logf("cleaned %v", res) - srv.usageStats() // for side effect of updating lastUsage - } - - } -} - -func (srv *server) runCleanSessionsLoop() { - for { - select { - case <-srv.shutdownCtx.Done(): - return - case <-time.After(time.Hour): - } - - srv.sessionsMu.Lock() - count := len(srv.sessions) - var deleted int - for token, metadata := range srv.sessions { - if time.Now().After(metadata.expiry) { - delete(srv.sessions, token) - deleted++ - srv.Sessions.Add(-1) - } - } - srv.sessionsMu.Unlock() - srv.logf("cleaned up %d/%d access tokens", deleted, count) - } -} - -func durFmt(d time.Duration) string { - days := int(d.Hours() / 24) - if days > 0 { - return fmt.Sprintf("%dd", days) - } - return d.String() -} - -func bytesFmt(n int64) string { - if n >= 1<<30 { - return fmt.Sprintf("%.1f GiB", float64(n)/(1<<30)) - } - if n >= 1<<20 { - return fmt.Sprintf("%.1f MiB", float64(n)/(1<<20)) - } - if n >= 1<<10 { - return fmt.Sprintf("%.1f KiB", float64(n)/(1<<10)) - } - return fmt.Sprintf("%d bytes", n) -} - -func (srv *server) serveUsage(w http.ResponseWriter, r *http.Request) { - if r.Method == "POST" { - // For side effect of updating lastUsage. - _, err := srv.usageStats() + if *debugListen != "" { + debugLn, err := net.Listen("tcp", *debugListen) if err != nil { - http.Error(w, "error getting usage stats: "+err.Error(), http.StatusInternalServerError) - return - } - } - - us := srv.lastUsage.Load() - if us == nil { - http.Error(w, "no usage stats available", http.StatusInternalServerError) - return - } - - // Print out an HTML table of the usage stats, sorted by age. - w.Header().Set("Content-Type", "text/html; charset=utf-8") - fmt.Fprintf(w, "

gocached usage stats

\n") - fmt.Fprintf(w, "

Current usage: %v of limit %v

\n", - us.All(), bytesFmt(srv.maxSize)) - - fmt.Fprintf(w, "\n") - fmt.Fprintf(w, "\n") - for _, d := range slices.Sorted(maps.Keys(us.ActionsLE)) { - var title string - if d == math.MaxInt64 { - title = "all" - } else { - title = "<= " + durFmt(d) - } - c := us.ActionsLE[d] - fmt.Fprintf(w, "\n", - title, c.Count, bytesFmt(c.Size)) - } - fmt.Fprintf(w, "
AgeCountSize
%s%d%s
\n") -} - -func (srv *server) serveSessions(w http.ResponseWriter, r *http.Request) { - if r.Method != "GET" { - http.Error(w, "bad method", http.StatusMethodNotAllowed) - return - } - - srv.sessionsMu.RLock() - // Make a copy of all session data (excluding the mutex). - sessions := make([]*sessionData, 0, len(srv.sessions)) - for _, v := range srv.sessions { - v.mu.Lock() - sessions = append(sessions, &sessionData{ - expiry: v.expiry, - globalNSWrite: v.globalNSWrite, - claims: v.claims, - stats: v.stats, - }) - v.mu.Unlock() - } - srv.sessionsMu.RUnlock() - - w.Header().Set("Content-Type", "text/html; charset=utf-8") - fmt.Fprintf(w, "

gocached sessions

\n") - fmt.Fprintf(w, "

JWT issuer: %s

\n", *jwtIssuer) - fmt.Fprintf(w, "

JWT claims required: %v

\n", srv.jwtClaims) - fmt.Fprintf(w, "

JWT global write claims required: %v

\n", srv.globalJWTClaims) - fmt.Fprintf(w, "

Number of sessions: %d

\n", len(sessions)) - - fmt.Fprintf(w, "\n") - fmt.Fprintf(w, "\n") - slices.SortFunc(sessions, func(a, b *sessionData) int { - return a.stats.LastUsed.Compare(b.stats.LastUsed) - }) - for _, d := range slices.Backward(sessions) { - lastUsed := "never" - if !d.stats.LastUsed.IsZero() { - lastUsed = durFmt(time.Since(d.stats.LastUsed)) + " ago" + log.Fatalf("debug listen: %v", err) } - statsJSON, _ := json.MarshalIndent(d.stats, "", " ") - claimsJSON, _ := json.MarshalIndent(d.claims, "", " ") - fmt.Fprintf(w, "\n", - lastUsed, d.expiry.Format(time.RFC3339), d.globalNSWrite, statsJSON, claimsJSON) + go func() { + log.Fatal(http.Serve(debugLn, http.HandlerFunc(srv.ServeHTTPDebug))) + }() } - fmt.Fprintf(w, "
Last usedExpiry timeGlobal writeStatsClaims
%s%s%v
%s
%s
\n") -} - -// expvarCounterMetric is a Prometheus counter metric backed by an expvar.Int. -type expvarCounterMetric struct { - desc *prometheus.Desc - v *expvar.Int -} - -var _ prometheus.Metric = (*expvarCounterMetric)(nil) - -func (m *expvarCounterMetric) Desc() *prometheus.Desc { return m.desc } - -func (m *expvarCounterMetric) Write(out *dto.Metric) error { - val := float64(m.v.Value()) - out.Counter = &dto.Counter{Value: &val} - return nil -} -// expvarGaugeMetric is a Prometheus gauge metric backed by an expvar.Int. -type expvarGaugeMetric struct { - desc *prometheus.Desc - v *expvar.Int -} - -var _ prometheus.Metric = (*expvarGaugeMetric)(nil) - -func (m *expvarGaugeMetric) Desc() *prometheus.Desc { return m.desc } - -func (m *expvarGaugeMetric) Write(out *dto.Metric) error { - val := float64(m.v.Value()) - out.Gauge = &dto.Gauge{Value: &val} - return nil -} - -// singleMetricCollector is a Prometheus collector that collects a single metric. -type singleMetricCollector struct { - metric prometheus.Metric + log.Printf("gocached: listening on %s ...", *listen) + log.Fatal(http.ListenAndServe(*listen, srv)) } - -var _ prometheus.Collector = singleMetricCollector{} - -func (c singleMetricCollector) Describe(ch chan<- *prometheus.Desc) { ch <- c.metric.Desc() } -func (c singleMetricCollector) Collect(ch chan<- prometheus.Metric) { ch <- c.metric } diff --git a/cmd/gocached/gocached_test.go b/cmd/gocached/gocached_test.go index f28a0ba..30dbde8 100644 --- a/cmd/gocached/gocached_test.go +++ b/cmd/gocached/gocached_test.go @@ -1,699 +1,14 @@ +// Copyright (c) Tailscale Inc & AUTHORS +// SPDX-License-Identifier: BSD-3-Clause + package main import ( - "bytes" - "context" - "crypto" - "crypto/ecdsa" - "crypto/elliptic" - "crypto/rand" - "encoding/json" - "expvar" - "fmt" - "io" - "math" - "net" - "net/http" - "net/http/httptest" - "os" - "path/filepath" - "slices" - "strings" - "sync" - "sync/atomic" "testing" - "time" - "github.com/go-jose/go-jose/v4" - "github.com/golang-jwt/jwt/v5" "github.com/google/go-cmp/cmp" - ijwt "github.com/tailscale/tb/cmd/gocached/internal/jwt" - "github.com/tailscale/tb/gocache/cachers" ) -// sha256OfEmpty is the SHA-256 hash of an empty string, used as a well-known -// value in SQLite to store bytes, as it's common. -const sha256OfEmpty = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" - -type tester struct { - t testing.TB - srv *server - hs *httptest.Server - - issuer string - audience string - createJWT func(claims jwt.MapClaims, signingKey *ecdsa.PrivateKey) string - - timeMu sync.Mutex - curTime time.Time -} - -func (t *tester) Logf(format string, args ...any) { - t.t.Logf(format, args...) -} - -func (t *tester) now() time.Time { - t.timeMu.Lock() - defer t.timeMu.Unlock() - return t.curTime -} - -func (t *tester) advanceClock(d time.Duration) { - t.timeMu.Lock() - defer t.timeMu.Unlock() - t.curTime = t.curTime.Add(d) -} - -func (t *tester) mkClient() *cachers.HTTPClient { - clientCacheDir := t.t.TempDir() - return &cachers.HTTPClient{ - BaseURL: t.hs.URL, - Disk: &cachers.DiskCache{ - Dir: clientCacheDir, - Logf: func(format string, args ...any) { - t.Logf("client-disk: "+format, args...) - }, - }, - } -} - -func (st *tester) usageStats() *usageStats { - st.t.Helper() - stats, err := st.srv.usageStats() - if err != nil { - st.t.Fatalf("usageStats: %v", err) - } - return stats -} - -func (st *tester) cleanOldObjects() countAndSize { - st.t.Helper() - stats, err := st.srv.cleanOldObjects(st.usageStats()) - if err != nil { - st.t.Fatalf("cleanOldObjects: %v", err) - } - return stats -} - -func (st *tester) diskFiles() []string { - st.t.Helper() - var ret []string - - err := filepath.Walk(st.srv.dir, func(path string, fi os.FileInfo, err error) error { - if err != nil { - return err - } - if !fi.Mode().IsRegular() || strings.HasPrefix(fi.Name(), ".") || strings.HasPrefix(fi.Name(), "gocached") { - return nil - } - ret = append(ret, fi.Name()) - return nil - }) - if err != nil { - st.t.Fatalf("Walk: %v", err) - } - slices.Sort(ret) - return ret -} - -// wantMetric is a helper to check an expvar.Int metric and reset it -// for future tests. -func (st *tester) wantMetric(m *expvar.Int, want int64) { - st.t.Helper() - if got := m.Value(); got != want { - st.t.Errorf("metric = %d, want %d", got, want) - } - m.Set(0) -} - -func (st *tester) wantPut(c *cachers.HTTPClient, actionID, outputID string, val string) { - ctx := context.Background() - st.t.Helper() - clientDiskPath, err := c.Put(ctx, actionID, outputID, int64(len(val)), strings.NewReader(val)) - if err != nil { - st.t.Fatalf("Put: %v", err) - } - if clientDiskPath == "" { - st.t.Fatal("Put returned empty disk path") - } - st.wantMetric(&st.srv.Puts, 1) - wrote, err := os.ReadFile(clientDiskPath) - if err != nil { - st.t.Fatalf("ReadFile: %v", err) - } - if string(wrote) != val { - st.t.Errorf("ReadFile got %q, want %q", wrote, val) - } -} - -func (st *tester) wantGet(c *cachers.HTTPClient, actionID, outputID, wantVal string) { - ctx := context.Background() - st.t.Helper() - gotOutputID, diskPath, err := c.Get(ctx, actionID) - if err != nil { - st.t.Fatalf("Get: %v", err) - } - if gotOutputID != outputID { - st.t.Errorf("Get got outputID %q, want %q", gotOutputID, outputID) - } - if diskPath == "" { - st.t.Fatal("Get returned empty disk path") - } - wrote, err := os.ReadFile(diskPath) - if err != nil { - st.t.Fatalf("ReadFile: %v", err) - } - if string(wrote) != wantVal { - st.t.Errorf("ReadFile got %q, want %q", wrote, wantVal) - } -} - -func (st *tester) wantGetMiss(c *cachers.HTTPClient, actionID string) { - ctx := context.Background() - st.t.Helper() - gotOutputID, diskPath, err := c.Get(ctx, actionID) - if err != nil { - st.t.Fatalf("Get: %v", err) - } - if gotOutputID != "" { - st.t.Errorf("Get got outputID %q; want empty", gotOutputID) - } - if diskPath != "" { - st.t.Fatalf("Get returned disk path %q; want empty", diskPath) - } -} - -func newServerTester(t testing.TB) *tester { - st := &tester{ - t: t, - curTime: time.Unix(1234, 0), - } - - var err error - dir := t.TempDir() - st.srv, err = newServer(dir) - if err != nil { - t.Fatalf("newServer: %v", err) - } - st.srv.logf = t.Logf - st.srv.verbose = true - st.srv.clock = st.now - - st.hs = httptest.NewServer(st.srv) - t.Cleanup(st.hs.Close) - - return st -} - -// startOIDCServer starts a mock OIDC server and configures gocached to use it -// for JWT validation. The provided publicKey is what JWT signatures will be -// validated against. Use st.createJWT to create signed JWTs. -func (st *tester) startOIDCServer(publicKey crypto.PublicKey) { - st.t.Helper() - mux := http.NewServeMux() - srv := httptest.NewServer(mux) - st.t.Cleanup(srv.Close) - - st.issuer = fmt.Sprintf("http://%s", srv.Listener.Addr().String()) - st.audience = gocachedAudience - - mux.HandleFunc("/.well-known/openid-configuration", func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(map[string]any{ - "issuer": st.issuer, - "jwks_uri": fmt.Sprintf("%s/jwks", st.issuer), - }) - }) - mux.HandleFunc("/jwks", func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - json.NewEncoder(w).Encode(map[string]any{ - "keys": []jose.JSONWebKey{ - { - Key: publicKey, - KeyID: "test-key", - Algorithm: "ES256", - Use: "sig", - }, - }, - }) - }) - - st.srv.jwtValidator = ijwt.NewJWTValidator(st.issuer, st.audience) - st.srv.jwtValidator.Logf = st.Logf - if err := st.srv.jwtValidator.RunUpdateJWKSLoop(st.srv.shutdownCtx); err != nil { - st.t.Fatalf("failed to start JWKS loop for JWT validator: %v", err) - } - - st.createJWT = func(claims jwt.MapClaims, signingKey *ecdsa.PrivateKey) string { - st.t.Helper() - unsignedTk := &jwt.Token{ - Header: map[string]any{ - "typ": "JWT", - "alg": jwt.SigningMethodES256.Alg(), - "kid": "test-key", - }, - Claims: claims, - Method: jwt.SigningMethodES256, - } - tk, err := unsignedTk.SignedString(signingKey) - if err != nil { - st.t.Fatalf("error signing token: %v", err) - } - - return tk - } -} - -func TestServer(t *testing.T) { - st := newServerTester(t) - - ctx := context.Background() - - // Make two clients (imagine: two different builder VMs) - c1 := st.mkClient() - c2 := st.mkClient() - - const testActionID = "0001" - const testActionIDMiss = "0002" // this one doesn't exist - const testActionIDBig = "0bbb" // non-inline object - const testActionIDEmpty = "0000" - const testOutputID = "9900" - const testOutputIDBig = "9bbb" - const testOutputIDEmpty = sha256OfEmpty - const testObjectValue = "test data" - testObjectValueBig := strings.Repeat("x", smallObjectSize+1) - - // Populate from the first client. - st.wantPut(c1, testActionID, testOutputID, testObjectValue) - st.wantPut(c1, testActionIDBig, testOutputIDBig, testObjectValueBig) - st.wantPut(c1, testActionIDEmpty, testOutputIDEmpty, "") - - // Read from the second client. - st.wantGet(c2, testActionID, testOutputID, testObjectValue) - st.wantGet(c2, testActionIDBig, testOutputIDBig, testObjectValueBig) - st.wantGet(c2, testActionIDEmpty, testOutputIDEmpty, "") - - // Check metrics - st.wantMetric(&st.srv.Gets, 3) - st.wantMetric(&st.srv.GetHits, 3) - st.wantMetric(&st.srv.GetHitsInline, 1) - - // Do the same get again from the same client. This shouldn't hit the network. - st.wantGet(c2, testActionID, testOutputID, testObjectValue) - st.wantMetric(&st.srv.Gets, 0) - - // Cache miss. This should hit the network and fail. - if _, _, err := c2.Get(ctx, testActionIDMiss); err != nil { - t.Fatalf("miss Get: %v", err) - } - st.wantMetric(&st.srv.Gets, 1) - st.wantMetric(&st.srv.GetHits, 0) - - // Check that access time gets updated. - // Do it from a fresh client without a disk cache. - st.wantMetric(&st.srv.GetAccessBumps, 0) - st.advanceClock(relAtimeSeconds * 2 * time.Second) // advance clock by 2 days - c3 := st.mkClient() - st.wantGet(c3, testActionID, testOutputID, testObjectValue) - st.wantMetric(&st.srv.GetAccessBumps, 1) - - // Get usage stats. - stats, err := st.srv.usageStats() - if err != nil { - t.Fatalf("usageStats: %v", err) - } - want := &usageStats{ - MissingBlobRows: 0, - ActionsLE: map[time.Duration]countAndSize{ - 24 * time.Hour: {Count: 1, Size: 9}, - 48 * time.Hour: {Count: 1, Size: 9}, - 96 * time.Hour: {Count: 3, Size: 1034}, - 168 * time.Hour: {Count: 3, Size: 1034}, - 336 * time.Hour: {Count: 3, Size: 1034}, - 720 * time.Hour: {Count: 3, Size: 1034}, - 2160 * time.Hour: {Count: 3, Size: 1034}, - math.MaxInt64: {Count: 3, Size: 1034}, - }, - } - if diff := cmp.Diff(stats, want); diff != "" { - t.Errorf("usageStats mismatch (-got +want):\n%s", diff) - } - - st.advanceClock(relAtimeSeconds * 2 * time.Second) // advance clock by 2 days -} - -func TestCleanCandidates(t *testing.T) { - st := newServerTester(t) - - // Populate some data. - c1 := st.mkClient() - st.wantPut(c1, "0001", "9901", "1") - st.advanceClock(24 * time.Hour) - st.wantPut(c1, "0002", "9902", "22") - st.advanceClock(24 * time.Hour) - st.wantPut(c1, "0003", "9903", "333") - st.advanceClock(24 * time.Hour) - st.wantPut(c1, "0004", "9904", strings.Repeat("x", smallObjectSize+1)) - - const day = 24 * time.Hour - - tests := []struct { - maxAge time.Duration - limit int64 - want []cleanCandidate - }{ - { - maxAge: 0, - limit: 100, - want: []cleanCandidate{ - {BlobID: 1, Age: 3 * day, BlobSize: 1}, - {BlobID: 2, Age: 2 * day, BlobSize: 2}, - {BlobID: 3, Age: 1 * day, BlobSize: 3}, - {BlobID: 4, Age: 0, BlobSize: smallObjectSize + 1}, - }, - }, - { - maxAge: 25 * time.Hour, - limit: 100, - want: []cleanCandidate{ - {BlobID: 1, Age: 3 * day, BlobSize: 1}, - {BlobID: 2, Age: 2 * day, BlobSize: 2}, - }, - }, - { - maxAge: 0, - limit: 2, - want: []cleanCandidate{ - {BlobID: 1, Age: 3 * day, BlobSize: 1}, - {BlobID: 2, Age: 2 * day, BlobSize: 2}, - }, - }, - } - for _, tt := range tests { - t.Run(fmt.Sprintf("maxAge=%v,limit=%d", tt.maxAge, tt.limit), func(t *testing.T) { - candidates, err := st.srv.cleanCandidates(tt.maxAge, tt.limit) - if err != nil { - t.Fatal(err) - } - if diff := cmp.Diff(candidates, tt.want); diff != "" { - t.Errorf("cleanCandidates mismatch (-got +want):\n%s", diff) - - } - }) - } -} - -func TestCleanOldObjectsByAge(t *testing.T) { - st := newServerTester(t) - st.srv.maxAge = 24 * time.Hour - - // Populate some data. - c1 := st.mkClient() - st.wantPut(c1, "0001", "9901", strings.Repeat("x", smallObjectSize+1)) - st.advanceClock(25 * time.Hour) - st.wantPut(c1, "0002", "9902", strings.Repeat("x", smallObjectSize+2)) - st.wantPut(c1, "0003", "9903", "small") - smallLen := int64(len("small")) - - st1 := st.usageStats() - if all, want := st1.All(), (countAndSize{Count: 3, Size: smallObjectSize*2 + 3 + smallLen}); all != want { - t.Errorf("usageStats: %v; want %v", all, want) - } - if got, want := st.diskFiles(), []string{"333092a3daf718ed8f38a94e302df139edd4e3b5da4239a497995683942cf28c", "c6d8e9905300876046729949cc95c2385221270d389176f7234fe7ac00c4e430"}; !slices.Equal(got, want) { - t.Errorf("diskFiles: %v; want %v", got, want) - } - - clean1 := st.cleanOldObjects() - if clean1.Count != 1 || clean1.Size != smallObjectSize+1 { - t.Errorf("cleanOldObjects got %v, want {Count: 1, Size: %d}", clean1, smallObjectSize+1) - } - clean2 := st.cleanOldObjects() - if clean2.Count != 0 || clean2.Size != 0 { - t.Errorf("cleanOldObjects got %v, want {Count: 0, Size: 0}", clean2) - } - - st2 := st.usageStats() - if all, want := st2.All(), (countAndSize{Count: 2, Size: smallObjectSize + 2 + smallLen}); all != want { - t.Errorf("usageStats after clean: %v; want %v", all, want) - } - if got, want := st.diskFiles(), []string{"333092a3daf718ed8f38a94e302df139edd4e3b5da4239a497995683942cf28c"}; !slices.Equal(got, want) { - t.Errorf("diskFiles after clean: %v; want %v", got, want) - } -} - -func TestCleanOldObjectsBySize(t *testing.T) { - st := newServerTester(t) - - // Populate some data. - c1 := st.mkClient() - st.wantPut(c1, "0001", "9901", "1") - st.advanceClock(time.Second) - st.wantPut(c1, "0002", "9902", "22") - st.advanceClock(time.Second) - st.wantPut(c1, "0003", "9903", "333") - st.advanceClock(time.Second) - st.wantPut(c1, "0004", "9904", "4444") - st.advanceClock(time.Second) - - st1 := st.usageStats() - if all, want := st1.All(), (countAndSize{Count: 4, Size: 10}); all != want { - t.Errorf("usageStats: %v; want %v", all, want) - } - - clean1 := st.cleanOldObjects() - if clean1.Count != 0 || clean1.Size != 0 { - t.Errorf("cleanOldObjects got %v, want no clean", clean1) - } - - st.srv.maxSize = 8 // the only way get to 8 or under is by deleting "1" and "22" (3 bytes) - - if got, want := st.cleanOldObjects(), (countAndSize{Count: 2, Size: 3}); got != want { - t.Errorf("cleanOldObjects got %v, want %v", got, want) - } - if got, want := st.usageStats().All(), (countAndSize{Count: 2, Size: 7}); got != want { - t.Errorf("usageStats: %v; want %v", got, want) - } -} - -func TestClientConnReuse(t *testing.T) { - st := newServerTester(t) - - var numDials atomic.Int32 - tr := http.DefaultTransport.(*http.Transport).Clone() - tr.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) { - num := numDials.Add(1) - t.Logf("DialContext #%d for %s %s", num, network, addr) - var std net.Dialer - return std.DialContext(ctx, network, addr) - } - t.Cleanup(func() { tr.CloseIdleConnections() }) - - c1 := st.mkClient() - c1.HTTPClient = &http.Client{Transport: tr} - const missAction = "0001" - st.wantGetMiss(c1, missAction) - st.wantGetMiss(c1, missAction) - st.wantGetMiss(c1, missAction) - st.wantPut(c1, "0001", "9901", "1") - st.wantGet(c1, "0001", "9901", "1") - if got := numDials.Load(); got != 1 { - t.Errorf("numDials = %d; want 1", got) - } -} - -func TestExchangeToken(t *testing.T) { - // Generate private keys outside of the loop for speed. - privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatalf("error generating OIDC server private key: %v", err) - } - otherPrivateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatalf("error generating OIDC server private key: %v", err) - } - - for name, tc := range map[string]struct { - mutateClaims func(jwt.MapClaims) - signingKey *ecdsa.PrivateKey - wantStatusCode int - wantWrite bool - }{ - // Base case: no mutation. - "valid_read": { - wantStatusCode: http.StatusOK, - wantWrite: false, - }, - // Additional claim needed for write scope. - "valid_write": { - mutateClaims: func(cl jwt.MapClaims) { - cl["ref"] = "refs/heads/main" - }, - wantStatusCode: http.StatusOK, - wantWrite: true, - }, - // Every other test makes one mutation from the base case that should cause failure. - "missing_sub": { - mutateClaims: func(cl jwt.MapClaims) { - delete(cl, "sub") - }, - wantStatusCode: http.StatusUnauthorized, - }, - "invalid_sub": { - mutateClaims: func(cl jwt.MapClaims) { - cl["sub"] = "user456" - }, - wantStatusCode: http.StatusUnauthorized, - }, - "invalid_iss": { - mutateClaims: func(cl jwt.MapClaims) { - cl["iss"] = "invalid_issuer" - }, - wantStatusCode: http.StatusUnauthorized, - }, - "invalid_aud": { - mutateClaims: func(cl jwt.MapClaims) { - cl["aud"] = "invalid_audience" - }, - wantStatusCode: http.StatusUnauthorized, - }, - "not_yet_valid": { - mutateClaims: func(cl jwt.MapClaims) { - cl["nbf"] = jwt.NewNumericDate(time.Now().Add(10 * time.Minute)) - }, - wantStatusCode: http.StatusUnauthorized, - }, - "expired": { - mutateClaims: func(cl jwt.MapClaims) { - cl["exp"] = jwt.NewNumericDate(time.Now().Add(-time.Minute)) - }, - wantStatusCode: http.StatusUnauthorized, - }, - "invalid_signature": { - signingKey: otherPrivateKey, - wantStatusCode: http.StatusUnauthorized, - }, - } { - t.Run(name, func(t *testing.T) { - st := newServerTester(t) - st.startOIDCServer(privateKey.Public()) - st.srv.jwtClaims = map[string]string{ - "sub": "user123", - } - st.srv.globalJWTClaims = map[string]string{ - "sub": "user123", - "ref": "refs/heads/main", - } - - // Generate JWT. - tokenClaims := jwt.MapClaims{ - "sub": "user123", - "iss": st.issuer, - "aud": st.audience, - "nbf": jwt.NewNumericDate(time.Now().Add(-time.Minute)), - "exp": jwt.NewNumericDate(time.Now().Add(time.Hour)), - } - if tc.mutateClaims != nil { - tc.mutateClaims(tokenClaims) - } - signingKey := privateKey - if tc.signingKey != nil { - signingKey = tc.signingKey - } - body, err := json.Marshal(map[string]any{ - "jwt": st.createJWT(tokenClaims, signingKey), - }) - if err != nil { - t.Fatalf("error marshaling request body: %v", err) - } - - // Exchange JWT for access token. - req, err := http.NewRequest("POST", st.hs.URL+"/auth/exchange-token", bytes.NewReader(body)) - if err != nil { - t.Fatalf("error creating request: %v", err) - } - resp, err := http.DefaultClient.Do(req) - if err != nil { - t.Fatalf("error making request: %v", err) - } - defer resp.Body.Close() - if resp.StatusCode != tc.wantStatusCode { - t.Fatalf("unexpected status code: want %d, got %d", tc.wantStatusCode, resp.StatusCode) - } - body, err = io.ReadAll(resp.Body) - if err != nil { - t.Fatalf("error reading response body: %v", err) - } - - if tc.wantStatusCode != http.StatusOK { - if string(body) != "unauthorized\n" { - t.Fatalf("unexpected error body: %s", string(body)) - } - - // No access token to do further checks with; test finished. - return - } - - // Check returned access token. - var d struct { - AccessToken string `json:"access_token"` - } - if err := json.Unmarshal(body, &d); err != nil { - t.Fatalf("error decoding response body: %v", err) - } - if d.AccessToken == "" { - t.Fatalf("expected access_token in response, got %s", string(body)) - } - - cl := st.mkClient() - if _, _, err := cl.Get(t.Context(), "abc123"); err == nil { - t.Fatalf("Get without access token succeeded unexpectedly") - } - - cl.AccessToken = d.AccessToken - st.wantGetMiss(cl, "abc123") - - if tc.wantWrite { - st.wantPut(cl, "abc123", "def456", "data789") - st.wantGet(cl, "abc123", "def456", "data789") - } else { - if _, err := cl.Put(t.Context(), "abc123", "def456", 0, nil); err == nil { - t.Fatalf("Put without write scope succeeded unexpectedly") - } - } - - // Check session stats. - reqStats, err := http.NewRequest("GET", st.hs.URL+"/session/stats", nil) - if err != nil { - t.Fatalf("error creating stats request: %v", err) - } - reqStats.Header.Set("Authorization", "Bearer "+d.AccessToken) - respStats, err := http.DefaultClient.Do(reqStats) - if err != nil { - t.Fatalf("error making stats request: %v", err) - } - defer respStats.Body.Close() - if respStats.StatusCode != http.StatusOK { - t.Fatalf("unexpected stats status code: want %d, got %d", http.StatusOK, respStats.StatusCode) - } - bodyStats, err := io.ReadAll(respStats.Body) - if err != nil { - t.Fatalf("error reading stats response body: %v", err) - } - var stats stats - if err := json.Unmarshal(bodyStats, &stats); err != nil { - t.Fatalf("error decoding stats response body: %v", err) - } - t.Logf("stats: %v", stats) - if stats.Gets == 0 { - t.Errorf("expected non-zero gets in session stats") - } - if stats.Puts == 0 && tc.wantWrite { - t.Errorf("expected non-zero puts in session stats") - } - }) - } -} - func TestJWTClaimFlag(t *testing.T) { for name, tc := range map[string]struct { input []string diff --git a/gocache/gocached/gocached.go b/gocache/gocached/gocached.go new file mode 100644 index 0000000..20bcfe7 --- /dev/null +++ b/gocache/gocached/gocached.go @@ -0,0 +1,1453 @@ +// Copyright (c) Tailscale Inc & AUTHORS +// SPDX-License-Identifier: BSD-3-Clause + +// The gocached package provides an HTTP server daemon that go-cacher can hit. +// It does cache tiering and evicts old large things from disk, and can fetch +// metadata and object contents from peer cache servers. +// +// It uses sqlite (the pure Go modernc.org/sqlite driver) to store metadata and +// indexes. +// +/* + +It speaks the same protocol as go-cacher-server, but requires +the "Want-Object: 1" header variant on the GET request. + + GET /action/ + Want-Object: 1 + + 200 OK + Content-Type: application/octet-stream + Content-Length: 1234 + Go-Output-Id: xxxxxxxxxxx + + + +And to insert an object: + + PUT // + Content-Length: 1234 + + + +*/ +package gocached + +import ( + "cmp" + "context" + "crypto/rand" + "crypto/sha256" + "database/sql" + "encoding/hex" + "encoding/json" + "errors" + "expvar" + "fmt" + "io" + "log" + "maps" + "math" + "net/http" + "net/http/pprof" + "os" + "path/filepath" + "reflect" + "slices" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/prometheus/client_golang/prometheus" + "github.com/prometheus/client_golang/prometheus/collectors" + "github.com/prometheus/client_golang/prometheus/promhttp" + dto "github.com/prometheus/client_model/go" + ijwt "github.com/tailscale/tb/gocache/internal/jwt" + _ "modernc.org/sqlite" +) + +const ( + // smallObjectSize is the maximum size of an object that we store inline in the + // database, rather than on disk. Empirically, about half of objects are 1KB or + // smaller. + smallObjectSize = 1 << 10 + + // tokenPrefix is the prefix for all gocached access tokens. + tokenPrefix = "gocached-token-" + + // gocachedAudience is the audience we require JWTs to have. Could be + // configurable in future, but for now just needs to be specific to gocached. + gocachedAudience = "gocached" +) + +const schemaVersion = 3 + +const schema = ` +PRAGMA journal_mode=WAL; +CREATE TABLE IF NOT EXISTS Actions ( + NamespaceID INTEGER NOT NULL, -- 0 for global trusted namespace + ActionID TEXT NOT NULL, + BlobID INTEGER NOT NULL, + AltOutputID TEXT NOT NULL DEFAULT '', -- if non-empty, the alternate object ID to use for this action; NULL means the blob's sha256 + CreateTime INTEGER NOT NULL, -- unix sec when inserted (locally or on a peer) + AccessTime INTEGER NOT NULL, -- unix sec of last access + + PRIMARY KEY (NamespaceID, ActionID), + + CHECK (ActionID = lower(ActionID)), + CHECK (ActionID GLOB '[0-9a-f]*'), + CHECK (CreateTime >= 0), + CHECK (AccessTime >= 0) +) STRICT; + +CREATE INDEX IF NOT EXISTS idx_actions_access ON Actions(AccessTime); +CREATE INDEX IF NOT EXISTS idx_actions_blobid ON Actions(BlobID); + +CREATE TABLE IF NOT EXISTS Blobs ( + BlobID INTEGER PRIMARY KEY AUTOINCREMENT, + SHA256 TEXT NOT NULL, + BlobSize INTEGER NOT NULL, -- size in bytes, either inline or on disk + SmallData BLOB, -- NULL if stored on disk + + CHECK (SmalLData IS NULL OR length(SmallData) = BlobSize) +) STRICT; + +CREATE UNIQUE INDEX IF NOT EXISTS idx_blobs_sha256 ON Blobs(SHA256); + +CREATE TABLE IF NOT EXISTS Namespaces ( + NamespaceID INTEGER PRIMARY KEY AUTOINCREMENT, + Namespace TEXT NOT NULL UNIQUE CHECK (Namespace = lower(Namespace)) +) STRICT; +` + +func openDB(dbDir string) (*sql.DB, error) { + dbPath := filepath.Join(dbDir, fmt.Sprintf("gocached-v%d.db", schemaVersion)) + db, err := sql.Open("sqlite", "file:"+dbPath+"?_pragma=busy_timeout(5000)") + if err != nil { + return nil, err + } + db.SetMaxOpenConns(4) + db.SetMaxIdleConns(4) + db.SetConnMaxLifetime(0) // no limit + if _, err := db.Exec(schema); err != nil { + return nil, err + } + return db, nil +} + +// start initializes the server, including defaults and background goroutines. +func (srv *Server) start() error { + if srv.dir == "" { + d, err := os.UserCacheDir() + if err != nil { + return fmt.Errorf("getting user cache dir: %w", err) + } + srv.dir = filepath.Join(d, "gocached") + srv.logf("Defaulting to cache dir %v ...", srv.dir) + } + if err := os.MkdirAll(srv.dir, 0750); err != nil { + return fmt.Errorf("creating cache dir: %w", err) + } + db, err := openDB(srv.dir) + if err != nil { + return fmt.Errorf("openDB: %w", err) + } + srv.db = db + + reg := prometheus.NewRegistry() + reg.MustRegister( + collectors.NewGoCollector(), + collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}), + collectors.NewBuildInfoCollector(), + ) + srv.registerMetrics(reg) + + srv.metricsHandler = promhttp.HandlerFor(reg, promhttp.HandlerOpts{ + ErrorLog: log.Default(), + }) + + srv.logf("gocached: scanning usage & cleaning as needed...") + us, err := srv.usageStats() + if err != nil { + return fmt.Errorf("getting usage stats: %w", err) + } + + srv.logf("gocached: current usage: %v of limit %v", us.All(), bytesFmt(srv.maxSize)) + if res, err := srv.cleanOldObjects(us); err != nil { + return fmt.Errorf("clean old objects: %w", err) + } else if res.Count > 0 { + srv.logf("gocached: cleaned %v", res) + } + + if srv.jwtIssuer != "" { + srv.jwtValidator = ijwt.NewJWTValidator(srv.logf, srv.jwtIssuer, gocachedAudience) + if err := srv.jwtValidator.RunUpdateJWKSLoop(srv.shutdownCtx); err != nil { + return fmt.Errorf("failed to fetch JWKS for JWT validator: %w", err) + } + + srv.logf("gocached: using JWT issuer %q with claims %v, global claims %v", srv.jwtIssuer, srv.jwtClaims, srv.globalJWTClaims) + + go srv.runCleanSessionsLoop() + } + + go srv.runCleanLoop() + + return nil +} + +func (srv *Server) registerMetrics(reg *prometheus.Registry) { + rv := reflect.ValueOf(&srv.m).Elem() + t := reflect.TypeOf(&srv.m).Elem() + for i := 0; i < t.NumField(); i++ { + sf := t.Field(i) + if sf.Type == reflect.TypeFor[expvar.Int]() { + expvarInt := rv.Field(i).Addr().Interface().(*expvar.Int) + typ := sf.Tag.Get("type") + name := sf.Tag.Get("name") + if typ == "" { + panic("missing type tag for " + sf.Name) + } + if name == "" { + panic("missing name tag for " + sf.Name) + } + help := sf.Tag.Get("help") + metricName := "gocached_" + name + + if tag := sf.Tag.Get("type"); tag != "" { + if tag == "gauge" { + reg.MustRegister(singleMetricCollector{&expvarGaugeMetric{ + desc: prometheus.NewDesc(metricName, help, nil, nil), + v: expvarInt, + }}) + } else if tag == "counter" { + reg.MustRegister(singleMetricCollector{&expvarCounterMetric{ + desc: prometheus.NewDesc(metricName, help, nil, nil), + v: expvarInt, + }}) + } + } + } + } +} + +// ServerOption configures a gocached Server. +type ServerOption func(*Server) + +// WithShutdownCtx sets the context used to signal server shutdown. Defaults to +// context.Background(). +func WithShutdownCtx(ctx context.Context) ServerOption { + return func(srv *Server) { + srv.shutdownCtx = ctx + } +} + +// WithDir sets the directory where the server stores its data. Defaults to the +// OS user cache directory under $XDG_CACHE_HOME/gocached or equivalent. +func WithDir(dir string) ServerOption { + return func(srv *Server) { + srv.dir = dir + } +} + +// WithVerbose enables verbose logging for the server. Defaults to false. +func WithVerbose(verbose bool) ServerOption { + return func(srv *Server) { + srv.verbose = verbose + } +} + +type logf func(format string, args ...any) + +// WithLogf sets a custom logging function for the server. Defaults to +// [log.Printf]. +func WithLogf(logf logf) ServerOption { + return func(srv *Server) { + srv.logf = logf + } +} + +// WithMaxSize sets the maximum size of the cache in bytes. Defaults to 0, which +// means no limit. +func WithMaxSize(maxSize int64) ServerOption { + return func(srv *Server) { + srv.maxSize = maxSize + } +} + +// WithMaxAge sets the maximum age of objects in the cache. Objects older than +// this duration will be cleaned periodically. Defaults to 0, which means no +// limit. +func WithMaxAge(maxAge time.Duration) ServerOption { + return func(srv *Server) { + srv.maxAge = maxAge + } +} + +// WithJWTAuth enables JWT-based authentication for the server. The issuer must +// be a reachable HTTP(S) server that serves its JWKS via a URL discoverable at +// /.well-known/openid-configuration, and any JWT presented to the server must +// exactly match the provided claims to start a session. No requests are allowed +// without authentication if JWT auth is enabled. +func WithJWTAuth(issuer string, claims map[string]string) ServerOption { + return func(srv *Server) { + srv.jwtIssuer = issuer + srv.jwtClaims = claims + } +} + +// WithGlobalNamespaceJWTClaims sets additional claims that a JWT must have to +// write to the cache's global namespace. It should be a superset of the claims +// provided to [WithJWTAuth]. +func WithGlobalNamespaceJWTClaims(claims map[string]string) ServerOption { + return func(srv *Server) { + srv.globalJWTClaims = claims + } +} + +// NewServer creates and starts a new gocached [Server] that is ready to serve +// requests. It defaults to requiring no authentication and storing its data in +// the OS user cache directory under $XDG_CACHE_HOME/gocached or equivalent. +func NewServer(opts ...ServerOption) (*Server, error) { + srv := &Server{ + shutdownCtx: context.Background(), + logf: log.Printf, + sessions: make(map[string]*sessionData), + clock: time.Now, + } + for _, opt := range opts { + opt(srv) + } + + err := srv.start() + if err != nil { + return nil, err + } + + return srv, nil +} + +// Server implements a gocached server. Use [NewServer] to create and start a +// valid instance. +type Server struct { + db *sql.DB + dir string // for SQLite DB + large blobs + verbose bool + logf logf + clock func() time.Time // if non-nil, alternate time.Now for testing + metricsHandler http.Handler + maxSize int64 // maximum size of the cache in bytes; 0 means no limit + maxAge time.Duration // maximum age of objects; 0 means no limit + shutdownCtx context.Context + + jwtValidator *ijwt.Validator // nil unless jwtIssuer is set + jwtIssuer string // issuer URL for JWTs + jwtClaims map[string]string // claims required for any JWT to start a session + globalJWTClaims map[string]string // additional claims required to write to global namespace + + sessionsMu sync.RWMutex // guards sessions + sessions map[string]*sessionData // maps access token -> session data. + + // sqliteWriteMu serializes access to SQLite. In theory the SQLite driver + // should serialize access with our 5000ms busy timeout, but empirically we + // sometimes seen DB busy errors. Just serialize it explicitly out of + // laziness for now. + sqliteWriteMu sync.Mutex + + lastUsage atomic.Pointer[usageStats] + + // Metrics. Exported fields for reflection, but within a private struct + // field to control the gocached Server API surface. + m struct { + ActiveGets expvar.Int `type:"gauge" name:"active_gets" help:"currently pending get requests; should usually be zero"` + ActivePuts expvar.Int `type:"gauge" name:"active_puts" help:"currently pending put requests; should usually be zero"` + Gets expvar.Int `type:"counter" name:"gets" help:"total number of gocache get requests"` // gets = getHits + getErrs + implicit misses + GetBytes expvar.Int `type:"counter" name:"get_bytes" help:"total bytes fetched from gocache gets that were cache hits"` + GetHits expvar.Int `type:"counter" name:"get_hits" help:"total number of successful gocache get requests"` + GetAccessBumps expvar.Int `type:"counter" name:"get_access_bumps" help:"number of times a get request updated the access time of object"` + GetHitsInline expvar.Int `type:"counter" name:"get_hits_inline" help:"cache hits served from inline database storage (small objects)"` + GetErrs expvar.Int `type:"counter" name:"get_errs" help:"number of gocache get request errors"` + Puts expvar.Int `type:"counter" name:"puts" help:"total number of gocache put requests"` + PutErrs expvar.Int `type:"counter" name:"put_errs" help:"number of gocache put request errors"` + PutsDup expvar.Int `type:"counter" name:"puts_dup" help:"total number of gocache put requests that are duplicates of a mapping we already had"` + PutsBytes expvar.Int `type:"counter" name:"put_bytes" help:"total bytes added from gocache puts"` + PutsInline expvar.Int `type:"counter" name:"put_inline" help:"subset of gocached_puts that were stored inline (small objects)"` + BlobCount expvar.Int `type:"gauge" name:"blob_count" help:"number of blobs currently stored in the cache"` + BlobBytes expvar.Int `type:"gauge" name:"blob_bytes" help:"sum of blob sizes currently stored in the cache"` + EvictedBlobs expvar.Int `type:"counter" name:"evicted_blobs" help:"number of blobs evicted from the cache"` + EvictedBytes expvar.Int `type:"counter" name:"evicted_bytes" help:"number of bytes evicted from the cache"` + Sessions expvar.Int `type:"gauge" name:"sessions" help:"number of active authenticated sessions"` + Auths expvar.Int `type:"counter" name:"auth_attempts" help:"number of successful token exchanges"` + AuthErrs expvar.Int `type:"counter" name:"auth_errs" help:"number of failed token exchanges"` + } +} + +// sessionData corresponds to a specific access token, and is only used if JWT +// auth is enabled. +type sessionData struct { + expiry time.Time // Session valid until. + globalNSWrite bool // Whether this session can write to the cache's global namespace. + claims map[string]any // Claims from the JWT used to create this session, stored for debug. + + mu sync.Mutex // Guards stats. + stats stats +} + +// stats holds per-request or per-session stats which get rolled up into server +// stats. See [server] struct for detailed definitions. +type stats struct { + LastUsed time.Time // Only applies to session stats. Last time the access token for this session was used. + Gets int64 + GetBytes int64 + GetHits int64 + GetAccessBumps int64 + GetHitsInline int64 + GetNanos int64 + GetErrs int64 + Puts int64 + PutErrs int64 + PutsDup int64 + PutsBytes int64 + PutsInline int64 + PutsNanos int64 +} + +func (srv *Server) now() time.Time { + return srv.clock() +} + +// ServeHTTPDebug serves debug HTTP endpoints. It is unauthenticated, so should +// only be used on a separate debug listener. +func (srv *Server) ServeHTTPDebug(w http.ResponseWriter, r *http.Request) { + if srv.verbose { + srv.logf("ServeHTTPDebug: %s %s", r.Method, r.RequestURI) + } + switch { + case r.URL.Path == "/": + w.Header().Set("Content-Type", "text/html; charset=utf-8") + io.WriteString(w, "

gocached

") + io.WriteString(w, "

This is a shared Go build cache server, hit by GOCACHEPROG clients.

") + io.WriteString(w, "

See /usage for usage stats.

") + io.WriteString(w, "

See /sessions for session data

") + io.WriteString(w, "

See /metrics for Prometheus metrics.

") + io.WriteString(w, "

See /debug/pprof/ for pprof

") + io.WriteString(w, "

See /debug/pprof/goroutine?debug=2 - full goroutines

") + case r.URL.Path == "/usage": + srv.serveUsage(w, r) + case r.URL.Path == "/sessions": + srv.serveSessions(w, r) + case r.URL.Path == "/metrics": + srv.metricsHandler.ServeHTTP(w, r) + case strings.HasPrefix(r.URL.Path, "/debug/pprof/profile"): + pprof.Profile(w, r) + case strings.HasPrefix(r.URL.Path, "/debug/pprof/cmdline"): + pprof.Cmdline(w, r) + case strings.HasPrefix(r.URL.Path, "/debug/pprof/symbol"): + pprof.Symbol(w, r) + case strings.HasPrefix(r.URL.Path, "/debug/pprof/trace"): + pprof.Trace(w, r) + case strings.HasPrefix(r.URL.Path, "/debug/pprof/"): + pprof.Index(w, r) + default: + http.Error(w, "not found", http.StatusNotFound) + } +} + +// ServeHTTP implements gocached's API via [http.Handler]. +func (srv *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if srv.verbose { + srv.logf("ServeHTTP: %s %s", r.Method, r.RequestURI) + } + + var sessionData *sessionData // remains nil for unauthenticated requests. + reqStats := &stats{} + defer func() { + // Call inside func to capture maybe-updated sessionData pointer. + srv.processRequestStats(reqStats, sessionData) + }() + + // Handle session auth first if enabled. + if srv.jwtValidator != nil { + // If JWT auth enabled, this is the only unauthenticated (non-debug) endpoint. + if r.Method == "POST" && r.URL.Path == "/auth/exchange-token" { + srv.handleTokenExchange(w, r) + return + } + + // Check for session data and error if none. + token := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ") + if !strings.HasPrefix(token, tokenPrefix) { + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + + var ok bool + sessionData, ok = srv.getSessionData(token) + if !ok || srv.now().After(sessionData.expiry) { + if srv.verbose { + reason := fmt.Sprintf("exists: %v", ok) + if sessionData != nil { + reason += fmt.Sprintf(", expiry: %v", sessionData.expiry) + } + srv.logf("unauthorized; %s", reason) + } + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + } + + if r.Method == "PUT" { + if sessionData != nil && !sessionData.globalNSWrite { + // TODO(tomhjp): support per-namespace writes. + http.Error(w, "forbidden", http.StatusForbidden) + return + } + srv.handlePut(w, r, reqStats) + return + } + if r.Method != "GET" && r.Method != "HEAD" { + http.Error(w, "bad method", http.StatusBadRequest) + return + } + if strings.HasPrefix(r.URL.Path, "/action/") { + srv.handleGetAction(w, r, reqStats) + return + } + if sessionData != nil && r.URL.Path == "/session/stats" { + srv.handleSessionStats(w, sessionData) + return + } + http.Error(w, "not found", http.StatusNotFound) +} + +func (srv *Server) getSessionData(token string) (*sessionData, bool) { + srv.sessionsMu.Lock() + defer srv.sessionsMu.Unlock() + sessionData, ok := srv.sessions[token] + return sessionData, ok +} + +func (srv *Server) addSessionData(token string, sessionData *sessionData) { + srv.sessionsMu.Lock() + defer srv.sessionsMu.Unlock() + srv.sessions[token] = sessionData + srv.m.Sessions.Add(1) +} + +func (srv *Server) processRequestStats(req *stats, sessionData *sessionData) { + srv.m.Gets.Add(req.Gets) + srv.m.GetBytes.Add(req.GetBytes) + srv.m.GetHits.Add(req.GetHits) + srv.m.GetAccessBumps.Add(req.GetAccessBumps) + srv.m.GetHitsInline.Add(req.GetHitsInline) + srv.m.GetErrs.Add(req.GetErrs) + srv.m.Puts.Add(req.Puts) + srv.m.PutErrs.Add(req.PutErrs) + srv.m.PutsDup.Add(req.PutsDup) + srv.m.PutsBytes.Add(req.PutsBytes) + srv.m.PutsInline.Add(req.PutsInline) + + if sessionData != nil { + sessionData.mu.Lock() + defer sessionData.mu.Unlock() + + sessionData.stats.LastUsed = srv.now().UTC() + sessionData.stats.Gets += req.Gets + sessionData.stats.GetBytes += req.GetBytes + sessionData.stats.GetHits += req.GetHits + sessionData.stats.GetAccessBumps += req.GetAccessBumps + sessionData.stats.GetHitsInline += req.GetHitsInline + sessionData.stats.GetErrs += req.GetErrs + sessionData.stats.GetNanos += req.GetNanos + sessionData.stats.Puts += req.Puts + sessionData.stats.PutErrs += req.PutErrs + sessionData.stats.PutsDup += req.PutsDup + sessionData.stats.PutsBytes += req.PutsBytes + sessionData.stats.PutsInline += req.PutsInline + sessionData.stats.PutsNanos += req.PutsNanos + } +} + +func getHexSuffix(r *http.Request, prefix string) (hexSuffix string, ok bool) { + hexSuffix, _ = strings.CutPrefix(r.RequestURI, prefix) + if !validHex(hexSuffix) { + return "", false + } + return hexSuffix, true +} + +func validHex(x string) bool { + if len(x) < 4 || len(x) > 1000 || len(x)%2 == 1 { + return false + } + for i := range x { + b := x[i] + if b >= '0' && b <= '9' || b >= 'a' && b <= 'f' { + continue + } + return false + } + return true +} + +// relAtimeSeconds is how old an access time needs to be before +// we do a DB write to update it. +const relAtimeSeconds = 60 * 60 * 24 // 1 day + +func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats *stats) { + srv.m.ActiveGets.Add(1) + defer srv.m.ActiveGets.Add(-1) + + start := srv.now() + defer func() { + stats.GetNanos += srv.now().Sub(start).Nanoseconds() + }() + stats.Gets++ + ctx := r.Context() + + httpErr := func(msg string, code int) { + http.Error(w, msg, code) + stats.GetErrs++ + } + + actionID, ok := getHexSuffix(r, "/action/") + if !ok { + httpErr("bad request", http.StatusBadRequest) + return + } + if r.Header.Get("Want-Object") != "1" { + httpErr("bad request: missing Want-Object header", http.StatusBadRequest) + return + } + + var sha256hex string + var size int64 + var smallData sql.NullString + var altObjectID string + var accessTime int64 + namespaceID := 0 // global for now; TODO(bradfitz): support namespaces + err := srv.db.QueryRow( + "SELECT b.SHA256, b.BlobSize, b.SmallData, a.AltOutputID, a.AccessTime FROM Actions a, Blobs b WHERE a.NameSpaceID = ? AND a.ActionID = ? AND a.BlobID = b.BlobID", + namespaceID, actionID).Scan( + &sha256hex, &size, &smallData, &altObjectID, &accessTime) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + http.Error(w, "not found", http.StatusNotFound) + return + } + srv.logf("QueryRow error: %v", err) + httpErr("QueryRow error", http.StatusInternalServerError) + return + } + + // If it's been more than a day since the last access, update the access time. + // This is similar to the Linux "relatime" behavior. + now := srv.now().Unix() + if accessTime < now-relAtimeSeconds { + // TODO(bradfitz): do this async? not worth blocking the caller. + // But we need a mechanism for tests to wait on async work. + srv.sqliteWriteMu.Lock() + _, err := srv.db.Exec("UPDATE Actions SET AccessTime = ? WHERE ActionID = ?", now, actionID) + srv.sqliteWriteMu.Unlock() + if err != nil { + srv.logf("Update AccessTime error: %v", err) + httpErr("internal server error", http.StatusInternalServerError) + return + } + stats.GetAccessBumps++ + } + + stats.GetHits++ + + outputID := cmp.Or(altObjectID, sha256hex) + + w.Header().Set("Content-Type", "application/octet-stream") + w.Header().Set("Content-Length", fmt.Sprint(size)) + w.Header().Set("Go-Output-Id", outputID) + + if r.Method == "HEAD" || size == 0 { + return + } + + if smallData.Valid { + // For small outputs stored inline in the database, we can return them directly. + stats.GetHitsInline++ + stats.GetBytes += size + io.WriteString(w, smallData.String) + return + } + + // Otherwise, for large objects that we know about, we can try to get them + // from our local disk or a peer. + + rc, err := srv.getObjectFromDiskOrPeer(ctx, sha256hex) + if err != nil { + srv.logf("Get object error: %v", actionID, err) + httpErr("Get object error", http.StatusInternalServerError) + return + } + if rc == nil { + // Our database suggested we should've had this object, + // but maybe somebody delete it by hand from the filesystem. + // Just treat it as a cache miss. The background cleanup + // will eventually remove the Action row from the DB + // after identifying it as a dangling reference. + http.Error(w, "not found", http.StatusNotFound) + return + } + stats.GetBytes += size + defer rc.Close() + io.Copy(w, rc) +} + +// getObjectFromDiskOrPeer retrieves the object for the given actionID, either +// from disk or a peer. This is used after a local DB lookup discovers the content +// exists but is not stored in SQLite. +// +// It returns (nil, nil) on miss. +func (srv *Server) getObjectFromDiskOrPeer(_ context.Context, sha256hex string) (rc io.ReadCloser, err error) { + if len(sha256hex) != sha256.Size*2 { + return nil, fmt.Errorf("invalid sha256hex %q", sha256hex) + } + diskPath := filepath.Join(srv.dir, sha256hex[:2], sha256hex) + f, err := os.Open(diskPath) + if err != nil { + if os.IsNotExist(err) { + // TODO(bradfitz): search peers, S3, etc. + // For now, just return nil, nil on miss. + return nil, nil + } + return nil, err + } + return f, nil +} + +func (s *Server) handlePut(w http.ResponseWriter, r *http.Request, stats *stats) { + s.m.ActivePuts.Add(1) + defer s.m.ActivePuts.Add(-1) + + start := s.now() + defer func() { + stats.PutsNanos += s.now().Sub(start).Nanoseconds() + }() + + if r.Method != "PUT" { + http.Error(w, "bad method", http.StatusMethodNotAllowed) + return + } + actionID, outputID, ok := strings.Cut(r.RequestURI[len("/"):], "/") + if !ok || !validHex(actionID) || !validHex(outputID) { + http.Error(w, "bad URI", http.StatusBadRequest) + return + } + if r.ContentLength == -1 { + http.Error(w, "missing Content-Length", http.StatusBadRequest) + return + } + + hasher := sha256.New() + hashingBody := io.TeeReader(r.Body, hasher) + + var smallData []byte + if r.ContentLength <= smallObjectSize { + // Store small objects inline in the database. + var err error + smallData, err = io.ReadAll(hashingBody) + if err != nil { + s.logf("Read content error: %v", err) + http.Error(w, "Read content error", http.StatusInternalServerError) + return + } + if int64(len(smallData)) != r.ContentLength { + // This check is redundant with net/http's validation, but + // for extra clarity. + http.Error(w, "bad content length", http.StatusInternalServerError) + return + } + } else { + // For larger objects, we store them on disk. + if err := s.writeDiskBlob(r.ContentLength, hashingBody); err != nil { + s.logf("Write disk blob error: %v", err) + http.Error(w, "Write disk blob error", http.StatusInternalServerError) + return + } + } + + sha256hex := fmt.Sprintf("%x", hasher.Sum(nil)) + blobSize := r.ContentLength + + s.sqliteWriteMu.Lock() + defer s.sqliteWriteMu.Unlock() + + var blobID int64 + err := s.db.QueryRow(`INSERT INTO Blobs (SHA256, BlobSize, SmallData) + VALUES (?, ?, ?) + ON CONFLICT(SHA256) DO UPDATE SET SHA256=excluded.SHA256 + RETURNING BlobID; +`, sha256hex, blobSize, smallData).Scan(&blobID) + if err != nil { + s.logf("Blobs insert error: %v", err) + stats.PutErrs++ + http.Error(w, "Blobs insert error", http.StatusInternalServerError) + return + } + + // Insert or update the action in the database. + nowUnix := s.now().Unix() + altObjectID := "" + namespace := 0 // global for now; TODO(bradfitz): support namespaces + if sha256hex != outputID { + altObjectID = outputID + } + res, err := s.db.Exec(`INSERT OR IGNORE INTO Actions (NamespaceID, ActionID, BlobID, AltOutputID, CreateTime, AccessTime) + VALUES (?, ?, ?, ?, ?, ?)`, + namespace, + actionID, + blobID, + altObjectID, + nowUnix, + nowUnix, + ) + if err != nil { + s.logf("Actions insert error: %v", err) + stats.PutErrs++ + http.Error(w, "Actions insert error", http.StatusInternalServerError) + return + } + + affected, err := res.RowsAffected() + if err != nil { + s.logf("Actions rows affected error: %v", err) + stats.PutErrs++ + http.Error(w, "Actions rows affected error", http.StatusInternalServerError) + return + } + + if affected == 0 { + stats.PutsDup++ + } + + stats.Puts++ + stats.PutsBytes += r.ContentLength + if smallData != nil { + stats.PutsInline++ + } + + w.WriteHeader(http.StatusNoContent) +} + +// handleTokenExchange handles POST /auth/exchange-token requests to exchange +// a JWT for an access token. Each access token represents a session that will +// last for one hour and have cache stats associated with it. +func (srv *Server) handleTokenExchange(w http.ResponseWriter, r *http.Request) { + var req struct { + JWT string `json:"jwt"` + } + // JWTs are often sent in HTTP headers, so 4KiB should ~always be enough. + if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, 4<<10)).Decode(&req); err != nil { + srv.m.AuthErrs.Add(1) + http.Error(w, "bad request: "+err.Error(), http.StatusBadRequest) + return + } + + jwtClaims, err := srv.jwtValidator.Validate(r.Context(), req.JWT) + if err != nil { + srv.m.AuthErrs.Add(1) + if srv.verbose { + srv.logf("token exchange: JWT validation error: %v", err) + } + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + + globalNSWrite, err := srv.evaluateClaims(jwtClaims) + if err != nil { + srv.m.AuthErrs.Add(1) + if srv.verbose { + srv.logf("token exchange: %v", err) + } + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + + const ttl = time.Hour + // 52 base32 characters, 256 bits of entropy. + accessToken := tokenPrefix + strings.ToLower(rand.Text()+rand.Text()) + srv.addSessionData(accessToken, &sessionData{ + expiry: srv.now().UTC().Add(ttl), + globalNSWrite: globalNSWrite, + claims: jwtClaims, + }) + + resp := map[string]any{ + "access_token": accessToken, + "token_type": "Bearer", + "expires_in": ttl.Seconds(), + } + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(resp); err != nil { + srv.m.AuthErrs.Add(1) + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + + srv.m.Auths.Add(1) +} + +func (srv *Server) evaluateClaims(claims map[string]any) (globalNSWrite bool, _ error) { + if missing := findMissingClaims(srv.jwtClaims, claims); len(missing) > 0 { + return false, fmt.Errorf("got claims %v; missing required claims: %v", claims, missing) + } + + if missing := findMissingClaims(srv.globalJWTClaims, claims); len(missing) == 0 { + return true, nil + } else if srv.verbose { + srv.logf("token exchange: missing global namespace write claims: %v", missing) + } + + return false, nil +} + +func findMissingClaims(wantClaims map[string]string, gotClaims map[string]any) map[string]any { + if wantClaims == nil { + return nil + } + + missing := make(map[string]any) + for k, want := range wantClaims { + if got, ok := gotClaims[k]; !ok || got != want { + missing[k] = want + } + } + return missing +} + +func (srv *Server) handleSessionStats(w http.ResponseWriter, sessionData *sessionData) { + w.Header().Set("Content-Type", "application/json") + if err := json.NewEncoder(w).Encode(sessionData.stats); err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } +} + +func (s *Server) sha256Filepath(hash [sha256.Size]byte) string { + hex := fmt.Sprintf("%x", hash) + return filepath.Join(s.dir, hex[:2], hex) +} + +func (s *Server) writeDiskBlob(size int64, r io.Reader) (err error) { + nowUnix := s.now().Unix() + tf, err := os.CreateTemp(s.dir, fmt.Sprintf("upload-%d-*", nowUnix)) + if err != nil { + return err + } + defer func() { + if err == nil { + return + } + tf.Close() + os.Remove(tf.Name()) + }() + hasher := sha256.New() + n, err := io.Copy(tf, io.LimitReader(io.TeeReader(r, hasher), size+1)) + if err != nil { + return err + } + if n != size { + return fmt.Errorf("wrote %d bytes; wanted %d", n, size) + } + if err := tf.Close(); err != nil { + return err + } + var hash [sha256.Size]byte + hasher.Sum(hash[:0]) + + target := s.sha256Filepath(hash) + if err := os.MkdirAll(filepath.Dir(target), 0750); err != nil { + return err + } + return os.Rename(tf.Name(), target) +} + +type countAndSize struct { + Count int64 // number of actions + Size int64 // total size of all actions' blobs (even if shared by other actions) +} + +func (cs countAndSize) String() string { + if cs.Count == 0 { + return "0 objects, 0 bytes" + } + return fmt.Sprintf("%d objects, %s", cs.Count, bytesFmt(cs.Size)) +} + +type usageStats struct { + // ActionsLE is a histogram of the actions in the DB by their access time. + // + // The key is a Prometheus-style histogram "less than" value. That is, if + // there are map keys for 24h and 48h, the latter includes the sum of the + // 24h values as well. + // + // The map keys are day-granularity, as the access time is only updated once + // it's over a day old. + // + // So the map keys are 24h, 48h, 96h, 168h (7d), 336h (14d), 720h + // (30d), and 2160h (90d) and math.MaxInt64 for infinity. + ActionsLE map[time.Duration]countAndSize + + // MissingBlobRows is the number of rows in the Actions table that + // reference a BlobID that doesn't exist in the Blobs table. + // This should always be zero in a healthy system. + MissingBlobRows int +} + +func (us *usageStats) All() countAndSize { return us.ActionsLE[math.MaxInt64] } + +const day = 24 * time.Hour + +var standardDurs = []time.Duration{ + 1 * day, + 2 * day, + 4 * day, + 7 * day, + 14 * day, + 30 * day, + 90 * day, + math.MaxInt64, +} + +func (s *Server) usageStats() (_ *usageStats, err error) { + defer func() { + if err != nil { + s.logf("usageStats error: %v", err) + } + }() + + st := &usageStats{ + ActionsLE: make(map[time.Duration]countAndSize), + } + + // Build the durations to use for the histogram. + // The math.MaxInt64 value is always included. + // If s.maxAge is set, we ignore sizes above that, except + // for the math.MaxInt64 value. + var durs []time.Duration + if s.maxAge == 0 { + durs = standardDurs + } else { + durs = make([]time.Duration, 0, len(standardDurs)+1) + durs = append(durs, s.maxAge) + for _, d := range standardDurs { + if d < s.maxAge || d == math.MaxInt64 { + durs = append(durs, d) + } + } + slices.Sort(durs) + } + + now := s.now().Unix() + rows, err := s.db.Query( + "SELECT a.BlobID, a.AccessTime, b.BlobSize FROM Actions a LEFT JOIN Blobs b ON a.BlobID = b.BlobID") + if err != nil { + return nil, fmt.Errorf("query Actions: %w", err) + } + var blobID int64 + var accessTime int64 + var blobSize sql.NullInt64 + for rows.Next() { + if err := rows.Scan(&blobID, &accessTime, &blobSize); err != nil { + return nil, fmt.Errorf("rows.Scan: %w", err) + } + if !blobSize.Valid { + st.MissingBlobRows++ + continue + } + + dur := time.Duration(now-accessTime) * time.Second + if dur < 0 { + dur = 0 + } + for _, d := range durs { + if dur < d { + was := st.ActionsLE[d] + was.Count++ + was.Size += blobSize.Int64 + st.ActionsLE[d] = was + } + } + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("rows.Next: %w", err) + } + + s.lastUsage.Store(st) + all := st.All() + s.m.BlobCount.Set(all.Count) + s.m.BlobBytes.Set(all.Size) + return st, nil +} + +type cleanCandidate struct { + BlobID int64 + Age time.Duration + BlobSize int64 // size of the blob, in bytes +} + +func (s *Server) cleanCandidates(olderThan time.Duration, limit int64) ([]cleanCandidate, error) { + now := s.now() + nowUnix := now.Unix() + cutoff := now.Add(-olderThan).Unix() + + rows, err := s.db.Query(` + SELECT b.BlobID, MAX(a.AccessTime), b.BlobSize + FROM Blobs b LEFT JOIN Actions a ON b.BlobID = a.BlobID + GROUP BY b.BlobID + HAVING MAX(a.AccessTime) <= ? + ORDER BY MAX(a.AccessTime) + LIMIT ?`, cutoff, limit) + if err != nil { + return nil, fmt.Errorf("query clean candidates: %w", err) + } + defer rows.Close() + + var candidates []cleanCandidate + var accessTime int64 + for rows.Next() { + var c cleanCandidate + if err := rows.Scan(&c.BlobID, &accessTime, &c.BlobSize); err != nil { + return nil, fmt.Errorf("rows.Scan: %w", err) + } + c.Age = time.Duration(nowUnix-accessTime) * time.Second + candidates = append(candidates, c) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("rows.Next: %w", err) + } + + return candidates, nil +} + +func (srv *Server) deleteBlobs(blobIDs ...int64) error { + srv.sqliteWriteMu.Lock() + defer srv.sqliteWriteMu.Unlock() + + tx, err := srv.db.Begin() + if err != nil { + return fmt.Errorf("delete blob Begin: %w", err) + } + defer tx.Rollback() + + var sumBytes int64 + for _, blobID := range blobIDs { + var sha256Hex string + var blobSize int64 + if err := tx.QueryRow("SELECT SHA256, BlobSize FROM Blobs WHERE BlobID = ?", blobID).Scan(&sha256Hex, &blobSize); err != nil && !errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("querying blob SHA256: %w", err) + } + sumBytes += blobSize + if _, err := tx.Exec("DELETE FROM Blobs WHERE BlobID = ?", blobID); err != nil { + return fmt.Errorf("deleting blob: %w", err) + } + if _, err := tx.Exec("DELETE FROM Actions WHERE BlobID = ?", blobID); err != nil { + return fmt.Errorf("deleting actions: %w", err) + } + var hash [sha256.Size]byte + if _, err := hex.Decode(hash[:], []byte(sha256Hex)); err == nil { + if err := os.Remove(srv.sha256Filepath(hash)); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("removing disk file: %w", err) + } + } + } + if err := tx.Commit(); err != nil { + return err + } + + srv.m.EvictedBlobs.Add(int64(len(blobIDs))) + srv.m.EvictedBytes.Add(sumBytes) + + return nil +} + +func (srv *Server) cleanOldObjects(us *usageStats) (countAndSize, error) { + var zero countAndSize + var ret countAndSize + + all := us.ActionsLE[math.MaxInt64] + if srv.verbose { + srv.logf("current usage stats: %v", all) + last := all + for _, d := range slices.Sorted(maps.Keys(us.ActionsLE)) { + if d == math.MaxInt64 { + continue // skip infinity + } + c := us.ActionsLE[d] + srv.logf(" <=%v: %v", durFmt(d), c) + if last == c { + break + } + last = c + } + } + + // First clean things that are just too old. + if srv.maxAge > 0 { + if toDelete := all.Count - us.ActionsLE[srv.maxAge].Count; toDelete > 0 { + srv.logf("Cleaning %d objects older than %v ...", toDelete, durFmt(srv.maxAge)) + candidates, err := srv.cleanCandidates(srv.maxAge, toDelete+1) + if err != nil { + return zero, fmt.Errorf("getting clean candidates: %v", err) + } + blobIDs := make([]int64, 0, len(candidates)) + var sumSize int64 + for _, c := range candidates { + blobIDs = append(blobIDs, c.BlobID) + sumSize += c.BlobSize + } + if err := srv.deleteBlobs(blobIDs...); err != nil { + return zero, fmt.Errorf("deleting old blobs: %v", err) + } + all.Count -= int64(len(candidates)) + all.Size -= sumSize + ret.Count += int64(len(candidates)) + ret.Size += sumSize + } + } + + for srv.maxSize > 0 && all.Size > srv.maxSize { + toClean := all.Size - srv.maxSize + if srv.verbose { + srv.logf("need to clean %v to get under max size of %v ...", + bytesFmt(toClean), bytesFmt(srv.maxSize)) + } + + var batchBytes int64 + var blobIDs []int64 + candidates, err := srv.cleanCandidates(0, 10000) + if err != nil { + return zero, fmt.Errorf("getting clean candidates: %v", err) + } + for _, c := range candidates { + blobIDs = append(blobIDs, c.BlobID) + batchBytes += c.BlobSize + if batchBytes >= toClean { + break + } + } + if err := srv.deleteBlobs(blobIDs...); err != nil { + return zero, fmt.Errorf("deleting old blobs: %v", err) + } + + ret.Count += int64(len(blobIDs)) + ret.Size += batchBytes + all.Count -= int64(len(blobIDs)) + all.Size -= batchBytes + + if len(blobIDs) == len(candidates) { + // We didn't find enough candidates to delete. + // Just stop here. + srv.logf("[unexpected] didn't find enough candidates to delete") + break + } + } + + return ret, nil +} + +func (srv *Server) runCleanLoop() { + for { + select { + case <-srv.shutdownCtx.Done(): + return + case <-time.After(5 * time.Minute): + } + + us, err := srv.usageStats() + if err != nil { + srv.logf("error getting usage stats: %v", err) + continue + } + + res, err := srv.cleanOldObjects(us) + if err != nil { + srv.logf("error cleaning old objects: %v", err) + continue + } + if res.Count > 0 { + srv.logf("cleaned %v", res) + srv.usageStats() // for side effect of updating lastUsage + } + + } +} + +func (srv *Server) runCleanSessionsLoop() { + for { + select { + case <-srv.shutdownCtx.Done(): + return + case <-time.After(time.Hour): + } + + srv.sessionsMu.Lock() + count := len(srv.sessions) + var deleted int + for token, metadata := range srv.sessions { + if time.Now().After(metadata.expiry) { + delete(srv.sessions, token) + deleted++ + srv.m.Sessions.Add(-1) + } + } + srv.sessionsMu.Unlock() + if count > 0 { + srv.logf("cleaned up %d/%d access tokens", deleted, count) + } + } +} + +func durFmt(d time.Duration) string { + days := int(d.Hours() / 24) + if days > 0 { + return fmt.Sprintf("%dd", days) + } + return d.String() +} + +func bytesFmt(n int64) string { + if n >= 1<<30 { + return fmt.Sprintf("%.1f GiB", float64(n)/(1<<30)) + } + if n >= 1<<20 { + return fmt.Sprintf("%.1f MiB", float64(n)/(1<<20)) + } + if n >= 1<<10 { + return fmt.Sprintf("%.1f KiB", float64(n)/(1<<10)) + } + return fmt.Sprintf("%d bytes", n) +} + +func (srv *Server) serveUsage(w http.ResponseWriter, r *http.Request) { + if r.Method == "POST" { + // For side effect of updating lastUsage. + _, err := srv.usageStats() + if err != nil { + http.Error(w, "error getting usage stats: "+err.Error(), http.StatusInternalServerError) + return + } + } + + us := srv.lastUsage.Load() + if us == nil { + http.Error(w, "no usage stats available", http.StatusInternalServerError) + return + } + + // Print out an HTML table of the usage stats, sorted by age. + w.Header().Set("Content-Type", "text/html; charset=utf-8") + fmt.Fprintf(w, "

gocached usage stats

\n") + fmt.Fprintf(w, "

Current usage: %v of limit %v

\n", + us.All(), bytesFmt(srv.maxSize)) + + fmt.Fprintf(w, "\n") + fmt.Fprintf(w, "\n") + for _, d := range slices.Sorted(maps.Keys(us.ActionsLE)) { + var title string + if d == math.MaxInt64 { + title = "all" + } else { + title = "<= " + durFmt(d) + } + c := us.ActionsLE[d] + fmt.Fprintf(w, "\n", + title, c.Count, bytesFmt(c.Size)) + } + fmt.Fprintf(w, "
AgeCountSize
%s%d%s
\n") +} + +func (srv *Server) serveSessions(w http.ResponseWriter, r *http.Request) { + if r.Method != "GET" { + http.Error(w, "bad method", http.StatusMethodNotAllowed) + return + } + + srv.sessionsMu.RLock() + // Make a copy of all session data (excluding the mutex). + sessions := make([]*sessionData, 0, len(srv.sessions)) + for _, v := range srv.sessions { + v.mu.Lock() + sessions = append(sessions, &sessionData{ + expiry: v.expiry, + globalNSWrite: v.globalNSWrite, + claims: v.claims, + stats: v.stats, + }) + v.mu.Unlock() + } + srv.sessionsMu.RUnlock() + + w.Header().Set("Content-Type", "text/html; charset=utf-8") + fmt.Fprintf(w, "

gocached sessions

\n") + fmt.Fprintf(w, "

JWT issuer: %s

\n", srv.jwtIssuer) + fmt.Fprintf(w, "

JWT claims required: %v

\n", srv.jwtClaims) + fmt.Fprintf(w, "

JWT global write claims required: %v

\n", srv.globalJWTClaims) + fmt.Fprintf(w, "

Number of sessions: %d

\n", len(sessions)) + + fmt.Fprintf(w, "\n") + fmt.Fprintf(w, "\n") + slices.SortFunc(sessions, func(a, b *sessionData) int { + return a.stats.LastUsed.Compare(b.stats.LastUsed) + }) + for _, d := range slices.Backward(sessions) { + lastUsed := "never" + if !d.stats.LastUsed.IsZero() { + lastUsed = durFmt(time.Since(d.stats.LastUsed)) + " ago" + } + statsJSON, _ := json.MarshalIndent(d.stats, "", " ") + claimsJSON, _ := json.MarshalIndent(d.claims, "", " ") + fmt.Fprintf(w, "\n", + lastUsed, d.expiry.Format(time.RFC3339), d.globalNSWrite, statsJSON, claimsJSON) + } + fmt.Fprintf(w, "
Last usedExpiry timeGlobal writeStatsClaims
%s%s%v
%s
%s
\n") +} + +// expvarCounterMetric is a Prometheus counter metric backed by an expvar.Int. +type expvarCounterMetric struct { + desc *prometheus.Desc + v *expvar.Int +} + +var _ prometheus.Metric = (*expvarCounterMetric)(nil) + +func (m *expvarCounterMetric) Desc() *prometheus.Desc { return m.desc } + +func (m *expvarCounterMetric) Write(out *dto.Metric) error { + val := float64(m.v.Value()) + out.Counter = &dto.Counter{Value: &val} + return nil +} + +// expvarGaugeMetric is a Prometheus gauge metric backed by an expvar.Int. +type expvarGaugeMetric struct { + desc *prometheus.Desc + v *expvar.Int +} + +var _ prometheus.Metric = (*expvarGaugeMetric)(nil) + +func (m *expvarGaugeMetric) Desc() *prometheus.Desc { return m.desc } + +func (m *expvarGaugeMetric) Write(out *dto.Metric) error { + val := float64(m.v.Value()) + out.Gauge = &dto.Gauge{Value: &val} + return nil +} + +// singleMetricCollector is a Prometheus collector that collects a single metric. +type singleMetricCollector struct { + metric prometheus.Metric +} + +var _ prometheus.Collector = singleMetricCollector{} + +func (c singleMetricCollector) Describe(ch chan<- *prometheus.Desc) { ch <- c.metric.Desc() } +func (c singleMetricCollector) Collect(ch chan<- prometheus.Metric) { ch <- c.metric } diff --git a/gocache/gocached/gocached_test.go b/gocache/gocached/gocached_test.go new file mode 100644 index 0000000..a609393 --- /dev/null +++ b/gocache/gocached/gocached_test.go @@ -0,0 +1,699 @@ +// Copyright (c) Tailscale Inc & AUTHORS +// SPDX-License-Identifier: BSD-3-Clause + +package gocached + +import ( + "bytes" + "context" + "crypto" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "encoding/json" + "expvar" + "fmt" + "io" + "math" + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "slices" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/go-jose/go-jose/v4" + "github.com/golang-jwt/jwt/v5" + "github.com/google/go-cmp/cmp" + "github.com/tailscale/tb/gocache/cachers" +) + +// sha256OfEmpty is the SHA-256 hash of an empty string, used as a well-known +// value in SQLite to store bytes, as it's common. +const sha256OfEmpty = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" + +type tester struct { + t testing.TB + srv *Server + hs *httptest.Server + + timeMu sync.Mutex + curTime time.Time +} + +func (t *tester) Logf(format string, args ...any) { + t.t.Logf(format, args...) +} + +func (t *tester) now() time.Time { + t.timeMu.Lock() + defer t.timeMu.Unlock() + return t.curTime +} + +func (t *tester) advanceClock(d time.Duration) { + t.timeMu.Lock() + defer t.timeMu.Unlock() + t.curTime = t.curTime.Add(d) +} + +func (t *tester) mkClient() *cachers.HTTPClient { + clientCacheDir := t.t.TempDir() + return &cachers.HTTPClient{ + BaseURL: t.hs.URL, + Disk: &cachers.DiskCache{ + Dir: clientCacheDir, + Logf: func(format string, args ...any) { + t.Logf("client-disk: "+format, args...) + }, + }, + } +} + +func (st *tester) usageStats() *usageStats { + st.t.Helper() + stats, err := st.srv.usageStats() + if err != nil { + st.t.Fatalf("usageStats: %v", err) + } + return stats +} + +func (st *tester) cleanOldObjects() countAndSize { + st.t.Helper() + stats, err := st.srv.cleanOldObjects(st.usageStats()) + if err != nil { + st.t.Fatalf("cleanOldObjects: %v", err) + } + return stats +} + +func (st *tester) diskFiles() []string { + st.t.Helper() + var ret []string + + err := filepath.Walk(st.srv.dir, func(path string, fi os.FileInfo, err error) error { + if err != nil { + return err + } + if !fi.Mode().IsRegular() || strings.HasPrefix(fi.Name(), ".") || strings.HasPrefix(fi.Name(), "gocached") { + return nil + } + ret = append(ret, fi.Name()) + return nil + }) + if err != nil { + st.t.Fatalf("Walk: %v", err) + } + slices.Sort(ret) + return ret +} + +// wantMetric is a helper to check an expvar.Int metric and reset it +// for future tests. +func (st *tester) wantMetric(m *expvar.Int, want int64) { + st.t.Helper() + if got := m.Value(); got != want { + st.t.Errorf("metric = %d, want %d", got, want) + } + m.Set(0) +} + +func (st *tester) wantPut(c *cachers.HTTPClient, actionID, outputID string, val string) { + ctx := context.Background() + st.t.Helper() + clientDiskPath, err := c.Put(ctx, actionID, outputID, int64(len(val)), strings.NewReader(val)) + if err != nil { + st.t.Fatalf("Put: %v", err) + } + if clientDiskPath == "" { + st.t.Fatal("Put returned empty disk path") + } + st.wantMetric(&st.srv.m.Puts, 1) + wrote, err := os.ReadFile(clientDiskPath) + if err != nil { + st.t.Fatalf("ReadFile: %v", err) + } + if string(wrote) != val { + st.t.Errorf("ReadFile got %q, want %q", wrote, val) + } +} + +func (st *tester) wantGet(c *cachers.HTTPClient, actionID, outputID, wantVal string) { + ctx := context.Background() + st.t.Helper() + gotOutputID, diskPath, err := c.Get(ctx, actionID) + if err != nil { + st.t.Fatalf("Get: %v", err) + } + if gotOutputID != outputID { + st.t.Errorf("Get got outputID %q, want %q", gotOutputID, outputID) + } + if diskPath == "" { + st.t.Fatal("Get returned empty disk path") + } + wrote, err := os.ReadFile(diskPath) + if err != nil { + st.t.Fatalf("ReadFile: %v", err) + } + if string(wrote) != wantVal { + st.t.Errorf("ReadFile got %q, want %q", wrote, wantVal) + } +} + +func (st *tester) wantGetMiss(c *cachers.HTTPClient, actionID string) { + ctx := context.Background() + st.t.Helper() + gotOutputID, diskPath, err := c.Get(ctx, actionID) + if err != nil { + st.t.Fatalf("Get: %v", err) + } + if gotOutputID != "" { + st.t.Errorf("Get got outputID %q; want empty", gotOutputID) + } + if diskPath != "" { + st.t.Fatalf("Get returned disk path %q; want empty", diskPath) + } +} + +func withClock(clk func() time.Time) ServerOption { + return func(cfg *Server) { + cfg.clock = clk + } +} + +func newServerTester(t testing.TB, extraOpts ...ServerOption) *tester { + st := &tester{ + t: t, + curTime: time.Unix(1234, 0), + } + + opts := []ServerOption{ + WithDir(t.TempDir()), + WithLogf(t.Logf), + WithVerbose(true), + withClock(st.now), + } + srv, err := NewServer(append(opts, extraOpts...)...) + if err != nil { + t.Fatalf("starting gocached: %v", err) + } + st.srv = srv + + st.hs = httptest.NewServer(st.srv) + t.Cleanup(st.hs.Close) + + return st +} + +type jwtFunc func(claims jwt.MapClaims, signingKey *ecdsa.PrivateKey) string + +// startOIDCServer starts a mock OIDC server that gocached can use for JWT auth. +// The provided publicKey is what JWT signatures will be validated against. +func startOIDCServer(t *testing.T, publicKey crypto.PublicKey) (iss string, jwtFunc jwtFunc) { + t.Helper() + mux := http.NewServeMux() + srv := httptest.NewServer(mux) + t.Cleanup(srv.Close) + + issuer := fmt.Sprintf("http://%s", srv.Listener.Addr().String()) + + mux.HandleFunc("/.well-known/openid-configuration", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "issuer": issuer, + "jwks_uri": fmt.Sprintf("%s/jwks", issuer), + }) + }) + mux.HandleFunc("/jwks", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(map[string]any{ + "keys": []jose.JSONWebKey{ + { + Key: publicKey, + KeyID: "test-key", + Algorithm: "ES256", + Use: "sig", + }, + }, + }) + }) + + return issuer, func(claims jwt.MapClaims, signingKey *ecdsa.PrivateKey) string { + t.Helper() + unsignedTk := &jwt.Token{ + Header: map[string]any{ + "typ": "JWT", + "alg": jwt.SigningMethodES256.Alg(), + "kid": "test-key", + }, + Claims: claims, + Method: jwt.SigningMethodES256, + } + tk, err := unsignedTk.SignedString(signingKey) + if err != nil { + t.Fatalf("error signing token: %v", err) + } + + return tk + } +} + +func TestServer(t *testing.T) { + st := newServerTester(t) + + ctx := context.Background() + + // Make two clients (imagine: two different builder VMs) + c1 := st.mkClient() + c2 := st.mkClient() + + const testActionID = "0001" + const testActionIDMiss = "0002" // this one doesn't exist + const testActionIDBig = "0bbb" // non-inline object + const testActionIDEmpty = "0000" + const testOutputID = "9900" + const testOutputIDBig = "9bbb" + const testOutputIDEmpty = sha256OfEmpty + const testObjectValue = "test data" + testObjectValueBig := strings.Repeat("x", smallObjectSize+1) + + // Populate from the first client. + st.wantPut(c1, testActionID, testOutputID, testObjectValue) + st.wantPut(c1, testActionIDBig, testOutputIDBig, testObjectValueBig) + st.wantPut(c1, testActionIDEmpty, testOutputIDEmpty, "") + + // Read from the second client. + st.wantGet(c2, testActionID, testOutputID, testObjectValue) + st.wantGet(c2, testActionIDBig, testOutputIDBig, testObjectValueBig) + st.wantGet(c2, testActionIDEmpty, testOutputIDEmpty, "") + + // Check metrics + st.wantMetric(&st.srv.m.Gets, 3) + st.wantMetric(&st.srv.m.GetHits, 3) + st.wantMetric(&st.srv.m.GetHitsInline, 1) + + // Do the same get again from the same client. This shouldn't hit the network. + st.wantGet(c2, testActionID, testOutputID, testObjectValue) + st.wantMetric(&st.srv.m.Gets, 0) + + // Cache miss. This should hit the network and fail. + if _, _, err := c2.Get(ctx, testActionIDMiss); err != nil { + t.Fatalf("miss Get: %v", err) + } + st.wantMetric(&st.srv.m.Gets, 1) + st.wantMetric(&st.srv.m.GetHits, 0) + + // Check that access time gets updated. + // Do it from a fresh client without a disk cache. + st.wantMetric(&st.srv.m.GetAccessBumps, 0) + st.advanceClock(relAtimeSeconds * 2 * time.Second) // advance clock by 2 days + c3 := st.mkClient() + st.wantGet(c3, testActionID, testOutputID, testObjectValue) + st.wantMetric(&st.srv.m.GetAccessBumps, 1) + + // Get usage stats. + stats, err := st.srv.usageStats() + if err != nil { + t.Fatalf("usageStats: %v", err) + } + want := &usageStats{ + MissingBlobRows: 0, + ActionsLE: map[time.Duration]countAndSize{ + 24 * time.Hour: {Count: 1, Size: 9}, + 48 * time.Hour: {Count: 1, Size: 9}, + 96 * time.Hour: {Count: 3, Size: 1034}, + 168 * time.Hour: {Count: 3, Size: 1034}, + 336 * time.Hour: {Count: 3, Size: 1034}, + 720 * time.Hour: {Count: 3, Size: 1034}, + 2160 * time.Hour: {Count: 3, Size: 1034}, + math.MaxInt64: {Count: 3, Size: 1034}, + }, + } + if diff := cmp.Diff(stats, want); diff != "" { + t.Errorf("usageStats mismatch (-got +want):\n%s", diff) + } + + st.advanceClock(relAtimeSeconds * 2 * time.Second) // advance clock by 2 days +} + +func TestCleanCandidates(t *testing.T) { + st := newServerTester(t) + + // Populate some data. + c1 := st.mkClient() + st.wantPut(c1, "0001", "9901", "1") + st.advanceClock(24 * time.Hour) + st.wantPut(c1, "0002", "9902", "22") + st.advanceClock(24 * time.Hour) + st.wantPut(c1, "0003", "9903", "333") + st.advanceClock(24 * time.Hour) + st.wantPut(c1, "0004", "9904", strings.Repeat("x", smallObjectSize+1)) + + const day = 24 * time.Hour + + tests := []struct { + maxAge time.Duration + limit int64 + want []cleanCandidate + }{ + { + maxAge: 0, + limit: 100, + want: []cleanCandidate{ + {BlobID: 1, Age: 3 * day, BlobSize: 1}, + {BlobID: 2, Age: 2 * day, BlobSize: 2}, + {BlobID: 3, Age: 1 * day, BlobSize: 3}, + {BlobID: 4, Age: 0, BlobSize: smallObjectSize + 1}, + }, + }, + { + maxAge: 25 * time.Hour, + limit: 100, + want: []cleanCandidate{ + {BlobID: 1, Age: 3 * day, BlobSize: 1}, + {BlobID: 2, Age: 2 * day, BlobSize: 2}, + }, + }, + { + maxAge: 0, + limit: 2, + want: []cleanCandidate{ + {BlobID: 1, Age: 3 * day, BlobSize: 1}, + {BlobID: 2, Age: 2 * day, BlobSize: 2}, + }, + }, + } + for _, tt := range tests { + t.Run(fmt.Sprintf("maxAge=%v,limit=%d", tt.maxAge, tt.limit), func(t *testing.T) { + candidates, err := st.srv.cleanCandidates(tt.maxAge, tt.limit) + if err != nil { + t.Fatal(err) + } + if diff := cmp.Diff(candidates, tt.want); diff != "" { + t.Errorf("cleanCandidates mismatch (-got +want):\n%s", diff) + + } + }) + } +} + +func TestCleanOldObjectsByAge(t *testing.T) { + st := newServerTester(t) + st.srv.maxAge = 24 * time.Hour + + // Populate some data. + c1 := st.mkClient() + st.wantPut(c1, "0001", "9901", strings.Repeat("x", smallObjectSize+1)) + st.advanceClock(25 * time.Hour) + st.wantPut(c1, "0002", "9902", strings.Repeat("x", smallObjectSize+2)) + st.wantPut(c1, "0003", "9903", "small") + smallLen := int64(len("small")) + + st1 := st.usageStats() + if all, want := st1.All(), (countAndSize{Count: 3, Size: smallObjectSize*2 + 3 + smallLen}); all != want { + t.Errorf("usageStats: %v; want %v", all, want) + } + if got, want := st.diskFiles(), []string{"333092a3daf718ed8f38a94e302df139edd4e3b5da4239a497995683942cf28c", "c6d8e9905300876046729949cc95c2385221270d389176f7234fe7ac00c4e430"}; !slices.Equal(got, want) { + t.Errorf("diskFiles: %v; want %v", got, want) + } + + clean1 := st.cleanOldObjects() + if clean1.Count != 1 || clean1.Size != smallObjectSize+1 { + t.Errorf("cleanOldObjects got %v, want {Count: 1, Size: %d}", clean1, smallObjectSize+1) + } + clean2 := st.cleanOldObjects() + if clean2.Count != 0 || clean2.Size != 0 { + t.Errorf("cleanOldObjects got %v, want {Count: 0, Size: 0}", clean2) + } + + st2 := st.usageStats() + if all, want := st2.All(), (countAndSize{Count: 2, Size: smallObjectSize + 2 + smallLen}); all != want { + t.Errorf("usageStats after clean: %v; want %v", all, want) + } + if got, want := st.diskFiles(), []string{"333092a3daf718ed8f38a94e302df139edd4e3b5da4239a497995683942cf28c"}; !slices.Equal(got, want) { + t.Errorf("diskFiles after clean: %v; want %v", got, want) + } +} + +func TestCleanOldObjectsBySize(t *testing.T) { + st := newServerTester(t) + + // Populate some data. + c1 := st.mkClient() + st.wantPut(c1, "0001", "9901", "1") + st.advanceClock(time.Second) + st.wantPut(c1, "0002", "9902", "22") + st.advanceClock(time.Second) + st.wantPut(c1, "0003", "9903", "333") + st.advanceClock(time.Second) + st.wantPut(c1, "0004", "9904", "4444") + st.advanceClock(time.Second) + + st1 := st.usageStats() + if all, want := st1.All(), (countAndSize{Count: 4, Size: 10}); all != want { + t.Errorf("usageStats: %v; want %v", all, want) + } + + clean1 := st.cleanOldObjects() + if clean1.Count != 0 || clean1.Size != 0 { + t.Errorf("cleanOldObjects got %v, want no clean", clean1) + } + + st.srv.maxSize = 8 // the only way get to 8 or under is by deleting "1" and "22" (3 bytes) + + if got, want := st.cleanOldObjects(), (countAndSize{Count: 2, Size: 3}); got != want { + t.Errorf("cleanOldObjects got %v, want %v", got, want) + } + if got, want := st.usageStats().All(), (countAndSize{Count: 2, Size: 7}); got != want { + t.Errorf("usageStats: %v; want %v", got, want) + } +} + +func TestClientConnReuse(t *testing.T) { + st := newServerTester(t) + + var numDials atomic.Int32 + tr := http.DefaultTransport.(*http.Transport).Clone() + tr.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) { + num := numDials.Add(1) + t.Logf("DialContext #%d for %s %s", num, network, addr) + var std net.Dialer + return std.DialContext(ctx, network, addr) + } + t.Cleanup(func() { tr.CloseIdleConnections() }) + + c1 := st.mkClient() + c1.HTTPClient = &http.Client{Transport: tr} + const missAction = "0001" + st.wantGetMiss(c1, missAction) + st.wantGetMiss(c1, missAction) + st.wantGetMiss(c1, missAction) + st.wantPut(c1, "0001", "9901", "1") + st.wantGet(c1, "0001", "9901", "1") + if got := numDials.Load(); got != 1 { + t.Errorf("numDials = %d; want 1", got) + } +} + +func TestExchangeToken(t *testing.T) { + // Generate private keys outside of the loop for speed. + privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("error generating OIDC server private key: %v", err) + } + otherPrivateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("error generating OIDC server private key: %v", err) + } + wantClaims := map[string]string{ + "sub": "user123", + } + wantGlobalClaims := map[string]string{ + "sub": "user123", + "ref": "refs/heads/main", + } + + for name, tc := range map[string]struct { + mutateClaims func(jwt.MapClaims) + signingKey *ecdsa.PrivateKey + wantStatusCode int + wantWrite bool + }{ + // Base case: no mutation. + "valid_read": { + wantStatusCode: http.StatusOK, + wantWrite: false, + }, + // Additional claim needed for write scope. + "valid_write": { + mutateClaims: func(cl jwt.MapClaims) { + cl["ref"] = "refs/heads/main" + }, + wantStatusCode: http.StatusOK, + wantWrite: true, + }, + // Every other test makes one mutation from the base case that should cause failure. + "missing_sub": { + mutateClaims: func(cl jwt.MapClaims) { + delete(cl, "sub") + }, + wantStatusCode: http.StatusUnauthorized, + }, + "invalid_sub": { + mutateClaims: func(cl jwt.MapClaims) { + cl["sub"] = "user456" + }, + wantStatusCode: http.StatusUnauthorized, + }, + "invalid_iss": { + mutateClaims: func(cl jwt.MapClaims) { + cl["iss"] = "invalid_issuer" + }, + wantStatusCode: http.StatusUnauthorized, + }, + "invalid_aud": { + mutateClaims: func(cl jwt.MapClaims) { + cl["aud"] = "invalid_audience" + }, + wantStatusCode: http.StatusUnauthorized, + }, + "not_yet_valid": { + mutateClaims: func(cl jwt.MapClaims) { + cl["nbf"] = jwt.NewNumericDate(time.Now().Add(10 * time.Minute)) + }, + wantStatusCode: http.StatusUnauthorized, + }, + "expired": { + mutateClaims: func(cl jwt.MapClaims) { + cl["exp"] = jwt.NewNumericDate(time.Now().Add(-time.Minute)) + }, + wantStatusCode: http.StatusUnauthorized, + }, + "invalid_signature": { + signingKey: otherPrivateKey, + wantStatusCode: http.StatusUnauthorized, + }, + } { + t.Run(name, func(t *testing.T) { + issuer, createJWT := startOIDCServer(t, privateKey.Public()) + st := newServerTester(t, + WithJWTAuth(issuer, wantClaims), + WithGlobalNamespaceJWTClaims(wantGlobalClaims), + ) + + // Generate JWT. + tokenClaims := jwt.MapClaims{ + "sub": "user123", + "num": 42, + "iss": issuer, + "aud": gocachedAudience, + "nbf": jwt.NewNumericDate(time.Now().Add(-time.Minute)), + "exp": jwt.NewNumericDate(time.Now().Add(time.Hour)), + } + if tc.mutateClaims != nil { + tc.mutateClaims(tokenClaims) + } + signingKey := privateKey + if tc.signingKey != nil { + signingKey = tc.signingKey + } + body, err := json.Marshal(map[string]any{ + "jwt": createJWT(tokenClaims, signingKey), + }) + if err != nil { + t.Fatalf("error marshaling request body: %v", err) + } + + // Exchange JWT for access token. + req, err := http.NewRequest("POST", st.hs.URL+"/auth/exchange-token", bytes.NewReader(body)) + if err != nil { + t.Fatalf("error creating request: %v", err) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("error making request: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != tc.wantStatusCode { + t.Fatalf("unexpected status code: want %d, got %d", tc.wantStatusCode, resp.StatusCode) + } + body, err = io.ReadAll(resp.Body) + if err != nil { + t.Fatalf("error reading response body: %v", err) + } + + if tc.wantStatusCode != http.StatusOK { + if string(body) != "unauthorized\n" { + t.Fatalf("unexpected error body: %s", string(body)) + } + + // No access token to do further checks with; test finished. + return + } + + // Check returned access token. + var d struct { + AccessToken string `json:"access_token"` + } + if err := json.Unmarshal(body, &d); err != nil { + t.Fatalf("error decoding response body: %v", err) + } + if d.AccessToken == "" { + t.Fatalf("expected access_token in response, got %s", string(body)) + } + + cl := st.mkClient() + if _, _, err := cl.Get(t.Context(), "abc123"); err == nil { + t.Fatalf("Get without access token succeeded unexpectedly") + } + + cl.AccessToken = d.AccessToken + st.wantGetMiss(cl, "abc123") + + if tc.wantWrite { + st.wantPut(cl, "abc123", "def456", "data789") + st.wantGet(cl, "abc123", "def456", "data789") + } else { + if _, err := cl.Put(t.Context(), "abc123", "def456", 0, nil); err == nil { + t.Fatalf("Put without write scope succeeded unexpectedly") + } + } + + // Check session stats. + reqStats, err := http.NewRequest("GET", st.hs.URL+"/session/stats", nil) + if err != nil { + t.Fatalf("error creating stats request: %v", err) + } + reqStats.Header.Set("Authorization", "Bearer "+d.AccessToken) + respStats, err := http.DefaultClient.Do(reqStats) + if err != nil { + t.Fatalf("error making stats request: %v", err) + } + defer respStats.Body.Close() + if respStats.StatusCode != http.StatusOK { + t.Fatalf("unexpected stats status code: want %d, got %d", http.StatusOK, respStats.StatusCode) + } + bodyStats, err := io.ReadAll(respStats.Body) + if err != nil { + t.Fatalf("error reading stats response body: %v", err) + } + var stats stats + if err := json.Unmarshal(bodyStats, &stats); err != nil { + t.Fatalf("error decoding stats response body: %v", err) + } + t.Logf("stats: %v", stats) + if stats.Gets == 0 { + t.Errorf("expected non-zero gets in session stats") + } + if stats.Puts == 0 && tc.wantWrite { + t.Errorf("expected non-zero puts in session stats") + } + }) + } +} diff --git a/cmd/gocached/internal/jwt/jwt.go b/gocache/internal/jwt/jwt.go similarity index 94% rename from cmd/gocached/internal/jwt/jwt.go rename to gocache/internal/jwt/jwt.go index 1344b73..31f13a0 100644 --- a/cmd/gocached/internal/jwt/jwt.go +++ b/gocache/internal/jwt/jwt.go @@ -1,3 +1,6 @@ +// Copyright (c) Tailscale Inc & AUTHORS +// SPDX-License-Identifier: BSD-3-Clause + package jwt import ( @@ -5,7 +8,6 @@ import ( "encoding/json" "fmt" "io" - "log" "net/http" "net/url" "path" @@ -31,9 +33,9 @@ var ( // public signing keys via the path defined by [oidcConfigWellKnownPath], and // the audience should be a value specific to the trust boundary that gocached // resides within. -func NewJWTValidator(issuer, audience string) *Validator { +func NewJWTValidator(logf func(format string, args ...any), issuer, audience string) *Validator { return &Validator{ - Logf: log.Printf, + logf: logf, issuer: issuer, parser: jwt.NewParser( jwt.WithValidMethods(supportedAlgorithms), @@ -63,7 +65,7 @@ func (v *Validator) RunUpdateJWKSLoop(ctx context.Context) error { // Validator provides methods for validating JWTs. Use [NewJWTValidator] to // construct a working Validator. type Validator struct { - Logf func(format string, args ...any) + logf func(format string, args ...any) issuer string parser *jwt.Parser @@ -127,13 +129,13 @@ func (v *Validator) runUpdateJWKSLoop(ctx context.Context) { // Non-fatal; in practice, the most recent keys will normally // still be valid for a long time, but JWT validation will start // erroring more loudly than this if not. - v.Logf("jwt: failed to update JWKS: %v", err) + v.logf("jwt: failed to update JWKS: %v", err) } } } func (v *Validator) updateJWKS(ctx context.Context) error { - v.Logf("jwt: fetching JWKS from issuer %q", v.issuer) + v.logf("jwt: fetching JWKS from issuer %q", v.issuer) u, err := url.Parse(v.issuer) if err != nil { return fmt.Errorf("failed to parse issuer URL %q: %w", v.issuer, err) From 4d1305784fd91ad834aa2b798e262c6c73a1113c Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Mon, 1 Dec 2025 07:41:17 -0800 Subject: [PATCH 30/67] gocached: fix AccessTime bump query to use an index We weren't specifying the NamespaceID so the primary key wasn't being used, as the namespace ID is the prefix of the primary key. Updates tailscale/corp#34696 Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@80f0f91d98ec135fd44a55e36066145b169b121b --- gocache/gocached/gocached.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gocache/gocached/gocached.go b/gocache/gocached/gocached.go index 20bcfe7..6158022 100644 --- a/gocache/gocached/gocached.go +++ b/gocache/gocached/gocached.go @@ -647,7 +647,7 @@ func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats // TODO(bradfitz): do this async? not worth blocking the caller. // But we need a mechanism for tests to wait on async work. srv.sqliteWriteMu.Lock() - _, err := srv.db.Exec("UPDATE Actions SET AccessTime = ? WHERE ActionID = ?", now, actionID) + _, err := srv.db.Exec("UPDATE Actions SET AccessTime = ? WHERE NamespaceID = ? AND ActionID = ?", now, namespaceID, actionID) srv.sqliteWriteMu.Unlock() if err != nil { srv.logf("Update AccessTime error: %v", err) From 7fde5d9cf4dcaf7cb062bbe314b166a4dd499daf Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sun, 30 Nov 2025 19:53:06 -0800 Subject: [PATCH 31/67] gocached: batch access time updates, reducing SQLite write mutex contention Updates tailscale/corp#34696 Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@fef72b933348f22ba68130d931cac90802fa9d92 --- gocache/gocached/gocached.go | 154 +++++++++++++++++++++++++++++------ 1 file changed, 127 insertions(+), 27 deletions(-) diff --git a/gocache/gocached/gocached.go b/gocache/gocached/gocached.go index 6158022..29956ae 100644 --- a/gocache/gocached/gocached.go +++ b/gocache/gocached/gocached.go @@ -345,13 +345,17 @@ type Server struct { jwtClaims map[string]string // claims required for any JWT to start a session globalJWTClaims map[string]string // additional claims required to write to global namespace - sessionsMu sync.RWMutex // guards sessions - sessions map[string]*sessionData // maps access token -> session data. + mu sync.RWMutex // guards following fields in this block + sessions map[string]*sessionData // maps access token -> session data. + accessDirty map[actionKey]int64 // action -> accessTime + accessFlushTimer *time.Timer // nil if no flush is scheduled // sqliteWriteMu serializes access to SQLite. In theory the SQLite driver // should serialize access with our 5000ms busy timeout, but empirically we // sometimes seen DB busy errors. Just serialize it explicitly out of // laziness for now. + // + // Lock ordering: sqliteWriteMu before mu. sqliteWriteMu sync.Mutex lastUsage atomic.Pointer[usageStats] @@ -521,15 +525,15 @@ func (srv *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { } func (srv *Server) getSessionData(token string) (*sessionData, bool) { - srv.sessionsMu.Lock() - defer srv.sessionsMu.Unlock() + srv.mu.Lock() + defer srv.mu.Unlock() sessionData, ok := srv.sessions[token] return sessionData, ok } func (srv *Server) addSessionData(token string, sessionData *sessionData) { - srv.sessionsMu.Lock() - defer srv.sessionsMu.Unlock() + srv.mu.Lock() + defer srv.mu.Unlock() srv.sessions[token] = sessionData srv.m.Sessions.Add(1) } @@ -576,6 +580,13 @@ func getHexSuffix(r *http.Request, prefix string) (hexSuffix string, ok bool) { return hexSuffix, true } +// actionKey is the comparable value type for the (NamespaceID, ActionID) +// primary key tuple used in the SQLite Actions table. +type actionKey struct { + NamespaceID int // 0 for global + ActionID string +} + func validHex(x string) bool { if len(x) < 4 || len(x) > 1000 || len(x)%2 == 1 { return false @@ -625,10 +636,13 @@ func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats var smallData sql.NullString var altObjectID string var accessTime int64 - namespaceID := 0 // global for now; TODO(bradfitz): support namespaces + var actionKey = actionKey{ + NamespaceID: 0, // global for now; TODO(bradfitz): support namespac + ActionID: actionID, + } err := srv.db.QueryRow( "SELECT b.SHA256, b.BlobSize, b.SmallData, a.AltOutputID, a.AccessTime FROM Actions a, Blobs b WHERE a.NameSpaceID = ? AND a.ActionID = ? AND a.BlobID = b.BlobID", - namespaceID, actionID).Scan( + actionKey.NamespaceID, actionKey.ActionID).Scan( &sha256hex, &size, &smallData, &altObjectID, &accessTime) if err != nil { if errors.Is(err, sql.ErrNoRows) { @@ -640,20 +654,7 @@ func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats return } - // If it's been more than a day since the last access, update the access time. - // This is similar to the Linux "relatime" behavior. - now := srv.now().Unix() - if accessTime < now-relAtimeSeconds { - // TODO(bradfitz): do this async? not worth blocking the caller. - // But we need a mechanism for tests to wait on async work. - srv.sqliteWriteMu.Lock() - _, err := srv.db.Exec("UPDATE Actions SET AccessTime = ? WHERE NamespaceID = ? AND ActionID = ?", now, namespaceID, actionID) - srv.sqliteWriteMu.Unlock() - if err != nil { - srv.logf("Update AccessTime error: %v", err) - httpErr("internal server error", http.StatusInternalServerError) - return - } + if srv.maybeBumpAccessTime(actionKey, accessTime) { stats.GetAccessBumps++ } @@ -700,6 +701,96 @@ func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats io.Copy(w, rc) } +// maybeBumpAccessTime reports whether it enqueued an access time bump for the +// given actionKey, if the provided prior access time (unix seconds) is old +// enough to warrant an update. +func (srv *Server) maybeBumpAccessTime(actionKey actionKey, priorAccessTimeUnixSec int64) (didBump bool) { + now := srv.now().Unix() + if priorAccessTimeUnixSec > now-relAtimeSeconds { + return false + } + // If it's been more than a day since the last access, update the access time. + // This is similar to the Linux "relatime" behavior. + return srv.enqueueAccessTimeBump(actionKey) +} + +// accessBatchSizeSoftLimit is the size of the accessDirty map at which +// we stop kicking the timer can down the road and let the thing flush, +// considering it large enough. +const accessBatchSizeSoftLimit = 1000 + +func (srv *Server) enqueueAccessTimeBump(action actionKey) (changed bool) { + srv.mu.Lock() + defer srv.mu.Unlock() + if _, ok := srv.accessDirty[action]; ok { + // Another caller already bumped this from old (over realTimeSeconds) + // to recent, so let that caller win, even if it's a few seconds old. + // Those few seconds don't matter much for cleaning purposes, and it's + // better for stats to only have one caller report the bump. + return false + } + if srv.accessDirty == nil { + srv.accessDirty = make(map[actionKey]int64) + } + srv.accessDirty[action] = srv.now().Unix() + if srv.accessFlushTimer == nil { + srv.accessFlushTimer = time.AfterFunc(5*time.Second, srv.flushAccessTimeBumps) + } else if len(srv.accessDirty) < accessBatchSizeSoftLimit { + srv.accessFlushTimer.Reset(5 * time.Second) + } + return true +} + +// flushAccessTimeBumps writes any pending access time updates to the database. +// +// It doesn't return an error to it can be easily used as a time.AfterFunc +// callback. +func (srv *Server) flushAccessTimeBumps() { + srv.flushAccessTimeBumpsWithErr() +} + +// flushAccessTimeBumpsWithErr is the same as flushAccessTimeBumps but returns +// any error that's used only for the internal rescheduling defer logic. +func (srv *Server) flushAccessTimeBumpsWithErr() (ret error) { + srv.sqliteWriteMu.Lock() // before srv.mu + defer srv.sqliteWriteMu.Unlock() + + srv.mu.Lock() + defer srv.mu.Unlock() + + defer func() { + if ret != nil { + srv.logf("flushAccessTimeBumps error: %v", ret) + srv.accessFlushTimer = time.AfterFunc(5*time.Second, srv.flushAccessTimeBumps) + } + }() + + srv.accessFlushTimer = nil + + if len(srv.accessDirty) == 0 { + return nil + } + + tx, err := srv.db.Begin() + if err != nil { + return fmt.Errorf("begin Tx: %w", err) + } + defer tx.Rollback() + + for action, accessTime := range srv.accessDirty { + _, err := tx.Exec("UPDATE Actions SET AccessTime = ? WHERE NamespaceID = ? AND ActionID = ?", accessTime, action.NamespaceID, action.ActionID) + if err != nil { + return fmt.Errorf("updating access time for ns=%d,action=%q: %w", action.NamespaceID, action.ActionID, err) + } + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("commit: %w", err) + } + + srv.accessDirty = nil + return nil +} + // getObjectFromDiskOrPeer retrieves the object for the given actionID, either // from disk or a peer. This is used after a local DB lookup discovers the content // exists but is not stored in SQLite. @@ -1045,6 +1136,16 @@ func (s *Server) usageStats() (_ *usageStats, err error) { slices.Sort(durs) } + // Flush any pending access time bumps before computing usage stats. + s.mu.RLock() + shouldFlush := len(s.accessDirty) > 0 + s.mu.RUnlock() + if shouldFlush { + // This acquires sqliteWriteMu, so we avoid that lock if there's nothing + // to do. + s.flushAccessTimeBumps() + } + now := s.now().Unix() rows, err := s.db.Query( "SELECT a.BlobID, a.AccessTime, b.BlobSize FROM Actions a LEFT JOIN Blobs b ON a.BlobID = b.BlobID") @@ -1276,7 +1377,6 @@ func (srv *Server) runCleanLoop() { srv.logf("cleaned %v", res) srv.usageStats() // for side effect of updating lastUsage } - } } @@ -1288,7 +1388,7 @@ func (srv *Server) runCleanSessionsLoop() { case <-time.After(time.Hour): } - srv.sessionsMu.Lock() + srv.mu.Lock() count := len(srv.sessions) var deleted int for token, metadata := range srv.sessions { @@ -1298,7 +1398,7 @@ func (srv *Server) runCleanSessionsLoop() { srv.m.Sessions.Add(-1) } } - srv.sessionsMu.Unlock() + srv.mu.Unlock() if count > 0 { srv.logf("cleaned up %d/%d access tokens", deleted, count) } @@ -1370,7 +1470,7 @@ func (srv *Server) serveSessions(w http.ResponseWriter, r *http.Request) { return } - srv.sessionsMu.RLock() + srv.mu.RLock() // Make a copy of all session data (excluding the mutex). sessions := make([]*sessionData, 0, len(srv.sessions)) for _, v := range srv.sessions { @@ -1383,7 +1483,7 @@ func (srv *Server) serveSessions(w http.ResponseWriter, r *http.Request) { }) v.mu.Unlock() } - srv.sessionsMu.RUnlock() + srv.mu.RUnlock() w.Header().Set("Content-Type", "text/html; charset=utf-8") fmt.Fprintf(w, "

gocached sessions

\n") From 9402fa512870b4f10ef9792e459c39b77388dfb2 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Wed, 11 Feb 2026 16:20:12 +0000 Subject: [PATCH 32/67] cachers: make HTTP client support lz4 decompression Updates tailscale/corp#37104 Migrated-from: bradfitz/go-tool-cache@d16f1e0999cb7f788cb166868d40172efac9f14b --- go.mod | 1 + go.sum | 2 + gocache/cachers/http.go | 41 ++++++++-- gocache/cachers/http_test.go | 145 +++++++++++++++++++++++++++++++++++ 4 files changed, 183 insertions(+), 6 deletions(-) create mode 100644 gocache/cachers/http_test.go diff --git a/go.mod b/go.mod index 779fdf3..6c0f0b9 100644 --- a/go.mod +++ b/go.mod @@ -5,6 +5,7 @@ require ( github.com/go-jose/go-jose/v4 v4.1.3 github.com/golang-jwt/jwt/v5 v5.3.0 github.com/google/go-cmp v0.7.0 + github.com/pierrec/lz4/v4 v4.1.25 github.com/prometheus/client_golang v1.23.0 github.com/prometheus/client_model v0.6.2 modernc.org/sqlite v1.38.2 diff --git a/go.sum b/go.sum index fee6bb1..e87ef0d 100644 --- a/go.sum +++ b/go.sum @@ -26,6 +26,8 @@ github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4= github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/pierrec/lz4/v4 v4.1.25 h1:kocOqRffaIbU5djlIBr7Wh+cx82C0vtFb0fOurZHqD0= +github.com/pierrec/lz4/v4 v4.1.25/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/prometheus/client_golang v1.23.0 h1:ust4zpdl9r4trLY/gSjlm07PuiBq2ynaXXlptpfy8Uc= diff --git a/gocache/cachers/http.go b/gocache/cachers/http.go index 1b809e2..e3ecd40 100644 --- a/gocache/cachers/http.go +++ b/gocache/cachers/http.go @@ -8,6 +8,9 @@ import ( "io" "log" "net/http" + "strconv" + + "github.com/pierrec/lz4/v4" ) // ActionValue is the JSON value returned by the cacher server for an GET /action request. @@ -56,6 +59,28 @@ func tryReadErrorMessage(res *http.Response) []byte { return msg } +// responseBody returns the response body and the uncompressed content length. +// If the response has Content-Encoding: lz4, the body is wrapped with an lz4 +// decompressor and the uncompressed length is read from X-Uncompressed-Length. +// For uncompressed responses, Content-Length is used directly. +func responseBody(res *http.Response) (body io.Reader, uncompressedLength int64, err error) { + if res.Header.Get("Content-Encoding") == "lz4" { + sizeStr := res.Header.Get("X-Uncompressed-Length") + if sizeStr == "" { + return nil, 0, fmt.Errorf("lz4-compressed response missing X-Uncompressed-Length header") + } + size, err := strconv.ParseInt(sizeStr, 10, 64) + if err != nil { + return nil, 0, fmt.Errorf("invalid X-Uncompressed-Length %q: %v", sizeStr, err) + } + return lz4.NewReader(res.Body), size, nil + } + if res.ContentLength == -1 { + return nil, 0, fmt.Errorf("no Content-Length from server") + } + return res.Body, res.ContentLength, nil +} + func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPath string, err error) { outputID, diskPath, err = c.Disk.Get(ctx, actionID) if err == nil && outputID != "" { @@ -67,6 +92,8 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa req.Header.Set("Authorization", "Bearer "+c.AccessToken) } + req.Header.Set("Accept-Encoding", "lz4") + // Set a header to indicate we want the object and metadata in one response. // Prior to 2025-08-09, the protocol was two separate requests. Rather than // change this repo's protocol and potentially break existing clients, @@ -98,10 +125,11 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa if outputID == "" { return "", "", fmt.Errorf("missing Go-Output-Id header in response") } - if res.ContentLength == -1 { - return "", "", fmt.Errorf("no Content-Length from server") + body, size, err := responseBody(res) + if err != nil { + return "", "", err } - diskPath, err = c.Disk.Put(ctx, actionID, outputID, res.ContentLength, res.Body) + diskPath, err = c.Disk.Put(ctx, actionID, outputID, size, body) case "application/json": // old two-hop protocol var av ActionValue @@ -116,6 +144,7 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa putBody = bytes.NewReader(nil) } else { req, _ = http.NewRequestWithContext(ctx, "GET", c.BaseURL+"/output/"+outputID, nil) + req.Header.Set("Accept-Encoding", "lz4") res, err = c.httpClient().Do(req) if err != nil { return "", "", err @@ -130,10 +159,10 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa log.Printf("error GET /output/%s: %v, %s", outputID, res.Status, msg) return "", "", fmt.Errorf("unexpected GET /output/%s status %v", outputID, res.Status) } - if res.ContentLength == -1 { - return "", "", fmt.Errorf("no Content-Length from server") + putBody, _, err = responseBody(res) + if err != nil { + return "", "", err } - putBody = res.Body } diskPath, err = c.Disk.Put(ctx, actionID, outputID, av.Size, putBody) } diff --git a/gocache/cachers/http_test.go b/gocache/cachers/http_test.go new file mode 100644 index 0000000..043e630 --- /dev/null +++ b/gocache/cachers/http_test.go @@ -0,0 +1,145 @@ +package cachers + +import ( + "bytes" + "context" + "encoding/binary" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "strconv" + "testing" + + "github.com/pierrec/lz4/v4" +) + +func lz4Compress(t *testing.T, data []byte) []byte { + t.Helper() + var buf bytes.Buffer + w := lz4.NewWriter(&buf) + if err := w.Apply(lz4.SizeOption(uint64(len(data)))); err != nil { + t.Fatal(err) + } + if _, err := w.Write(data); err != nil { + t.Fatal(err) + } + if err := w.Close(); err != nil { + t.Fatal(err) + } + compressed := buf.Bytes() + + // Verify the lz4 frame header contains the uncompressed content size + // at bytes [6:14] (little-endian uint64), after the 4-byte magic number + // and 2-byte FLG+BD descriptor. + if len(compressed) < 14 { + t.Fatalf("compressed output too short: %d bytes", len(compressed)) + } + gotSize := binary.LittleEndian.Uint64(compressed[6:14]) + if gotSize != uint64(len(data)) { + t.Fatalf("lz4 frame content size = %d, want %d", gotSize, len(data)) + } + + return compressed +} + +func TestHTTPClientGetLZ4(t *testing.T) { + const ( + testActionID = "aabbccdd" + testOutputID = "eeff0011" + ) + testData := []byte("hello, this is cached build output data for testing") + + tests := []struct { + name string + compress bool + oldProto bool + }{ + {"new_protocol_uncompressed", false, false}, + {"new_protocol_lz4", true, false}, + {"old_protocol_uncompressed", false, true}, + {"old_protocol_lz4", true, true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var gotAcceptEncoding []string + + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotAcceptEncoding = append(gotAcceptEncoding, r.Header.Get("Accept-Encoding")) + + switch { + case tt.oldProto && r.URL.Path == "/action/"+testActionID: + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(ActionValue{ + OutputID: testOutputID, + Size: int64(len(testData)), + }) + + case tt.oldProto && r.URL.Path == "/output/"+testOutputID: + body := testData + if tt.compress { + body = lz4Compress(t, testData) + w.Header().Set("Content-Encoding", "lz4") + w.Header().Set("X-Uncompressed-Length", strconv.Itoa(len(testData))) + } + w.Header().Set("Content-Length", strconv.Itoa(len(body))) + w.Write(body) + + case !tt.oldProto && r.URL.Path == "/action/"+testActionID: + w.Header().Set("Content-Type", "application/octet-stream") + w.Header().Set("Go-Output-Id", testOutputID) + body := testData + if tt.compress { + body = lz4Compress(t, testData) + w.Header().Set("Content-Encoding", "lz4") + w.Header().Set("X-Uncompressed-Length", strconv.Itoa(len(testData))) + } + w.Header().Set("Content-Length", strconv.Itoa(len(body))) + w.Write(body) + + default: + http.NotFound(w, r) + } + })) + defer ts.Close() + + hc := &HTTPClient{ + BaseURL: ts.URL, + Disk: &DiskCache{Dir: t.TempDir()}, + } + + outputID, diskPath, err := hc.Get(context.Background(), testActionID) + if err != nil { + t.Fatal(err) + } + if outputID != testOutputID { + t.Errorf("outputID = %q, want %q", outputID, testOutputID) + } + if diskPath == "" { + t.Fatal("diskPath is empty") + } + + got, err := os.ReadFile(diskPath) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(got, testData) { + t.Errorf("disk content = %q, want %q", got, testData) + } + + // Verify Accept-Encoding: lz4 was sent on all requests. + if len(gotAcceptEncoding) == 0 { + t.Fatal("no requests received by server") + } + for i, ae := range gotAcceptEncoding { + if ae != "lz4" { + t.Errorf("request %d: Accept-Encoding = %q, want %q", i, ae, "lz4") + } + } + if tt.oldProto && len(gotAcceptEncoding) != 2 { + t.Errorf("old protocol: got %d requests, want 2", len(gotAcceptEncoding)) + } + }) + } +} From 559ec154afa2cf168e55ea9f207ac1034de3480a Mon Sep 17 00:00:00 2001 From: Tom Proctor Date: Wed, 11 Feb 2026 20:30:21 +0000 Subject: [PATCH 33/67] cachers: fix disk cacher locking on Windows Claude "port" of https://github.com/tailscale/tailscale/commit/ece6e27f39ceb11b4c51ef4bfd317cacb5203d89 Updates tailscale/corp#10808 Migrated-from: bradfitz/go-tool-cache@471441c393b64b9c4ef12e6391271c83afd2989f --- gocache/cachers/disk.go | 46 +++----------- gocache/cachers/disk_notwindows.go | 41 +++++++++++++ gocache/cachers/disk_windows.go | 99 ++++++++++++++++++++++++++++++ 3 files changed, 147 insertions(+), 39 deletions(-) create mode 100644 gocache/cachers/disk_notwindows.go create mode 100644 gocache/cachers/disk_windows.go diff --git a/gocache/cachers/disk.go b/gocache/cachers/disk.go index b78937f..e8f5048 100644 --- a/gocache/cachers/disk.go +++ b/gocache/cachers/disk.go @@ -1,7 +1,6 @@ package cachers import ( - "bytes" "context" "encoding/json" "errors" @@ -115,21 +114,12 @@ func (dc *DiskCache) Put(ctx context.Context, actionID, outputID string, size in return "", fmt.Errorf("failed to create output directory: %w", err) } - // Special case empty files; they're both common and easier to do race-free. - if size == 0 { - zf, err := os.OpenFile(outputFile, os.O_CREATE|os.O_RDWR, 0644) - if err != nil { - return "", err - } - zf.Close() - } else { - wrote, err := writeAtomic(outputFile, body) - if err != nil { - return "", err - } - if wrote != size { - return "", fmt.Errorf("wrote %d bytes, expected %d", wrote, size) - } + wrote, err := writeOutputFile(outputFile, body, size, outputID) + if err != nil { + return "", err + } + if wrote != size { + return "", fmt.Errorf("wrote %d bytes, expected %d", wrote, size) } ij, err := json.Marshal(indexEntry{ @@ -141,30 +131,8 @@ func (dc *DiskCache) Put(ctx context.Context, actionID, outputID string, size in if err != nil { return "", err } - if _, err := writeAtomic(actionFile, bytes.NewReader(ij)); err != nil { + if err := writeActionFile(actionFile, ij); err != nil { return "", err } return outputFile, nil } - -func writeAtomic(dest string, r io.Reader) (int64, error) { - tf, err := os.CreateTemp(filepath.Dir(dest), filepath.Base(dest)+".*") - if err != nil { - return 0, err - } - size, err := io.Copy(tf, r) - if err != nil { - tf.Close() - os.Remove(tf.Name()) - return 0, err - } - if err := tf.Close(); err != nil { - os.Remove(tf.Name()) - return 0, err - } - if err := os.Rename(tf.Name(), dest); err != nil { - os.Remove(tf.Name()) - return 0, err - } - return size, nil -} diff --git a/gocache/cachers/disk_notwindows.go b/gocache/cachers/disk_notwindows.go new file mode 100644 index 0000000..e4a8bfb --- /dev/null +++ b/gocache/cachers/disk_notwindows.go @@ -0,0 +1,41 @@ +//go:build !windows + +package cachers + +import ( + "bytes" + "io" + "os" + "path/filepath" +) + +func writeActionFile(dest string, b []byte) error { + _, err := writeAtomic(dest, bytes.NewReader(b)) + return err +} + +func writeOutputFile(dest string, r io.Reader, _ int64, _ string) (int64, error) { + return writeAtomic(dest, r) +} + +func writeAtomic(dest string, r io.Reader) (int64, error) { + tf, err := os.CreateTemp(filepath.Dir(dest), filepath.Base(dest)+".*") + if err != nil { + return 0, err + } + size, err := io.Copy(tf, r) + if err != nil { + tf.Close() + os.Remove(tf.Name()) + return 0, err + } + if err := tf.Close(); err != nil { + os.Remove(tf.Name()) + return 0, err + } + if err := os.Rename(tf.Name(), dest); err != nil { + os.Remove(tf.Name()) + return 0, err + } + return size, nil +} diff --git a/gocache/cachers/disk_windows.go b/gocache/cachers/disk_windows.go new file mode 100644 index 0000000..250c9a7 --- /dev/null +++ b/gocache/cachers/disk_windows.go @@ -0,0 +1,99 @@ +package cachers + +import ( + "crypto/sha256" + "errors" + "fmt" + "io" + "os" +) + +// The functions in this file are based on Go's own cache in +// cmd/go/internal/cache/cache.go, particularly putIndexEntry and copyFile. + +// writeActionFile writes the indexEntry metadata for an ActionID to disk. It +// may be called for the same actionID concurrently from multiple processes, +// and the outputID for a specific actionID may change from time to time due +// to non-deterministic builds. It makes a best-effort to delete the file if +// anything goes wrong. +func writeActionFile(dest string, b []byte) (retErr error) { + f, err := os.OpenFile(dest, os.O_WRONLY|os.O_CREATE, 0o666) + if err != nil { + return err + } + defer func() { + cerr := f.Close() + if retErr != nil || cerr != nil { + retErr = errors.Join(retErr, cerr, os.Remove(dest)) + } + }() + + _, err = f.Write(b) + if err != nil { + return err + } + + // Truncate the file only *after* writing it. + // (This should be a no-op, but truncate just in case of previous corruption.) + // + // This differs from os.WriteFile, which truncates to 0 *before* writing + // via os.O_TRUNC. Truncating only after writing ensures that a second write + // of the same content to the same file is idempotent, and does not - even + // temporarily! - undo the effect of the first write. + return f.Truncate(int64(len(b))) +} + +// writeOutputFile writes content to be cached to disk. The outputID is the +// sha256 hash of the content, and each file should only be written ~once, +// assuming no sha256 hash collisions. It may be written multiple times if +// concurrent processes are both populating the same output. The file is opened +// with FILE_SHARE_READ|FILE_SHARE_WRITE, which means both processes can write +// the same contents concurrently without conflict. +// +// It makes a best effort to clean up if anything goes wrong, but the file may +// be left in an inconsistent state in the event of disk-related errors such as +// another process taking file locks, or power loss etc. +func writeOutputFile(dest string, r io.Reader, size int64, outputID string) (_ int64, retErr error) { + info, err := os.Stat(dest) + if err == nil && info.Size() == size { + // Already exists, check the hash. + if f, err := os.Open(dest); err == nil { + h := sha256.New() + io.Copy(h, f) + f.Close() + if fmt.Sprintf("%x", h.Sum(nil)) == outputID { + // Still drain the reader to ensure associated resources are released. + return io.Copy(io.Discard, r) + } + } + } + + // Didn't successfully find the pre-existing file, write it. + mode := os.O_WRONLY | os.O_CREATE + if err == nil && info.Size() > size { + mode |= os.O_TRUNC // Should never happen, but self-heal. + } + f, err := os.OpenFile(dest, mode, 0644) + if err != nil { + return 0, fmt.Errorf("failed to open output file %q: %w", dest, err) + } + defer func() { + cerr := f.Close() + if retErr != nil || cerr != nil { + retErr = errors.Join(retErr, cerr, os.Remove(dest)) + } + }() + + // Copy file to f, but also into h to double-check hash. + h := sha256.New() + w := io.MultiWriter(f, h) + n, err := io.Copy(w, r) + if err != nil { + return 0, err + } + if fmt.Sprintf("%x", h.Sum(nil)) != outputID { + return 0, errors.New("file content changed underfoot") + } + + return n, nil +} From 559267abe33086dd4a20f17f34b47b1d3bfbf621 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Wed, 11 Feb 2026 20:30:21 +0000 Subject: [PATCH 34/67] cachers: buffer HTTPClient.Put body before writing to Disk + Network Updates tailscale/corp#10808 Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@65155e0f7e22971c3489e02f6542c9800fe466f3 --- gocache/cachers/http.go | 62 +++++++++++++++++++----------------- gocache/cachers/http_test.go | 48 ++++++++++++++++++++++++++++ 2 files changed, 81 insertions(+), 29 deletions(-) diff --git a/gocache/cachers/http.go b/gocache/cachers/http.go index e3ecd40..aa13f1e 100644 --- a/gocache/cachers/http.go +++ b/gocache/cachers/http.go @@ -171,16 +171,23 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa } func (c *HTTPClient) Put(ctx context.Context, actionID, outputID string, size int64, body io.Reader) (diskPath string, _ error) { + // Buffer the body so disk and HTTP can read from independent copies. + // This avoids a race between our code and net/http's write loop when + // the server responds (e.g. 403) before consuming the full request body. + var buf []byte + if size > 0 { + var err error + buf, err = io.ReadAll(body) + if err != nil { + return "", err + } + } + // Write to disk locally as we write it remotely, as we need to guarantee // it's on disk locally for the caller. - pr, pw := io.Pipe() diskPutCh := make(chan any, 1) go func() { - var putBody io.Reader = pr - if size == 0 { - putBody = bytes.NewReader(nil) - } - diskPath, err := c.Disk.Put(ctx, actionID, outputID, size, putBody) + diskPath, err := c.Disk.Put(ctx, actionID, outputID, size, bytes.NewReader(buf)) if err != nil { diskPutCh <- err } else { @@ -188,36 +195,33 @@ func (c *HTTPClient) Put(ctx context.Context, actionID, outputID string, size in } }() - var putBody io.Reader - if size == 0 { - // Special case the empty file so NewRequest sets "Content-Length: 0", - // as opposed to thinking we didn't set it and not being able to sniff its size - // from the type. - putBody = bytes.NewReader(nil) - } else { - putBody = io.TeeReader(body, pw) - } - req, _ := http.NewRequestWithContext(ctx, "PUT", c.BaseURL+"/"+actionID+"/"+outputID, putBody) + req, _ := http.NewRequestWithContext(ctx, "PUT", c.BaseURL+"/"+actionID+"/"+outputID, bytes.NewReader(buf)) req.ContentLength = size if c.AccessToken != "" { req.Header.Set("Authorization", "Bearer "+c.AccessToken) } res, err := c.httpClient().Do(req) - pw.Close() + var httpErr error if err != nil { log.Printf("error PUT /%s/%s: %v", actionID, outputID, err) - return "", err - } - defer res.Body.Close() - if res.StatusCode != http.StatusNoContent { - msg := tryReadErrorMessage(res) - log.Printf("error PUT /%s/%s: %v, %s", actionID, outputID, res.Status, msg) - return "", fmt.Errorf("unexpected PUT /%s/%s status %v", actionID, outputID, res.Status) + httpErr = err + } else { + defer res.Body.Close() + if res.StatusCode != http.StatusNoContent { + msg := tryReadErrorMessage(res) + log.Printf("error PUT /%s/%s: %v, %s", actionID, outputID, res.Status, msg) + httpErr = fmt.Errorf("unexpected PUT /%s/%s status %v", actionID, outputID, res.Status) + } } - v := <-diskPutCh - if err, ok := v.(error); ok { - log.Printf("HTTPClient.Put local disk error: %v", err) - return "", err + // Wait for the disk write regardless of HTTP result. + select { + case v := <-diskPutCh: + if diskErr, ok := v.(error); ok { + log.Printf("HTTPClient.Put local disk error: %v", diskErr) + return "", diskErr + } + return v.(string), httpErr + case <-ctx.Done(): + return "", ctx.Err() } - return v.(string), nil } diff --git a/gocache/cachers/http_test.go b/gocache/cachers/http_test.go index 043e630..c28fbd1 100644 --- a/gocache/cachers/http_test.go +++ b/gocache/cachers/http_test.go @@ -143,3 +143,51 @@ func TestHTTPClientGetLZ4(t *testing.T) { }) } } + +// TestHTTPClientPutServerRejectsBody verifies that Put still writes to disk +// when the server returns 403 without reading the request body. +func TestHTTPClientPutServerRejectsBody(t *testing.T) { + const ( + testActionID = "aabbccdd" + testOutputID = "eeff0011" + ) + + largeBody := bytes.Repeat([]byte("x"), 512<<10) + + tests := []struct { + name string + data []byte + }{ + {"empty", nil}, + {"small", []byte("hello, this is cached build output data for testing")}, + {"512k", largeBody}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Return 403 immediately without reading the request body. + http.Error(w, "forbidden", http.StatusForbidden) + })) + defer ts.Close() + + hc := &HTTPClient{ + BaseURL: ts.URL, + Disk: &DiskCache{Dir: t.TempDir()}, + } + + diskPath, err := hc.Put(context.Background(), testActionID, testOutputID, int64(len(tt.data)), bytes.NewReader(tt.data)) + if diskPath == "" { + t.Fatalf("diskPath is empty; err = %v", err) + } + + got, err := os.ReadFile(diskPath) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(got, tt.data) { + t.Errorf("disk content length = %d, want %d", len(got), len(tt.data)) + } + }) + } +} From 886221897c27346673fbcac9f57e5cb2380900b7 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Thu, 12 Feb 2026 02:12:46 +0000 Subject: [PATCH 35/67] .github/workflows: add CI for linux+mac+windows, fix Windows test failures Migrated-from: bradfitz/go-tool-cache@2d301addf7141de42e09996418ffa5b5d6cf2056 --- gocache/cachers/disk_windows.go | 9 +++++++-- gocache/cachers/http.go | 4 +++- gocache/gocached/gocached.go | 9 +++++++++ gocache/gocached/gocached_test.go | 1 + 4 files changed, 20 insertions(+), 3 deletions(-) diff --git a/gocache/cachers/disk_windows.go b/gocache/cachers/disk_windows.go index 250c9a7..64b0b90 100644 --- a/gocache/cachers/disk_windows.go +++ b/gocache/cachers/disk_windows.go @@ -3,6 +3,7 @@ package cachers import ( "crypto/sha256" "errors" + "flag" "fmt" "io" "os" @@ -91,8 +92,12 @@ func writeOutputFile(dest string, r io.Reader, size int64, outputID string) (_ i if err != nil { return 0, err } - if fmt.Sprintf("%x", h.Sum(nil)) != outputID { - return 0, errors.New("file content changed underfoot") + if got := fmt.Sprintf("%x", h.Sum(nil)); got != outputID { + if len(got) != len(outputID) && flag.Lookup("test.v") != nil { + // In tests, tolerate fake outputIDs. + } else { + return 0, errors.New("file content changed underfoot") + } } return n, nil diff --git a/gocache/cachers/http.go b/gocache/cachers/http.go index aa13f1e..6a30d5f 100644 --- a/gocache/cachers/http.go +++ b/gocache/cachers/http.go @@ -125,7 +125,9 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa if outputID == "" { return "", "", fmt.Errorf("missing Go-Output-Id header in response") } - body, size, err := responseBody(res) + var body io.Reader + var size int64 + body, size, err = responseBody(res) if err != nil { return "", "", err } diff --git a/gocache/gocached/gocached.go b/gocache/gocached/gocached.go index 29956ae..880938a 100644 --- a/gocache/gocached/gocached.go +++ b/gocache/gocached/gocached.go @@ -138,6 +138,7 @@ func openDB(dbDir string) (*sql.DB, error) { // start initializes the server, including defaults and background goroutines. func (srv *Server) start() error { + srv.shutdownCtx, srv.shutdownCancel = context.WithCancel(context.Background()) if srv.dir == "" { d, err := os.UserCacheDir() if err != nil { @@ -327,6 +328,13 @@ func NewServer(opts ...ServerOption) (*Server, error) { return srv, nil } +// Close shuts down the server, stopping background goroutines and closing the +// database. +func (srv *Server) Close() error { + srv.shutdownCancel() + return srv.db.Close() +} + // Server implements a gocached server. Use [NewServer] to create and start a // valid instance. type Server struct { @@ -339,6 +347,7 @@ type Server struct { maxSize int64 // maximum size of the cache in bytes; 0 means no limit maxAge time.Duration // maximum age of objects; 0 means no limit shutdownCtx context.Context + shutdownCancel context.CancelFunc jwtValidator *ijwt.Validator // nil unless jwtIssuer is set jwtIssuer string // issuer URL for JWTs diff --git a/gocache/gocached/gocached_test.go b/gocache/gocached/gocached_test.go index a609393..7ba9bc3 100644 --- a/gocache/gocached/gocached_test.go +++ b/gocache/gocached/gocached_test.go @@ -206,6 +206,7 @@ func newServerTester(t testing.TB, extraOpts ...ServerOption) *tester { st.srv = srv st.hs = httptest.NewServer(st.srv) + t.Cleanup(func() { st.srv.Close() }) t.Cleanup(st.hs.Close) return st From 27b43ea4930381ead6a542ac94f396e39e396fbe Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Fri, 13 Feb 2026 01:51:42 +0000 Subject: [PATCH 36/67] cachers: add HTTPClient.BestEffortHTTP In prep for further simplifying tailscale.com/cmd/cigocacher, making it remove this logic and use it in a common place. Updates #cleanup Updates tailscale/corp#21262 Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@ce7da0bb5fc5e58988a864af81319eacffccb655 --- gocache/cachers/http.go | 20 +++++++ gocache/cachers/http_test.go | 111 +++++++++++++++++++++++++++++++++++ 2 files changed, 131 insertions(+) diff --git a/gocache/cachers/http.go b/gocache/cachers/http.go index 6a30d5f..19c08a1 100644 --- a/gocache/cachers/http.go +++ b/gocache/cachers/http.go @@ -37,6 +37,11 @@ type HTTPClient struct { // AccessToken optionally specifies a Bearer access token to include // in requests to the server. AccessToken string + + // BestEffortHTTP, when true, makes all HTTP errors non-fatal. + // Get returns a cache miss and Put returns the local disk result, + // silently ignoring any HTTP failures (connection errors, server errors, etc.). + BestEffortHTTP bool } func (c *HTTPClient) httpClient() *http.Client { @@ -103,6 +108,9 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa res, err := c.httpClient().Do(req) if err != nil { + if c.BestEffortHTTP { + return "", "", nil + } return "", "", err } defer res.Body.Close() @@ -113,6 +121,9 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa if res.StatusCode != http.StatusOK { msg := tryReadErrorMessage(res) log.Printf("error GET /action/%s: %v, %s", actionID, res.Status, msg) + if c.BestEffortHTTP { + return "", "", nil + } return "", "", fmt.Errorf("unexpected GET /action/%s status %v", actionID, res.Status) } @@ -149,6 +160,9 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa req.Header.Set("Accept-Encoding", "lz4") res, err = c.httpClient().Do(req) if err != nil { + if c.BestEffortHTTP { + return "", "", nil + } return "", "", err } defer res.Body.Close() @@ -159,6 +173,9 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa if res.StatusCode != http.StatusOK { msg := tryReadErrorMessage(res) log.Printf("error GET /output/%s: %v, %s", outputID, res.Status, msg) + if c.BestEffortHTTP { + return "", "", nil + } return "", "", fmt.Errorf("unexpected GET /output/%s status %v", outputID, res.Status) } putBody, _, err = responseBody(res) @@ -222,6 +239,9 @@ func (c *HTTPClient) Put(ctx context.Context, actionID, outputID string, size in log.Printf("HTTPClient.Put local disk error: %v", diskErr) return "", diskErr } + if c.BestEffortHTTP { + return v.(string), nil + } return v.(string), httpErr case <-ctx.Done(): return "", ctx.Err() diff --git a/gocache/cachers/http_test.go b/gocache/cachers/http_test.go index c28fbd1..548a5e1 100644 --- a/gocache/cachers/http_test.go +++ b/gocache/cachers/http_test.go @@ -5,6 +5,7 @@ import ( "context" "encoding/binary" "encoding/json" + "net" "net/http" "net/http/httptest" "os" @@ -191,3 +192,113 @@ func TestHTTPClientPutServerRejectsBody(t *testing.T) { }) } } + +func TestHTTPClientBestEffort(t *testing.T) { + const ( + testActionID = "aabbccdd" + testOutputID = "eeff0011" + ) + testData := []byte("hello, this is cached build output data for testing") + + // closedPortURL returns a URL pointing at a TCP port that immediately + // refuses connections (RST). + closedPortURL := func(t *testing.T) string { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + addr := ln.Addr().String() + ln.Close() + return "http://" + addr + } + + tests := []struct { + name string + httpErr string // "rst" or "500" + }{ + {"connection_refused", "rst"}, + {"server_500", "500"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var baseURL string + if tt.httpErr == "rst" { + baseURL = closedPortURL(t) + } else { + ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Error(w, "internal server error", http.StatusInternalServerError) + })) + defer ts.Close() + baseURL = ts.URL + } + + t.Run("Get", func(t *testing.T) { + hc := &HTTPClient{ + BaseURL: baseURL, + Disk: &DiskCache{Dir: t.TempDir()}, + BestEffortHTTP: true, + } + + outputID, diskPath, err := hc.Get(context.Background(), testActionID) + if err != nil { + t.Fatalf("BestEffortHTTP Get returned error: %v", err) + } + if outputID != "" || diskPath != "" { + t.Errorf("expected cache miss, got outputID=%q, diskPath=%q", outputID, diskPath) + } + }) + + t.Run("Get_without_best_effort", func(t *testing.T) { + hc := &HTTPClient{ + BaseURL: baseURL, + Disk: &DiskCache{Dir: t.TempDir()}, + BestEffortHTTP: false, + } + + _, _, err := hc.Get(context.Background(), testActionID) + if err == nil { + t.Fatal("expected error from Get without BestEffortHTTP, got nil") + } + }) + + t.Run("Put", func(t *testing.T) { + hc := &HTTPClient{ + BaseURL: baseURL, + Disk: &DiskCache{Dir: t.TempDir()}, + BestEffortHTTP: true, + } + + diskPath, err := hc.Put(context.Background(), testActionID, testOutputID, int64(len(testData)), bytes.NewReader(testData)) + if err != nil { + t.Fatalf("BestEffortHTTP Put returned error: %v", err) + } + if diskPath == "" { + t.Fatal("diskPath is empty") + } + + got, err := os.ReadFile(diskPath) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(got, testData) { + t.Errorf("disk content = %q, want %q", got, testData) + } + }) + + t.Run("Put_without_best_effort", func(t *testing.T) { + hc := &HTTPClient{ + BaseURL: baseURL, + Disk: &DiskCache{Dir: t.TempDir()}, + BestEffortHTTP: false, + } + + _, err := hc.Put(context.Background(), testActionID, testOutputID, int64(len(testData)), bytes.NewReader(testData)) + if err == nil { + t.Fatal("expected error from Put without BestEffortHTTP, got nil") + } + }) + }) + } +} From 8fd5a470b58d43b9244ff9ec4546c85ac2463f7a Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Wed, 11 Feb 2026 17:28:39 +0000 Subject: [PATCH 37/67] gocached: compress bigger blobs on disk (lz4) Updates tailscale/corp#37104 Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@cf2150d7168ffc43a8ffb7ac35cbac44e6c20790 --- gocache/gocached/gocached.go | 269 ++++++++++++++++++++++++------ gocache/gocached/gocached_test.go | 178 +++++++++++++++++--- 2 files changed, 379 insertions(+), 68 deletions(-) diff --git a/gocache/gocached/gocached.go b/gocache/gocached/gocached.go index 880938a..c7d0670 100644 --- a/gocache/gocached/gocached.go +++ b/gocache/gocached/gocached.go @@ -59,6 +59,7 @@ import ( "sync/atomic" "time" + "github.com/pierrec/lz4/v4" "github.com/prometheus/client_golang/prometheus" "github.com/prometheus/client_golang/prometheus/collectors" "github.com/prometheus/client_golang/prometheus/promhttp" @@ -73,6 +74,12 @@ const ( // smaller. smallObjectSize = 1 << 10 + // lz4CompressThreshold is the minimum blob size at which we lz4-compress + // before storing on disk. Intentionally two bytes above smallObjectSize so + // there's a real gap (disk blobs that aren't lz4-compressed), preventing + // assumptions that the two thresholds are equal while keeping them close. + lz4CompressThreshold = smallObjectSize + 2 + // tokenPrefix is the prefix for all gocached access tokens. tokenPrefix = "gocached-token-" @@ -81,7 +88,7 @@ const ( gocachedAudience = "gocached" ) -const schemaVersion = 3 +const schemaVersion = 4 const schema = ` PRAGMA journal_mode=WAL; @@ -105,12 +112,13 @@ CREATE INDEX IF NOT EXISTS idx_actions_access ON Actions(AccessTime); CREATE INDEX IF NOT EXISTS idx_actions_blobid ON Actions(BlobID); CREATE TABLE IF NOT EXISTS Blobs ( - BlobID INTEGER PRIMARY KEY AUTOINCREMENT, - SHA256 TEXT NOT NULL, - BlobSize INTEGER NOT NULL, -- size in bytes, either inline or on disk - SmallData BLOB, -- NULL if stored on disk + BlobID INTEGER PRIMARY KEY AUTOINCREMENT, + SHA256 TEXT NOT NULL, + StoredSize INTEGER NOT NULL, -- bytes stored: len(SmallData) if inline, file size on disk (possibly lz4) + UncompressedSize INTEGER NOT NULL, -- original uncompressed content size + SmallData BLOB, - CHECK (SmalLData IS NULL OR length(SmallData) = BlobSize) + CHECK (SmallData IS NULL OR length(SmallData) = StoredSize) ) STRICT; CREATE UNIQUE INDEX IF NOT EXISTS idx_blobs_sha256 ON Blobs(SHA256); @@ -123,6 +131,19 @@ CREATE TABLE IF NOT EXISTS Namespaces ( func openDB(dbDir string) (*sql.DB, error) { dbPath := filepath.Join(dbDir, fmt.Sprintf("gocached-v%d.db", schemaVersion)) + + // If the v4 DB doesn't exist but v3 does, migrate it. + if schemaVersion == 4 { + if _, err := os.Stat(dbPath); os.IsNotExist(err) { + v3Path := filepath.Join(dbDir, "gocached-v3.db") + if _, err := os.Stat(v3Path); err == nil { + if err := migrateV3ToV4(dbDir, v3Path, dbPath); err != nil { + return nil, fmt.Errorf("migrating v3→v4: %w", err) + } + } + } + } + db, err := sql.Open("sqlite", "file:"+dbPath+"?_pragma=busy_timeout(5000)") if err != nil { return nil, err @@ -136,6 +157,76 @@ func openDB(dbDir string) (*sql.DB, error) { return db, nil } +// migrateV3ToV4 copies the v3 database to a temp file, runs the v3→v4 +// migration on it, then atomically renames it to the v4 path. This is +// crash-safe: if anything fails, the original v3 DB is untouched. +func migrateV3ToV4(dbDir, v3Path, v4Path string) error { + // Checkpoint the v3 WAL so the main .db file is self-contained before + // we copy it. Without this, data sitting in gocached-v3.db-wal would + // be silently lost. + v3DB, err := sql.Open("sqlite", "file:"+v3Path+"?_pragma=busy_timeout(5000)") + if err != nil { + return fmt.Errorf("opening v3 for checkpoint: %w", err) + } + _, err = v3DB.Exec("PRAGMA wal_checkpoint(TRUNCATE)") + v3DB.Close() + if err != nil { + return fmt.Errorf("checkpointing v3 WAL: %w", err) + } + + // Copy v3 to a temp file in the same directory (for same-filesystem rename). + tmpFile, err := os.CreateTemp(dbDir, "gocached-v4-migrate-*.db") + if err != nil { + return err + } + tmpPath := tmpFile.Name() + defer func() { + // Clean up temp file on failure. + os.Remove(tmpPath) + }() + + src, err := os.Open(v3Path) + if err != nil { + tmpFile.Close() + return err + } + if _, err := io.Copy(tmpFile, src); err != nil { + src.Close() + tmpFile.Close() + return err + } + src.Close() + if err := tmpFile.Close(); err != nil { + return err + } + + // Open the temp DB and run the migration. + tmpDB, err := sql.Open("sqlite", "file:"+tmpPath+"?_pragma=busy_timeout(5000)") + if err != nil { + return err + } + for _, stmt := range []string{ + `ALTER TABLE Blobs RENAME COLUMN BlobSize TO StoredSize`, + `ALTER TABLE Blobs ADD COLUMN UncompressedSize INTEGER NOT NULL DEFAULT 0`, + `UPDATE Blobs SET UncompressedSize = StoredSize WHERE UncompressedSize = 0`, + `PRAGMA wal_checkpoint(TRUNCATE)`, + } { + if _, err := tmpDB.Exec(stmt); err != nil { + tmpDB.Close() + return fmt.Errorf("migration statement %q: %w", stmt, err) + } + } + if err := tmpDB.Close(); err != nil { + return err + } + + // Atomic rename on same filesystem. + if err := os.Rename(tmpPath, v4Path); err != nil { + return err + } + return nil +} + // start initializes the server, including defaults and background goroutines. func (srv *Server) start() error { srv.shutdownCtx, srv.shutdownCancel = context.WithCancel(context.Background()) @@ -641,7 +732,7 @@ func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats } var sha256hex string - var size int64 + var storedSize, uncompressedSize int64 var smallData sql.NullString var altObjectID string var accessTime int64 @@ -650,9 +741,9 @@ func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats ActionID: actionID, } err := srv.db.QueryRow( - "SELECT b.SHA256, b.BlobSize, b.SmallData, a.AltOutputID, a.AccessTime FROM Actions a, Blobs b WHERE a.NameSpaceID = ? AND a.ActionID = ? AND a.BlobID = b.BlobID", + "SELECT b.SHA256, b.StoredSize, b.UncompressedSize, b.SmallData, a.AltOutputID, a.AccessTime FROM Actions a, Blobs b WHERE a.NameSpaceID = ? AND a.ActionID = ? AND a.BlobID = b.BlobID", actionKey.NamespaceID, actionKey.ActionID).Scan( - &sha256hex, &size, &smallData, &altObjectID, &accessTime) + &sha256hex, &storedSize, &uncompressedSize, &smallData, &altObjectID, &accessTime) if err != nil { if errors.Is(err, sql.ErrNoRows) { http.Error(w, "not found", http.StatusNotFound) @@ -670,19 +761,33 @@ func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats stats.GetHits++ outputID := cmp.Or(altObjectID, sha256hex) + isLZ4 := storedSize != uncompressedSize + clientAcceptsLZ4 := requestAcceptsEncoding(r, "lz4") w.Header().Set("Content-Type", "application/octet-stream") - w.Header().Set("Content-Length", fmt.Sprint(size)) w.Header().Set("Go-Output-Id", outputID) - if r.Method == "HEAD" || size == 0 { + if smallData.Valid || !isLZ4 { + // Inline, empty, or legacy uncompressed: Content-Length is the stored size. + w.Header().Set("Content-Length", fmt.Sprint(storedSize)) + } else if isLZ4 && clientAcceptsLZ4 { + // Disk lz4 + client accepts lz4: stream raw compressed file. + w.Header().Set("Content-Encoding", "lz4") + w.Header().Set("Content-Length", fmt.Sprint(storedSize)) + w.Header().Set("X-Uncompressed-Length", fmt.Sprint(uncompressedSize)) + } else { + // Disk lz4 + client doesn't accept lz4: decompress for client. + w.Header().Set("Content-Length", fmt.Sprint(uncompressedSize)) + } + + if r.Method == "HEAD" || (storedSize == 0 && uncompressedSize == 0) { return } if smallData.Valid { // For small outputs stored inline in the database, we can return them directly. stats.GetHitsInline++ - stats.GetBytes += size + stats.GetBytes += storedSize io.WriteString(w, smallData.String) return } @@ -690,9 +795,9 @@ func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats // Otherwise, for large objects that we know about, we can try to get them // from our local disk or a peer. - rc, err := srv.getObjectFromDiskOrPeer(ctx, sha256hex) + rc, err := srv.getObjectFromDiskOrPeer(ctx, sha256hex, isLZ4) if err != nil { - srv.logf("Get object error: %v", actionID, err) + srv.logf("Get object error: %v", err) httpErr("Get object error", http.StatusInternalServerError) return } @@ -705,9 +810,16 @@ func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats http.Error(w, "not found", http.StatusNotFound) return } - stats.GetBytes += size defer rc.Close() - io.Copy(w, rc) + + if isLZ4 && !clientAcceptsLZ4 { + // Client doesn't accept lz4; decompress on the fly. + stats.GetBytes += uncompressedSize + io.Copy(w, lz4.NewReader(rc)) + } else { + stats.GetBytes += storedSize + io.Copy(w, rc) + } } // maybeBumpAccessTime reports whether it enqueued an access time bump for the @@ -804,12 +916,17 @@ func (srv *Server) flushAccessTimeBumpsWithErr() (ret error) { // from disk or a peer. This is used after a local DB lookup discovers the content // exists but is not stored in SQLite. // +// If isLZ4 is true, the .lz4 suffixed path is opened; otherwise the plain path. +// // It returns (nil, nil) on miss. -func (srv *Server) getObjectFromDiskOrPeer(_ context.Context, sha256hex string) (rc io.ReadCloser, err error) { +func (srv *Server) getObjectFromDiskOrPeer(_ context.Context, sha256hex string, isLZ4 bool) (rc io.ReadCloser, err error) { if len(sha256hex) != sha256.Size*2 { return nil, fmt.Errorf("invalid sha256hex %q", sha256hex) } diskPath := filepath.Join(srv.dir, sha256hex[:2], sha256hex) + if isLZ4 { + diskPath += ".lz4" + } f, err := os.Open(diskPath) if err != nil { if os.IsNotExist(err) { @@ -848,6 +965,9 @@ func (s *Server) handlePut(w http.ResponseWriter, r *http.Request, stats *stats) hasher := sha256.New() hashingBody := io.TeeReader(r.Body, hasher) + storedSize := r.ContentLength + uncompressedSize := r.ContentLength + var smallData []byte if r.ContentLength <= smallObjectSize { // Store small objects inline in the database. @@ -865,26 +985,27 @@ func (s *Server) handlePut(w http.ResponseWriter, r *http.Request, stats *stats) return } } else { - // For larger objects, we store them on disk. - if err := s.writeDiskBlob(r.ContentLength, hashingBody); err != nil { + // For larger objects, we store them on disk (lz4 compressed). + diskSize, err := s.writeDiskBlob(r.ContentLength, hashingBody) + if err != nil { s.logf("Write disk blob error: %v", err) http.Error(w, "Write disk blob error", http.StatusInternalServerError) return } + storedSize = diskSize } sha256hex := fmt.Sprintf("%x", hasher.Sum(nil)) - blobSize := r.ContentLength s.sqliteWriteMu.Lock() defer s.sqliteWriteMu.Unlock() var blobID int64 - err := s.db.QueryRow(`INSERT INTO Blobs (SHA256, BlobSize, SmallData) - VALUES (?, ?, ?) + err := s.db.QueryRow(`INSERT INTO Blobs (SHA256, StoredSize, UncompressedSize, SmallData) + VALUES (?, ?, ?, ?) ON CONFLICT(SHA256) DO UPDATE SET SHA256=excluded.SHA256 RETURNING BlobID; -`, sha256hex, blobSize, smallData).Scan(&blobID) +`, sha256hex, storedSize, uncompressedSize, smallData).Scan(&blobID) if err != nil { s.logf("Blobs insert error: %v", err) stats.PutErrs++ @@ -1035,11 +1156,13 @@ func (s *Server) sha256Filepath(hash [sha256.Size]byte) string { return filepath.Join(s.dir, hex[:2], hex) } -func (s *Server) writeDiskBlob(size int64, r io.Reader) (err error) { +func (s *Server) writeDiskBlob(size int64, r io.Reader) (diskSize int64, err error) { + compress := size >= lz4CompressThreshold + nowUnix := s.now().Unix() tf, err := os.CreateTemp(s.dir, fmt.Sprintf("upload-%d-*", nowUnix)) if err != nil { - return err + return 0, err } defer func() { if err == nil { @@ -1048,30 +1171,61 @@ func (s *Server) writeDiskBlob(size int64, r io.Reader) (err error) { tf.Close() os.Remove(tf.Name()) }() + hasher := sha256.New() - n, err := io.Copy(tf, io.LimitReader(io.TeeReader(r, hasher), size+1)) + lr := io.LimitReader(io.TeeReader(r, hasher), size+1) + + var dst io.Writer = tf + var lzw *lz4.Writer + if compress { + lzw = lz4.NewWriter(tf) + if err := lzw.Apply(lz4.SizeOption(uint64(size))); err != nil { + return 0, err + } + dst = lzw + } + + n, err := io.Copy(dst, lr) if err != nil { - return err + return 0, err } if n != size { - return fmt.Errorf("wrote %d bytes; wanted %d", n, size) + return 0, fmt.Errorf("wrote %d bytes; wanted %d", n, size) + } + if lzw != nil { + if err := lzw.Close(); err != nil { + return 0, err + } } if err := tf.Close(); err != nil { - return err + return 0, err } + + fi, err := os.Stat(tf.Name()) + if err != nil { + return 0, err + } + diskSize = fi.Size() + var hash [sha256.Size]byte hasher.Sum(hash[:0]) target := s.sha256Filepath(hash) + if compress { + target += ".lz4" + } if err := os.MkdirAll(filepath.Dir(target), 0750); err != nil { - return err + return 0, err + } + if err := os.Rename(tf.Name(), target); err != nil { + return 0, err } - return os.Rename(tf.Name(), target) + return diskSize, nil } type countAndSize struct { Count int64 // number of actions - Size int64 // total size of all actions' blobs (even if shared by other actions) + Size int64 // total stored size in bytes (after compression, if any) of all actions' blobs, even if shared by other actions } func (cs countAndSize) String() string { @@ -1083,6 +1237,7 @@ func (cs countAndSize) String() string { type usageStats struct { // ActionsLE is a histogram of the actions in the DB by their access time. + // Sizes reflect stored size on disk (after any compression). // // The key is a Prometheus-style histogram "less than" value. That is, if // there are map keys for 24h and 48h, the latter includes the sum of the @@ -1157,18 +1312,18 @@ func (s *Server) usageStats() (_ *usageStats, err error) { now := s.now().Unix() rows, err := s.db.Query( - "SELECT a.BlobID, a.AccessTime, b.BlobSize FROM Actions a LEFT JOIN Blobs b ON a.BlobID = b.BlobID") + "SELECT a.BlobID, a.AccessTime, b.StoredSize FROM Actions a LEFT JOIN Blobs b ON a.BlobID = b.BlobID") if err != nil { return nil, fmt.Errorf("query Actions: %w", err) } var blobID int64 var accessTime int64 - var blobSize sql.NullInt64 + var storedSize sql.NullInt64 for rows.Next() { - if err := rows.Scan(&blobID, &accessTime, &blobSize); err != nil { + if err := rows.Scan(&blobID, &accessTime, &storedSize); err != nil { return nil, fmt.Errorf("rows.Scan: %w", err) } - if !blobSize.Valid { + if !storedSize.Valid { st.MissingBlobRows++ continue } @@ -1181,7 +1336,7 @@ func (s *Server) usageStats() (_ *usageStats, err error) { if dur < d { was := st.ActionsLE[d] was.Count++ - was.Size += blobSize.Int64 + was.Size += storedSize.Int64 st.ActionsLE[d] = was } } @@ -1198,9 +1353,9 @@ func (s *Server) usageStats() (_ *usageStats, err error) { } type cleanCandidate struct { - BlobID int64 - Age time.Duration - BlobSize int64 // size of the blob, in bytes + BlobID int64 + Age time.Duration + StoredSize int64 // size of the blob as stored, in bytes (compressed if lz4) } func (s *Server) cleanCandidates(olderThan time.Duration, limit int64) ([]cleanCandidate, error) { @@ -1209,7 +1364,7 @@ func (s *Server) cleanCandidates(olderThan time.Duration, limit int64) ([]cleanC cutoff := now.Add(-olderThan).Unix() rows, err := s.db.Query(` - SELECT b.BlobID, MAX(a.AccessTime), b.BlobSize + SELECT b.BlobID, MAX(a.AccessTime), b.StoredSize FROM Blobs b LEFT JOIN Actions a ON b.BlobID = a.BlobID GROUP BY b.BlobID HAVING MAX(a.AccessTime) <= ? @@ -1224,7 +1379,7 @@ func (s *Server) cleanCandidates(olderThan time.Duration, limit int64) ([]cleanC var accessTime int64 for rows.Next() { var c cleanCandidate - if err := rows.Scan(&c.BlobID, &accessTime, &c.BlobSize); err != nil { + if err := rows.Scan(&c.BlobID, &accessTime, &c.StoredSize); err != nil { return nil, fmt.Errorf("rows.Scan: %w", err) } c.Age = time.Duration(nowUnix-accessTime) * time.Second @@ -1250,11 +1405,11 @@ func (srv *Server) deleteBlobs(blobIDs ...int64) error { var sumBytes int64 for _, blobID := range blobIDs { var sha256Hex string - var blobSize int64 - if err := tx.QueryRow("SELECT SHA256, BlobSize FROM Blobs WHERE BlobID = ?", blobID).Scan(&sha256Hex, &blobSize); err != nil && !errors.Is(err, sql.ErrNoRows) { + var storedSize int64 + if err := tx.QueryRow("SELECT SHA256, StoredSize FROM Blobs WHERE BlobID = ?", blobID).Scan(&sha256Hex, &storedSize); err != nil && !errors.Is(err, sql.ErrNoRows) { return fmt.Errorf("querying blob SHA256: %w", err) } - sumBytes += blobSize + sumBytes += storedSize if _, err := tx.Exec("DELETE FROM Blobs WHERE BlobID = ?", blobID); err != nil { return fmt.Errorf("deleting blob: %w", err) } @@ -1263,7 +1418,12 @@ func (srv *Server) deleteBlobs(blobIDs ...int64) error { } var hash [sha256.Size]byte if _, err := hex.Decode(hash[:], []byte(sha256Hex)); err == nil { - if err := os.Remove(srv.sha256Filepath(hash)); err != nil && !os.IsNotExist(err) { + base := srv.sha256Filepath(hash) + // Try removing both lz4 and plain paths; one or neither may exist. + if err := os.Remove(base + ".lz4"); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("removing disk file: %w", err) + } + if err := os.Remove(base); err != nil && !os.IsNotExist(err) { return fmt.Errorf("removing disk file: %w", err) } } @@ -1311,7 +1471,7 @@ func (srv *Server) cleanOldObjects(us *usageStats) (countAndSize, error) { var sumSize int64 for _, c := range candidates { blobIDs = append(blobIDs, c.BlobID) - sumSize += c.BlobSize + sumSize += c.StoredSize } if err := srv.deleteBlobs(blobIDs...); err != nil { return zero, fmt.Errorf("deleting old blobs: %v", err) @@ -1338,7 +1498,7 @@ func (srv *Server) cleanOldObjects(us *usageStats) (countAndSize, error) { } for _, c := range candidates { blobIDs = append(blobIDs, c.BlobID) - batchBytes += c.BlobSize + batchBytes += c.StoredSize if batchBytes >= toClean { break } @@ -1422,6 +1582,19 @@ func durFmt(d time.Duration) string { return d.String() } +// requestAcceptsEncoding reports whether r's Accept-Encoding header includes +// the given encoding. It handles optional quality parameters (e.g. +// "gzip;q=1.0, lz4;q=0.5") by stripping them before comparison. +func requestAcceptsEncoding(r *http.Request, encoding string) bool { + for part := range strings.SplitSeq(r.Header.Get("Accept-Encoding"), ",") { + name, _, _ := strings.Cut(strings.TrimSpace(part), ";") + if strings.EqualFold(strings.TrimSpace(name), encoding) { + return true + } + } + return false +} + func bytesFmt(n int64) string { if n >= 1<<30 { return fmt.Sprintf("%.1f GiB", float64(n)/(1<<30)) diff --git a/gocache/gocached/gocached_test.go b/gocache/gocached/gocached_test.go index 7ba9bc3..1d9f425 100644 --- a/gocache/gocached/gocached_test.go +++ b/gocache/gocached/gocached_test.go @@ -30,6 +30,7 @@ import ( "github.com/go-jose/go-jose/v4" "github.com/golang-jwt/jwt/v5" "github.com/google/go-cmp/cmp" + "github.com/pierrec/lz4/v4" "github.com/tailscale/tb/gocache/cachers" ) @@ -181,6 +182,23 @@ func (st *tester) wantGetMiss(c *cachers.HTTPClient, actionID string) { } } +// lz4Size returns the lz4-compressed size of data. +func lz4Size(t testing.TB, data []byte) int64 { + t.Helper() + var buf bytes.Buffer + w := lz4.NewWriter(&buf) + if err := w.Apply(lz4.SizeOption(uint64(len(data)))); err != nil { + t.Fatal(err) + } + if _, err := w.Write(data); err != nil { + t.Fatal(err) + } + if err := w.Close(); err != nil { + t.Fatal(err) + } + return int64(buf.Len()) +} + func withClock(clk func() time.Time) ServerOption { return func(cfg *Server) { cfg.clock = clk @@ -323,17 +341,19 @@ func TestServer(t *testing.T) { if err != nil { t.Fatalf("usageStats: %v", err) } + bigStored := int64(len(testObjectValueBig)) // below lz4CompressThreshold, stored uncompressed + totalSize := int64(9) + bigStored // 9 (inline "test data") + big + 0 (empty) want := &usageStats{ MissingBlobRows: 0, ActionsLE: map[time.Duration]countAndSize{ 24 * time.Hour: {Count: 1, Size: 9}, 48 * time.Hour: {Count: 1, Size: 9}, - 96 * time.Hour: {Count: 3, Size: 1034}, - 168 * time.Hour: {Count: 3, Size: 1034}, - 336 * time.Hour: {Count: 3, Size: 1034}, - 720 * time.Hour: {Count: 3, Size: 1034}, - 2160 * time.Hour: {Count: 3, Size: 1034}, - math.MaxInt64: {Count: 3, Size: 1034}, + 96 * time.Hour: {Count: 3, Size: totalSize}, + 168 * time.Hour: {Count: 3, Size: totalSize}, + 336 * time.Hour: {Count: 3, Size: totalSize}, + 720 * time.Hour: {Count: 3, Size: totalSize}, + 2160 * time.Hour: {Count: 3, Size: totalSize}, + math.MaxInt64: {Count: 3, Size: totalSize}, }, } if diff := cmp.Diff(stats, want); diff != "" { @@ -367,26 +387,26 @@ func TestCleanCandidates(t *testing.T) { maxAge: 0, limit: 100, want: []cleanCandidate{ - {BlobID: 1, Age: 3 * day, BlobSize: 1}, - {BlobID: 2, Age: 2 * day, BlobSize: 2}, - {BlobID: 3, Age: 1 * day, BlobSize: 3}, - {BlobID: 4, Age: 0, BlobSize: smallObjectSize + 1}, + {BlobID: 1, Age: 3 * day, StoredSize: 1}, + {BlobID: 2, Age: 2 * day, StoredSize: 2}, + {BlobID: 3, Age: 1 * day, StoredSize: 3}, + {BlobID: 4, Age: 0, StoredSize: smallObjectSize + 1}, // below lz4CompressThreshold, stored uncompressed }, }, { maxAge: 25 * time.Hour, limit: 100, want: []cleanCandidate{ - {BlobID: 1, Age: 3 * day, BlobSize: 1}, - {BlobID: 2, Age: 2 * day, BlobSize: 2}, + {BlobID: 1, Age: 3 * day, StoredSize: 1}, + {BlobID: 2, Age: 2 * day, StoredSize: 2}, }, }, { maxAge: 0, limit: 2, want: []cleanCandidate{ - {BlobID: 1, Age: 3 * day, BlobSize: 1}, - {BlobID: 2, Age: 2 * day, BlobSize: 2}, + {BlobID: 1, Age: 3 * day, StoredSize: 1}, + {BlobID: 2, Age: 2 * day, StoredSize: 2}, }, }, } @@ -416,17 +436,21 @@ func TestCleanOldObjectsByAge(t *testing.T) { st.wantPut(c1, "0003", "9903", "small") smallLen := int64(len("small")) + stored1 := int64(smallObjectSize + 1) // below lz4CompressThreshold, stored uncompressed + stored2 := lz4Size(t, []byte(strings.Repeat("x", smallObjectSize+2))) // at lz4CompressThreshold, lz4-compressed + st1 := st.usageStats() - if all, want := st1.All(), (countAndSize{Count: 3, Size: smallObjectSize*2 + 3 + smallLen}); all != want { + if all, want := st1.All(), (countAndSize{Count: 3, Size: stored1 + stored2 + smallLen}); all != want { t.Errorf("usageStats: %v; want %v", all, want) } - if got, want := st.diskFiles(), []string{"333092a3daf718ed8f38a94e302df139edd4e3b5da4239a497995683942cf28c", "c6d8e9905300876046729949cc95c2385221270d389176f7234fe7ac00c4e430"}; !slices.Equal(got, want) { + // First file is uncompressed (no .lz4 suffix), second is lz4-compressed. + if got, want := st.diskFiles(), []string{"333092a3daf718ed8f38a94e302df139edd4e3b5da4239a497995683942cf28c.lz4", "c6d8e9905300876046729949cc95c2385221270d389176f7234fe7ac00c4e430"}; !slices.Equal(got, want) { t.Errorf("diskFiles: %v; want %v", got, want) } clean1 := st.cleanOldObjects() - if clean1.Count != 1 || clean1.Size != smallObjectSize+1 { - t.Errorf("cleanOldObjects got %v, want {Count: 1, Size: %d}", clean1, smallObjectSize+1) + if clean1.Count != 1 || clean1.Size != stored1 { + t.Errorf("cleanOldObjects got %v, want {Count: 1, Size: %d}", clean1, stored1) } clean2 := st.cleanOldObjects() if clean2.Count != 0 || clean2.Size != 0 { @@ -434,10 +458,10 @@ func TestCleanOldObjectsByAge(t *testing.T) { } st2 := st.usageStats() - if all, want := st2.All(), (countAndSize{Count: 2, Size: smallObjectSize + 2 + smallLen}); all != want { + if all, want := st2.All(), (countAndSize{Count: 2, Size: stored2 + smallLen}); all != want { t.Errorf("usageStats after clean: %v; want %v", all, want) } - if got, want := st.diskFiles(), []string{"333092a3daf718ed8f38a94e302df139edd4e3b5da4239a497995683942cf28c"}; !slices.Equal(got, want) { + if got, want := st.diskFiles(), []string{"333092a3daf718ed8f38a94e302df139edd4e3b5da4239a497995683942cf28c.lz4"}; !slices.Equal(got, want) { t.Errorf("diskFiles after clean: %v; want %v", got, want) } } @@ -476,6 +500,120 @@ func TestCleanOldObjectsBySize(t *testing.T) { } } +func TestLZ4Storage(t *testing.T) { + st := newServerTester(t) + c := st.mkClient() + + type testCase struct { + name string + size int + wantLZ4 bool // expect .lz4 file on disk + wantDisk bool // expect any disk file (false = inline in DB) + } + tests := []testCase{ + {"empty", 0, false, false}, + {"tiny_1b", 1, false, false}, + {"inline_max", smallObjectSize, false, false}, + {"disk_no_lz4", smallObjectSize + 1, false, true}, // on disk but below lz4CompressThreshold + {"disk_at_lz4_threshold", lz4CompressThreshold, true, true}, // smallest lz4-compressed disk blob + {"disk_2k", 2048, true, true}, + {"disk_64k", 64 << 10, true, true}, + } + + for i, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + actionID := fmt.Sprintf("%04x", i+0x10) + outputID := fmt.Sprintf("%04x", i+0x90) + + data := make([]byte, tt.size) + if len(data) > 0 { + data[0] = 'X' + data[len(data)-1] = 'X' + } + val := string(data) + + // PUT + st.wantPut(c, actionID, outputID, val) + + // GET via client (sends Accept-Encoding: lz4) — verify round-trip. + c2 := st.mkClient() // fresh client, no disk cache + st.wantGet(c2, actionID, outputID, val) + + // Raw HTTP GET with Accept-Encoding: lz4 + req, _ := http.NewRequest("GET", st.hs.URL+"/action/"+actionID, nil) + req.Header.Set("Want-Object", "1") + req.Header.Set("Accept-Encoding", "lz4") + res, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("raw GET: %v", err) + } + body, _ := io.ReadAll(res.Body) + res.Body.Close() + + if tt.wantLZ4 { + if got := res.Header.Get("Content-Encoding"); got != "lz4" { + t.Errorf("Accept lz4: Content-Encoding = %q, want %q", got, "lz4") + } + if got := res.Header.Get("X-Uncompressed-Length"); got != fmt.Sprint(tt.size) { + t.Errorf("Accept lz4: X-Uncompressed-Length = %q, want %q", got, fmt.Sprint(tt.size)) + } + if got := res.Header.Get("Content-Length"); got == fmt.Sprint(tt.size) { + t.Errorf("Accept lz4: Content-Length = %q, should be compressed (smaller)", got) + } + } else { + if got := res.Header.Get("Content-Encoding"); got != "" { + t.Errorf("Accept lz4: Content-Encoding = %q, want empty", got) + } + if got := res.Header.Get("X-Uncompressed-Length"); got != "" { + t.Errorf("Accept lz4: X-Uncompressed-Length = %q, want empty", got) + } + // Uncompressed: body should be the raw data. + if string(body) != val { + t.Errorf("Accept lz4: body length = %d, want %d", len(body), len(val)) + } + } + + // Raw HTTP GET without Accept-Encoding: lz4 — server must decompress. + req2, _ := http.NewRequest("GET", st.hs.URL+"/action/"+actionID, nil) + req2.Header.Set("Want-Object", "1") + // Deliberately no Accept-Encoding. + res2, err := http.DefaultClient.Do(req2) + if err != nil { + t.Fatalf("raw GET (no lz4): %v", err) + } + body2, _ := io.ReadAll(res2.Body) + res2.Body.Close() + + if got := res2.Header.Get("Content-Encoding"); got != "" { + t.Errorf("No Accept lz4: Content-Encoding = %q, want empty", got) + } + if got := res2.Header.Get("Content-Length"); got != fmt.Sprint(tt.size) { + t.Errorf("No Accept lz4: Content-Length = %q, want %q", got, fmt.Sprint(tt.size)) + } + if string(body2) != val { + t.Errorf("No Accept lz4: body length = %d, want %d", len(body2), len(val)) + } + }) + } + + // Verify disk state: expect exactly one plain file (the disk_no_lz4 case) + // and the rest with .lz4 suffix. + var plain, compressed int + for _, f := range st.diskFiles() { + if strings.HasSuffix(f, ".lz4") { + compressed++ + } else { + plain++ + } + } + if plain != 1 { + t.Errorf("disk files: got %d plain (non-lz4), want 1", plain) + } + if compressed != 3 { + t.Errorf("disk files: got %d .lz4, want 3", compressed) + } +} + func TestClientConnReuse(t *testing.T) { st := newServerTester(t) From 50cdce254500bf4519ea2d6bb101f0a0a9ea5ea0 Mon Sep 17 00:00:00 2001 From: Tom Proctor Date: Mon, 16 Feb 2026 20:25:12 +0000 Subject: [PATCH 38/67] cachers: suppress 401 and 403 error logs for best-effort (bradfitz/go-tool-cache#31) The logs get very chatty in a couple of well-known scenarios if best-effort is enabled. At some point we'll fix those scenarios to behave better, but for now I want to unblock using this version of the client from tailscale/tailscale. Updates bradfitz/go-tool-cache#29 Signed-off-by: Tom Proctor Migrated-from: bradfitz/go-tool-cache@7bca4ab53cd417a96bd4f678e3339e15063df986 --- gocache/cachers/http.go | 26 ++++++++++++++++++++++++-- 1 file changed, 24 insertions(+), 2 deletions(-) diff --git a/gocache/cachers/http.go b/gocache/cachers/http.go index 19c08a1..331341f 100644 --- a/gocache/cachers/http.go +++ b/gocache/cachers/http.go @@ -120,7 +120,16 @@ func (c *HTTPClient) Get(ctx context.Context, actionID string) (outputID, diskPa } if res.StatusCode != http.StatusOK { msg := tryReadErrorMessage(res) - log.Printf("error GET /action/%s: %v, %s", actionID, res.Status, msg) + if c.BestEffortHTTP && res.StatusCode == http.StatusUnauthorized { + // Known error code that will repeatedly happen, so avoid + // filling the logs with these errors. + // + // 401: can happen when gocached restarts during a session, as it + // doesn't persist access tokens. + // TODO(tomhjp): make the client retry auth in the background. + } else { + log.Printf("error GET /action/%s: %v, %s", actionID, res.Status, msg) + } if c.BestEffortHTTP { return "", "", nil } @@ -228,7 +237,20 @@ func (c *HTTPClient) Put(ctx context.Context, actionID, outputID string, size in defer res.Body.Close() if res.StatusCode != http.StatusNoContent { msg := tryReadErrorMessage(res) - log.Printf("error PUT /%s/%s: %v, %s", actionID, outputID, res.Status, msg) + if c.BestEffortHTTP && (res.StatusCode == http.StatusUnauthorized || res.StatusCode == http.StatusForbidden) { + // Known error codes that will repeatedly happen, so avoid + // filling the logs with these errors. + // + // 401: can happen when gocached restarts during a session, as it + // doesn't persist access tokens. + // TODO(tomhjp): make the client retry auth in the background. + // + // 403: can happen when authed with a JWT that didn't grant global + // write permissions. + // TODO(tomhjp): support namespaces so all sessions can safely write. + } else { + log.Printf("error PUT /%s/%s: %v, %s", actionID, outputID, res.Status, msg) + } httpErr = fmt.Errorf("unexpected PUT /%s/%s status %v", actionID, outputID, res.Status) } } From 4a8e70bdb73079a8652c5f92f8794a6581e51e75 Mon Sep 17 00:00:00 2001 From: Tom Proctor Date: Wed, 18 Feb 2026 13:19:23 +0000 Subject: [PATCH 39/67] gocached: use prepared statement for atime updates (bradfitz/go-tool-cache#32) Now that our sqlite library supports optimisations for prepared statements, we can get about 3x speed-up for atime bumps. Running BenchmarkFlushAccessTimes on my laptop, I get: ``` Before: 38472702 ns/op After: 13989412 ns/op ``` Updates tailscale/corp#34696 (cherry picked from commit 4ecd4802deb1b68ac948b9fd340608d6978cca70) Signed-off-by: Tom Proctor Co-authored-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@8ef59dd13462906e936224ab4c3b3f978b5fa077 --- go.mod | 10 +++---- go.sum | 48 +++++++++++++++++-------------- gocache/gocached/gocached.go | 47 +++++++++++++++++++++++++----- gocache/gocached/gocached_test.go | 26 +++++++++++++++++ 4 files changed, 97 insertions(+), 34 deletions(-) diff --git a/go.mod b/go.mod index 6c0f0b9..ca41a97 100644 --- a/go.mod +++ b/go.mod @@ -8,7 +8,7 @@ require ( github.com/pierrec/lz4/v4 v4.1.25 github.com/prometheus/client_golang v1.23.0 github.com/prometheus/client_model v0.6.2 - modernc.org/sqlite v1.38.2 + modernc.org/sqlite v1.45.0 ) require ( @@ -18,14 +18,14 @@ require ( github.com/google/uuid v1.6.0 // indirect github.com/mattn/go-isatty v0.0.20 // indirect github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect - github.com/ncruces/go-strftime v0.1.9 // indirect + github.com/ncruces/go-strftime v1.0.0 // indirect github.com/prometheus/common v0.65.0 // indirect github.com/prometheus/procfs v0.16.1 // indirect github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect - golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b // indirect - golang.org/x/sys v0.34.0 // indirect + golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect + golang.org/x/sys v0.37.0 // indirect google.golang.org/protobuf v1.36.6 // indirect - modernc.org/libc v1.66.3 // indirect + modernc.org/libc v1.67.6 // indirect modernc.org/mathutil v1.7.1 // indirect modernc.org/memory v1.11.0 // indirect ) diff --git a/go.sum b/go.sum index e87ef0d..0c34f05 100644 --- a/go.sum +++ b/go.sum @@ -16,6 +16,8 @@ github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17k github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= +github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= github.com/klauspost/compress v1.18.0 h1:c/Cqfb0r+Yi+JtIEq73FWXVkRonBlf0CRNYc8Zttxdo= github.com/klauspost/compress v1.18.0/go.mod h1:2Pp+KzxcywXVXMr50+X0Q/Lsb43OQHYWRCY2AiWywWQ= github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc= @@ -24,8 +26,8 @@ github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWE github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= -github.com/ncruces/go-strftime v0.1.9 h1:bY0MQC28UADQmHmaF5dgpLmImcShSi2kHU9XLdhx/f4= -github.com/ncruces/go-strftime v0.1.9/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= +github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= github.com/pierrec/lz4/v4 v4.1.25 h1:kocOqRffaIbU5djlIBr7Wh+cx82C0vtFb0fOurZHqD0= github.com/pierrec/lz4/v4 v4.1.25/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= @@ -44,33 +46,35 @@ github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOf github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= -golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b h1:M2rDM6z3Fhozi9O7NWsxAkg/yqS/lQJ6PmkyIV3YP+o= -golang.org/x/exp v0.0.0-20250620022241-b7579e27df2b/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8= -golang.org/x/mod v0.25.0 h1:n7a+ZbQKQA/Ysbyb0/6IbB1H/X41mKgbhfv7AfG/44w= -golang.org/x/mod v0.25.0/go.mod h1:IXM97Txy2VM4PJ3gI61r1YEk/gAj6zAHN3AdZt6S9Ww= -golang.org/x/sync v0.15.0 h1:KWH3jNZsfyT6xfAfKiz6MRNmd46ByHDYaZ7KSkCtdW8= -golang.org/x/sync v0.15.0/go.mod h1:1dzgHSNfp02xaA81J2MS99Qcpr2w7fw1gpm99rleRqA= +golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY= +golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70= +golang.org/x/mod v0.29.0 h1:HV8lRxZC4l2cr3Zq1LvtOsi/ThTgWnUk/y64QSs8GwA= +golang.org/x/mod v0.29.0/go.mod h1:NyhrlYXJ2H4eJiRy/WDBO6HMqZQ6q9nk4JzS3NuCK+w= +golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug= +golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.34.0 h1:H5Y5sJ2L2JRdyv7ROF1he/lPdvFsd0mJHFw2ThKHxLA= -golang.org/x/sys v0.34.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= -golang.org/x/tools v0.34.0 h1:qIpSLOxeCYGg9TrcJokLBG4KFA6d795g0xkBkiESGlo= -golang.org/x/tools v0.34.0/go.mod h1:pAP9OwEaY1CAW3HOmg3hLZC5Z0CCmzjAF2UQMSqNARg= +golang.org/x/sys v0.37.0 h1:fdNQudmxPjkdUTPnLn5mdQv7Zwvbvpaxqs831goi9kQ= +golang.org/x/sys v0.37.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= +golang.org/x/tools v0.38.0 h1:Hx2Xv8hISq8Lm16jvBZ2VQf+RLmbd7wVUsALibYI/IQ= +golang.org/x/tools v0.38.0/go.mod h1:yEsQ/d/YK8cjh0L6rZlY8tgtlKiBNTL14pGDJPJpYQs= google.golang.org/protobuf v1.36.6 h1:z1NpPI8ku2WgiWnf+t9wTPsn6eP1L7ksHUlkfLvd9xY= google.golang.org/protobuf v1.36.6/go.mod h1:jduwjTPXsFjZGTmRluh+L6NjiWu7pchiJ2/5YcXBHnY= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= -modernc.org/cc/v4 v4.26.2 h1:991HMkLjJzYBIfha6ECZdjrIYz2/1ayr+FL8GN+CNzM= -modernc.org/cc/v4 v4.26.2/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0= -modernc.org/ccgo/v4 v4.28.0 h1:rjznn6WWehKq7dG4JtLRKxb52Ecv8OUGah8+Z/SfpNU= -modernc.org/ccgo/v4 v4.28.0/go.mod h1:JygV3+9AV6SmPhDasu4JgquwU81XAKLd3OKTUDNOiKE= -modernc.org/fileutil v1.3.8 h1:qtzNm7ED75pd1C7WgAGcK4edm4fvhtBsEiI/0NQ54YM= -modernc.org/fileutil v1.3.8/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc= +modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis= +modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0= +modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc= +modernc.org/ccgo/v4 v4.30.1/go.mod h1:bIOeI1JL54Utlxn+LwrFyjCx2n2RDiYEaJVSrgdrRfM= +modernc.org/fileutil v1.3.40 h1:ZGMswMNc9JOCrcrakF1HrvmergNLAmxOPjizirpfqBA= +modernc.org/fileutil v1.3.40/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc= modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI= modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito= +modernc.org/gc/v3 v3.1.1 h1:k8T3gkXWY9sEiytKhcgyiZ2L0DTyCQ/nvX+LoCljoRE= +modernc.org/gc/v3 v3.1.1/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY= modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks= modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI= -modernc.org/libc v1.66.3 h1:cfCbjTUcdsKyyZZfEUKfoHcP3S0Wkvz3jgSzByEWVCQ= -modernc.org/libc v1.66.3/go.mod h1:XD9zO8kt59cANKvHPXpx7yS2ELPheAey0vjIuZOhOU8= +modernc.org/libc v1.67.6 h1:eVOQvpModVLKOdT+LvBPjdQqfrZq+pC39BygcT+E7OI= +modernc.org/libc v1.67.6/go.mod h1:JAhxUVlolfYDErnwiqaLvUqc8nfb2r6S6slAgZOnaiE= modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= @@ -79,8 +83,8 @@ modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8= modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns= modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w= modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE= -modernc.org/sqlite v1.38.2 h1:Aclu7+tgjgcQVShZqim41Bbw9Cho0y/7WzYptXqkEek= -modernc.org/sqlite v1.38.2/go.mod h1:cPTJYSlgg3Sfg046yBShXENNtPrWrDX8bsbAQBzgQ5E= +modernc.org/sqlite v1.45.0 h1:r51cSGzKpbptxnby+EIIz5fop4VuE4qFoVEjNvWoObs= +modernc.org/sqlite v1.45.0/go.mod h1:CzbrU2lSB1DKUusvwGz7rqEKIq+NUd8GWuBBZDs9/nA= modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0= modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A= modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y= diff --git a/gocache/gocached/gocached.go b/gocache/gocached/gocached.go index c7d0670..8053eaa 100644 --- a/gocache/gocached/gocached.go +++ b/gocache/gocached/gocached.go @@ -53,6 +53,7 @@ import ( "os" "path/filepath" "reflect" + "runtime" "slices" "strings" "sync" @@ -148,8 +149,9 @@ func openDB(dbDir string) (*sql.DB, error) { if err != nil { return nil, err } - db.SetMaxOpenConns(4) - db.SetMaxIdleConns(4) + numConns := min(runtime.NumCPU(), 4) + db.SetMaxOpenConns(numConns) + db.SetMaxIdleConns(numConns) db.SetConnMaxLifetime(0) // no limit if _, err := db.Exec(schema); err != nil { return nil, err @@ -229,7 +231,7 @@ func migrateV3ToV4(dbDir, v3Path, v4Path string) error { // start initializes the server, including defaults and background goroutines. func (srv *Server) start() error { - srv.shutdownCtx, srv.shutdownCancel = context.WithCancel(context.Background()) + srv.shutdownCtx, srv.shutdownCancel = context.WithCancel(srv.shutdownCtx) if srv.dir == "" { d, err := os.UserCacheDir() if err != nil { @@ -423,7 +425,18 @@ func NewServer(opts ...ServerOption) (*Server, error) { // database. func (srv *Server) Close() error { srv.shutdownCancel() - return srv.db.Close() + srv.mu.Lock() + defer srv.mu.Unlock() + + var err error + if srv.updateAccessTimeStmt != nil { + err = srv.updateAccessTimeStmt.Close() + } + if srv.writeConn != nil { + err = errors.Join(err, srv.writeConn.Close()) + } + + return errors.Join(err, srv.db.Close()) } // Server implements a gocached server. Use [NewServer] to create and start a @@ -456,7 +469,9 @@ type Server struct { // laziness for now. // // Lock ordering: sqliteWriteMu before mu. - sqliteWriteMu sync.Mutex + sqliteWriteMu sync.Mutex + writeConn *sql.Conn // nil until first used; single connection for writes + updateAccessTimeStmt *sql.Stmt // nil until first used; for updating access times on writeConn lastUsage atomic.Pointer[usageStats] @@ -892,14 +907,32 @@ func (srv *Server) flushAccessTimeBumpsWithErr() (ret error) { return nil } - tx, err := srv.db.Begin() + if srv.writeConn == nil { + var err error + srv.writeConn, err = srv.db.Conn(context.Background()) + if err != nil { + return fmt.Errorf("getting writeConn: %w", err) + } + } + + if srv.updateAccessTimeStmt == nil { + var err error + srv.updateAccessTimeStmt, err = srv.writeConn.PrepareContext(context.Background(), "UPDATE Actions SET AccessTime = ? WHERE NamespaceID = ? AND ActionID = ?") + if err != nil { + return fmt.Errorf("prepare updateAccessTimeStmt: %w", err) + } + } + + tx, err := srv.writeConn.BeginTx(context.Background(), nil) if err != nil { return fmt.Errorf("begin Tx: %w", err) } defer tx.Rollback() + txStmt := tx.Stmt(srv.updateAccessTimeStmt) + for action, accessTime := range srv.accessDirty { - _, err := tx.Exec("UPDATE Actions SET AccessTime = ? WHERE NamespaceID = ? AND ActionID = ?", accessTime, action.NamespaceID, action.ActionID) + _, err := txStmt.Exec(accessTime, action.NamespaceID, action.ActionID) if err != nil { return fmt.Errorf("updating access time for ns=%d,action=%q: %w", action.NamespaceID, action.ActionID, err) } diff --git a/gocache/gocached/gocached_test.go b/gocache/gocached/gocached_test.go index 1d9f425..97e770d 100644 --- a/gocache/gocached/gocached_test.go +++ b/gocache/gocached/gocached_test.go @@ -836,3 +836,29 @@ func TestExchangeToken(t *testing.T) { }) } } + +func BenchmarkFlushAccessTimes(b *testing.B) { + st := newServerTester(b, WithVerbose(false)) + s := st.srv + + cl := st.mkClient() + var actionIDs []string + for n := range 5000 { + aid := fmt.Sprintf("abcd%04x", n) + st.wantPut(cl, aid, "def456", "data789") + actionIDs = append(actionIDs, aid) + } + + for b.Loop() { + s.mu.Lock() + s.accessDirty = make(map[actionKey]int64) + for _, aid := range actionIDs { + s.accessDirty[actionKey{ActionID: aid}] = 123 + } + s.mu.Unlock() + + if err := s.flushAccessTimeBumpsWithErr(); err != nil { + b.Fatalf("flushAccessTimeBumpsWithErr: %v", err) + } + } +} From e0d4185f1d8fe42293d51c679628c46cc325b040 Mon Sep 17 00:00:00 2001 From: Tom Proctor Date: Tue, 3 Mar 2026 20:50:42 +0000 Subject: [PATCH 40/67] gocached: allow multiple JWT issuers (bradfitz/go-tool-cache#33) To allow a single gocached server to be shared between clients in different contexts, support multiple JWT issuers. For example, it could be configured to support both GitHub identity tokens and AWS IAM outbound identity federation, with distinct claims in each case. As written, this commit breaks the gocached package API, but we're not releasing proper semantic versions of the library, so I haven't made efforts not to. I'm happy to receive feedback on that if it is going to cause issues for anyone though. Updates tailscale/corp#37839 Signed-off-by: Tom Proctor Migrated-from: bradfitz/go-tool-cache@75febc375a6f42646a3a3b959adb3ba11290a9cf --- cmd/gocached/gocached.go | 9 +- gocache/gocached/gocached.go | 95 ++++++++++------ gocache/gocached/gocached_test.go | 176 +++++++++++++++++++++++++++++- gocache/internal/jwt/jwt.go | 149 +++++++++++++++---------- gocache/logger/logger.go | 7 ++ 5 files changed, 341 insertions(+), 95 deletions(-) create mode 100644 gocache/logger/logger.go diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index d516622..8ffb86d 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -71,10 +71,11 @@ func main() { maps.Copy(globalClaims, jwtClaims) maps.Copy(globalClaims, globalJWTClaims) - opts = append(opts, - gocached.WithJWTAuth(*jwtIssuer, jwtClaims), - gocached.WithGlobalNamespaceJWTClaims(globalClaims), - ) + opts = append(opts, gocached.WithJWTAuth(gocached.JWTIssuerConfig{ + Issuer: *jwtIssuer, + RequiredClaims: jwtClaims, + GlobalWriteClaims: globalClaims, + })) } srv, err := gocached.NewServer(opts...) diff --git a/gocache/gocached/gocached.go b/gocache/gocached/gocached.go index 8053eaa..754793e 100644 --- a/gocache/gocached/gocached.go +++ b/gocache/gocached/gocached.go @@ -66,6 +66,7 @@ import ( "github.com/prometheus/client_golang/prometheus/promhttp" dto "github.com/prometheus/client_model/go" ijwt "github.com/tailscale/tb/gocache/internal/jwt" + "github.com/tailscale/tb/gocache/logger" _ "modernc.org/sqlite" ) @@ -274,13 +275,19 @@ func (srv *Server) start() error { srv.logf("gocached: cleaned %v", res) } - if srv.jwtIssuer != "" { - srv.jwtValidator = ijwt.NewJWTValidator(srv.logf, srv.jwtIssuer, gocachedAudience) + if len(srv.jwtIssuers) > 0 { + issuerURLs := make([]string, 0, len(srv.jwtIssuers)) + for iss := range srv.jwtIssuers { + issuerURLs = append(issuerURLs, iss) + } + srv.jwtValidator = ijwt.NewJWTValidator(srv.logf, gocachedAudience, issuerURLs) if err := srv.jwtValidator.RunUpdateJWKSLoop(srv.shutdownCtx); err != nil { return fmt.Errorf("failed to fetch JWKS for JWT validator: %w", err) } - srv.logf("gocached: using JWT issuer %q with claims %v, global claims %v", srv.jwtIssuer, srv.jwtClaims, srv.globalJWTClaims) + for iss, entry := range srv.jwtIssuers { + srv.logf("gocached: using JWT issuer %q with required claims %v, global write claims %v", iss, entry.requiredClaims, entry.globalWriteClaims) + } go srv.runCleanSessionsLoop() } @@ -351,11 +358,9 @@ func WithVerbose(verbose bool) ServerOption { } } -type logf func(format string, args ...any) - // WithLogf sets a custom logging function for the server. Defaults to // [log.Printf]. -func WithLogf(logf logf) ServerOption { +func WithLogf(logf logger.Logf) ServerOption { return func(srv *Server) { srv.logf = logf } @@ -378,24 +383,40 @@ func WithMaxAge(maxAge time.Duration) ServerOption { } } -// WithJWTAuth enables JWT-based authentication for the server. The issuer must -// be a reachable HTTP(S) server that serves its JWKS via a URL discoverable at -// /.well-known/openid-configuration, and any JWT presented to the server must -// exactly match the provided claims to start a session. No requests are allowed -// without authentication if JWT auth is enabled. -func WithJWTAuth(issuer string, claims map[string]string) ServerOption { - return func(srv *Server) { - srv.jwtIssuer = issuer - srv.jwtClaims = claims - } +// JWTIssuerConfig configures a single OIDC issuer for JWT-based authentication. +type JWTIssuerConfig struct { + // Issuer is the OIDC issuer URL. It must be a reachable HTTP(S) server + // that serves its JWKS via a URL discoverable at + // /.well-known/openid-configuration. + Issuer string + + // RequiredClaims are claims that any JWT from this issuer must have to + // start a session. All key-value pairs must match exactly. + RequiredClaims map[string]string + + // GlobalWriteClaims are claims that a JWT from this issuer must have to + // write to the cache's global namespace. It should be a superset of + // RequiredClaims. + GlobalWriteClaims map[string]string } -// WithGlobalNamespaceJWTClaims sets additional claims that a JWT must have to -// write to the cache's global namespace. It should be a superset of the claims -// provided to [WithJWTAuth]. -func WithGlobalNamespaceJWTClaims(claims map[string]string) ServerOption { +// WithJWTAuth enables JWT-based authentication for the server. Each issuer must +// be a reachable HTTP(S) server that serves its JWKS via a URL discoverable at +// /.well-known/openid-configuration, and any JWT presented to the server must +// exactly match the issuer's required claims to start a session. No requests are +// allowed without authentication if JWT auth is enabled. It can be called multiple +// times; configs accumulate. +func WithJWTAuth(issuers ...JWTIssuerConfig) ServerOption { return func(srv *Server) { - srv.globalJWTClaims = claims + if srv.jwtIssuers == nil { + srv.jwtIssuers = make(map[string]*jwtIssuerConfig) + } + for _, ic := range issuers { + srv.jwtIssuers[ic.Issuer] = &jwtIssuerConfig{ + requiredClaims: ic.RequiredClaims, + globalWriteClaims: ic.GlobalWriteClaims, + } + } } } @@ -445,7 +466,7 @@ type Server struct { db *sql.DB dir string // for SQLite DB + large blobs verbose bool - logf logf + logf logger.Logf clock func() time.Time // if non-nil, alternate time.Now for testing metricsHandler http.Handler maxSize int64 // maximum size of the cache in bytes; 0 means no limit @@ -453,10 +474,8 @@ type Server struct { shutdownCtx context.Context shutdownCancel context.CancelFunc - jwtValidator *ijwt.Validator // nil unless jwtIssuer is set - jwtIssuer string // issuer URL for JWTs - jwtClaims map[string]string // claims required for any JWT to start a session - globalJWTClaims map[string]string // additional claims required to write to global namespace + jwtValidator *ijwt.Validator // nil unless jwtIssuers is non-empty + jwtIssuers map[string]*jwtIssuerConfig // keyed by issuer URL mu sync.RWMutex // guards following fields in this block sessions map[string]*sessionData // maps access token -> session data. @@ -501,6 +520,12 @@ type Server struct { } } +// jwtIssuerConfig holds per-issuer claim requirements for JWT auth. +type jwtIssuerConfig struct { + requiredClaims map[string]string + globalWriteClaims map[string]string +} + // sessionData corresponds to a specific access token, and is only used if JWT // auth is enabled. type sessionData struct { @@ -1149,11 +1174,17 @@ func (srv *Server) handleTokenExchange(w http.ResponseWriter, r *http.Request) { } func (srv *Server) evaluateClaims(claims map[string]any) (globalNSWrite bool, _ error) { - if missing := findMissingClaims(srv.jwtClaims, claims); len(missing) > 0 { + iss, _ := claims["iss"].(string) + cfg, ok := srv.jwtIssuers[iss] + if !ok { + return false, fmt.Errorf("got claims %v; unknown issuer %q", claims, iss) + } + + if missing := findMissingClaims(cfg.requiredClaims, claims); len(missing) > 0 { return false, fmt.Errorf("got claims %v; missing required claims: %v", claims, missing) } - if missing := findMissingClaims(srv.globalJWTClaims, claims); len(missing) == 0 { + if missing := findMissingClaims(cfg.globalWriteClaims, claims); len(missing) == 0 { return true, nil } else if srv.verbose { srv.logf("token exchange: missing global namespace write claims: %v", missing) @@ -1702,9 +1733,11 @@ func (srv *Server) serveSessions(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/html; charset=utf-8") fmt.Fprintf(w, "

gocached sessions

\n") - fmt.Fprintf(w, "

JWT issuer: %s

\n", srv.jwtIssuer) - fmt.Fprintf(w, "

JWT claims required: %v

\n", srv.jwtClaims) - fmt.Fprintf(w, "

JWT global write claims required: %v

\n", srv.globalJWTClaims) + for iss, cfg := range srv.jwtIssuers { + fmt.Fprintf(w, "

JWT issuer: %s

\n", iss) + fmt.Fprintf(w, "

JWT claims required: %v

\n", cfg.requiredClaims) + fmt.Fprintf(w, "

JWT global write claims required: %v

\n", cfg.globalWriteClaims) + } fmt.Fprintf(w, "

Number of sessions: %d

\n", len(sessions)) fmt.Fprintf(w, "\n") diff --git a/gocache/gocached/gocached_test.go b/gocache/gocached/gocached_test.go index 97e770d..956f8d9 100644 --- a/gocache/gocached/gocached_test.go +++ b/gocache/gocached/gocached_test.go @@ -722,8 +722,11 @@ func TestExchangeToken(t *testing.T) { t.Run(name, func(t *testing.T) { issuer, createJWT := startOIDCServer(t, privateKey.Public()) st := newServerTester(t, - WithJWTAuth(issuer, wantClaims), - WithGlobalNamespaceJWTClaims(wantGlobalClaims), + WithJWTAuth(JWTIssuerConfig{ + Issuer: issuer, + RequiredClaims: wantClaims, + GlobalWriteClaims: wantGlobalClaims, + }), ) // Generate JWT. @@ -837,6 +840,175 @@ func TestExchangeToken(t *testing.T) { } } +func TestMultiIssuerAuth(t *testing.T) { + // Generate separate keys for each issuer. + keyA, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("error generating key A: %v", err) + } + keyB, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("error generating key B: %v", err) + } + keyC, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("error generating key C: %v", err) + } + + issuerA, createJWTA := startOIDCServer(t, keyA.Public()) + issuerB, createJWTB := startOIDCServer(t, keyB.Public()) + issuerC, createJWTC := startOIDCServer(t, keyC.Public()) + + st := newServerTester(t, + WithJWTAuth( + JWTIssuerConfig{ + Issuer: issuerA, + RequiredClaims: map[string]string{"sub": "userA"}, + GlobalWriteClaims: map[string]string{ + "sub": "userA", + "ref": "refs/heads/main", + }, + }, + JWTIssuerConfig{ + Issuer: issuerB, + RequiredClaims: map[string]string{"sub": "userB"}, + GlobalWriteClaims: map[string]string{ + "sub": "userB", + "ref": "refs/heads/main", + }, + }, + ), + ) + + makeJWTBody := func(jwtString string) []byte { + body, err := json.Marshal(map[string]any{"jwt": jwtString}) + if err != nil { + t.Fatalf("error marshaling request body: %v", err) + } + return body + } + + exchangeToken := func(jwtBody []byte) (*http.Response, string) { + t.Helper() + req, err := http.NewRequest("POST", st.hs.URL+"/auth/exchange-token", bytes.NewReader(jwtBody)) + if err != nil { + t.Fatalf("error creating request: %v", err) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("error making request: %v", err) + } + body, err := io.ReadAll(resp.Body) + resp.Body.Close() + if err != nil { + t.Fatalf("error reading response body: %v", err) + } + if resp.StatusCode == http.StatusOK { + var d struct { + AccessToken string `json:"access_token"` + } + if err := json.Unmarshal(body, &d); err != nil { + t.Fatalf("error decoding response body: %v", err) + } + return resp, d.AccessToken + } + return resp, "" + } + + baseClaims := func(iss string) jwt.MapClaims { + return jwt.MapClaims{ + "iss": iss, + "aud": gocachedAudience, + "nbf": jwt.NewNumericDate(time.Now().Add(-time.Minute)), + "exp": jwt.NewNumericDate(time.Now().Add(time.Hour)), + } + } + + // Issuer A: valid read-only token. + t.Run("issuerA_read", func(t *testing.T) { + claims := baseClaims(issuerA) + claims["sub"] = "userA" + resp, accessToken := exchangeToken(makeJWTBody(createJWTA(claims, keyA))) + if resp.StatusCode != http.StatusOK { + t.Fatalf("unexpected status code: want %d, got %d", http.StatusOK, resp.StatusCode) + } + if accessToken == "" { + t.Fatal("expected access token") + } + + cl := st.mkClient() + cl.AccessToken = accessToken + st.wantGetMiss(cl, "aabb01") + }) + + // Issuer A: valid write token. + t.Run("issuerA_write", func(t *testing.T) { + claims := baseClaims(issuerA) + claims["sub"] = "userA" + claims["ref"] = "refs/heads/main" + resp, accessToken := exchangeToken(makeJWTBody(createJWTA(claims, keyA))) + if resp.StatusCode != http.StatusOK { + t.Fatalf("unexpected status code: want %d, got %d", http.StatusOK, resp.StatusCode) + } + cl := st.mkClient() + cl.AccessToken = accessToken + st.wantPut(cl, "aabb02", "ccdd02", "hello-from-A") + st.wantGet(cl, "aabb02", "ccdd02", "hello-from-A") + }) + + // Issuer B: valid read-only token. + t.Run("issuerB_read", func(t *testing.T) { + claims := baseClaims(issuerB) + claims["sub"] = "userB" + resp, accessToken := exchangeToken(makeJWTBody(createJWTB(claims, keyB))) + if resp.StatusCode != http.StatusOK { + t.Fatalf("unexpected status code: want %d, got %d", http.StatusOK, resp.StatusCode) + } + if accessToken == "" { + t.Fatal("expected access token") + } + cl := st.mkClient() + cl.AccessToken = accessToken + // Can read data written by issuer A. + st.wantGet(cl, "aabb02", "ccdd02", "hello-from-A") + }) + + // Issuer B: valid write token. + t.Run("issuerB_write", func(t *testing.T) { + claims := baseClaims(issuerB) + claims["sub"] = "userB" + claims["ref"] = "refs/heads/main" + resp, accessToken := exchangeToken(makeJWTBody(createJWTB(claims, keyB))) + if resp.StatusCode != http.StatusOK { + t.Fatalf("unexpected status code: want %d, got %d", http.StatusOK, resp.StatusCode) + } + cl := st.mkClient() + cl.AccessToken = accessToken + st.wantPut(cl, "aabb03", "ccdd03", "hello-from-B") + st.wantGet(cl, "aabb03", "ccdd03", "hello-from-B") + }) + + // Issuer B: wrong required claims (sub doesn't match). + t.Run("issuerB_wrong_sub", func(t *testing.T) { + claims := baseClaims(issuerB) + claims["sub"] = "userA" // issuer B requires sub=userB + resp, _ := exchangeToken(makeJWTBody(createJWTB(claims, keyB))) + if resp.StatusCode != http.StatusUnauthorized { + t.Fatalf("unexpected status code: want %d, got %d", http.StatusUnauthorized, resp.StatusCode) + } + }) + + // Issuer C: not configured, should be rejected. + t.Run("issuerC_rejected", func(t *testing.T) { + claims := baseClaims(issuerC) + claims["sub"] = "userC" + resp, _ := exchangeToken(makeJWTBody(createJWTC(claims, keyC))) + if resp.StatusCode != http.StatusUnauthorized { + t.Fatalf("unexpected status code: want %d, got %d", http.StatusUnauthorized, resp.StatusCode) + } + }) +} + func BenchmarkFlushAccessTimes(b *testing.B) { st := newServerTester(b, WithVerbose(false)) s := st.srv diff --git a/gocache/internal/jwt/jwt.go b/gocache/internal/jwt/jwt.go index 31f13a0..ae68c27 100644 --- a/gocache/internal/jwt/jwt.go +++ b/gocache/internal/jwt/jwt.go @@ -6,6 +6,7 @@ package jwt import ( "context" "encoding/json" + "errors" "fmt" "io" "net/http" @@ -17,6 +18,7 @@ import ( "github.com/go-jose/go-jose/v4" "github.com/golang-jwt/jwt/v5" + "github.com/tailscale/tb/gocache/logger" ) const oidcConfigWellKnownPath string = "/.well-known/openid-configuration" @@ -26,30 +28,64 @@ var ( supportedAlgorithms = []string{"HS256", "RS256", "ES256"} ) +// issuer holds the per-issuer state: its own JWT parser and signing keys. +type issuer struct { + iss string + parser *jwt.Parser + signingKeys atomic.Value // []jose.JSONWebKey +} + +// keyFunc is how github.com/golang-jwt/jwt gets the public key it needs to +// verify a JWT signature. Each issuer entry only checks its own keys. +func (ie *issuer) keyFunc(t *jwt.Token) (any, error) { + var kid string + if v, ok := t.Header["kid"]; ok { + kid, _ = v.(string) + } + if kid == "" { + return nil, fmt.Errorf("no kid found in token header") + } + + signingKeys := ie.signingKeys.Load().([]jose.JSONWebKey) + for _, k := range signingKeys { + if k.KeyID == kid { + return k.Key, nil + } + } + + return nil, fmt.Errorf("unknown key ID: %s", kid) +} + // NewJWTValidator constructs a [Validator] for validating JWTs. Must call // [RunUpdateJWKSLoop] before validating any JWTs. Every JWT must exactly match -// the provided issuer and audience values in its "iss" and "aud" claims -// respectively. The issuer must be a reachable HTTP server that serves the JWT -// public signing keys via the path defined by [oidcConfigWellKnownPath], and -// the audience should be a value specific to the trust boundary that gocached -// resides within. -func NewJWTValidator(logf func(format string, args ...any), issuer, audience string) *Validator { +// one of the provided issuers and the audience value in its "iss" and "aud" +// claims respectively. Each issuer must be a reachable HTTP server that serves +// the JWT public signing keys via the path defined by [oidcConfigWellKnownPath], +// and the audience should be a value specific to the trust boundary that +// gocached resides within. +func NewJWTValidator(logf logger.Logf, audience string, issuerURLs []string) *Validator { + var issuers []*issuer + for _, iss := range issuerURLs { + issuers = append(issuers, &issuer{ + iss: iss, + parser: jwt.NewParser( + jwt.WithValidMethods(supportedAlgorithms), + jwt.WithIssuer(iss), + jwt.WithAudience(audience), + jwt.WithLeeway(10*time.Second), + jwt.WithIssuedAt(), + ), + }) + } return &Validator{ - logf: logf, - issuer: issuer, - parser: jwt.NewParser( - jwt.WithValidMethods(supportedAlgorithms), - jwt.WithIssuer(issuer), - jwt.WithAudience(audience), - jwt.WithLeeway(10*time.Second), - jwt.WithIssuedAt(), - ), + logf: logf, + issuers: issuers, } } // RunUpdateJWKSLoop fetches the JWKS synchronously once to surface any config // errors early, and then starts a background goroutine that periodically fetches -// the JWKS from the issuer to keep the signing keys up to date. Must be called +// the JWKS from all issuers to keep the signing keys up to date. Must be called // before validating any JWTs. func (v *Validator) RunUpdateJWKSLoop(ctx context.Context) error { // Initial fetch to error early on misconfiguration. @@ -65,56 +101,43 @@ func (v *Validator) RunUpdateJWKSLoop(ctx context.Context) error { // Validator provides methods for validating JWTs. Use [NewJWTValidator] to // construct a working Validator. type Validator struct { - logf func(format string, args ...any) - issuer string - parser *jwt.Parser - - signingKeys atomic.Value // []jose.JSONWebKey + logf logger.Logf + issuers []*issuer // TODO(tomhjp): metrics } // Validate returns an error if the provided JWT fails validation for an invalid -// signature or standard claim (iss, aud, iat, nbf, exp). It returns the token's -// verified claims if validation succeeds. The caller should then make policy -// decisions based on other claims such as "sub" or other custom claims. +// signature or standard claim (iss, aud, iat, nbf, exp). It tries each +// configured issuer and returns the verified claims from the first successful +// parse. If all issuers fail, it returns the last error. func (v *Validator) Validate(ctx context.Context, jwtString string) (map[string]any, error) { - tk, err := v.parser.Parse(jwtString, v.keyFunc) - if err != nil { - return nil, fmt.Errorf("failed to parse token: %w", err) - } - - if !tk.Valid { - return nil, fmt.Errorf("invalid token") - } + var lastErr error + for _, ie := range v.issuers { + tk, err := ie.parser.Parse(jwtString, ie.keyFunc) + if err != nil { + lastErr = err + continue + } - gotClaims, ok := tk.Claims.(jwt.MapClaims) - if !ok { - return nil, fmt.Errorf("unexpected claims type: %T", tk.Claims) - } + if !tk.Valid { + lastErr = fmt.Errorf("invalid token") + continue + } - return gotClaims, nil -} + gotClaims, ok := tk.Claims.(jwt.MapClaims) + if !ok { + lastErr = fmt.Errorf("unexpected claims type: %T", tk.Claims) + continue + } -// keyFunc is how github.com/golang-jwt/jwt gets the public key it needs to -// verify a JWT signature. -func (v *Validator) keyFunc(t *jwt.Token) (any, error) { - var kid string - if v, ok := t.Header["kid"]; ok { - kid, _ = v.(string) - } - if kid == "" { - return nil, fmt.Errorf("no kid found in token header") + return gotClaims, nil } - signingKeys := v.signingKeys.Load().([]jose.JSONWebKey) - for _, k := range signingKeys { - if k.KeyID == kid { - return k.Key, nil - } + if lastErr != nil { + return nil, fmt.Errorf("failed to parse token: %w", lastErr) } - - return nil, fmt.Errorf("unknown key ID: %s", kid) + return nil, fmt.Errorf("no issuers configured") } func (v *Validator) runUpdateJWKSLoop(ctx context.Context) { @@ -135,10 +158,20 @@ func (v *Validator) runUpdateJWKSLoop(ctx context.Context) { } func (v *Validator) updateJWKS(ctx context.Context) error { - v.logf("jwt: fetching JWKS from issuer %q", v.issuer) - u, err := url.Parse(v.issuer) + var errs []error + for _, ie := range v.issuers { + if err := ie.updateJWKS(ctx, v.logf); err != nil { + errs = append(errs, fmt.Errorf("issuer %q: %w", ie.iss, err)) + } + } + return errors.Join(errs...) +} + +func (ie *issuer) updateJWKS(ctx context.Context, logf logger.Logf) error { + logf("jwt: fetching JWKS from issuer %q", ie.iss) + u, err := url.Parse(ie.iss) if err != nil { - return fmt.Errorf("failed to parse issuer URL %q: %w", v.issuer, err) + return fmt.Errorf("failed to parse issuer URL %q: %w", ie.iss, err) } u.Path = path.Join(u.Path, oidcConfigWellKnownPath) @@ -174,7 +207,7 @@ func (v *Validator) updateJWKS(ctx context.Context) error { signingKeys = append(signingKeys, k) } - v.signingKeys.Store(signingKeys) + ie.signingKeys.Store(signingKeys) return nil } diff --git a/gocache/logger/logger.go b/gocache/logger/logger.go new file mode 100644 index 0000000..69b955e --- /dev/null +++ b/gocache/logger/logger.go @@ -0,0 +1,7 @@ +// Copyright (c) Tailscale Inc & AUTHORS +// SPDX-License-Identifier: BSD-3-Clause + +package logger + +// Logf is a logging function type. It is implemented by log.Printf. +type Logf func(format string, args ...any) From 3c5642505d29defcb6094fbb3a864846e3c85f2b Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Sun, 15 Mar 2026 04:40:10 +0000 Subject: [PATCH 41/67] cmd/gocached: add parentdeath support Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@9fc622d67f65e47033ff373807635ece5dedffce --- cmd/gocached/gocached.go | 7 +++++++ go.mod | 1 + go.sum | 2 ++ 3 files changed, 10 insertions(+) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index 8ffb86d..018c679 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -12,9 +12,11 @@ import ( "maps" "net" "net/http" + "os" "strings" "time" + "github.com/bradfitz/parentdeath" "github.com/tailscale/tb/gocache/gocached" _ "modernc.org/sqlite" ) @@ -55,6 +57,11 @@ func main() { flag.Var(&globalJWTClaims, "global-jwt-claim", "an additional claim in the form x=y that a JWT must have to allow writing to the cache's global namespace; may be specified more than once") flag.Parse() + parentdeath.Monitor(func() { + log.Printf("gocached: parent process died, exiting") + os.Exit(0) + }) + opts := []gocached.ServerOption{ gocached.WithDir(*dir), gocached.WithVerbose(*verbose), diff --git a/go.mod b/go.mod index ca41a97..e008369 100644 --- a/go.mod +++ b/go.mod @@ -2,6 +2,7 @@ module github.com/tailscale/tb go 1.24.0 require ( + github.com/bradfitz/parentdeath v0.0.0-20260315043412-764506aeb900 github.com/go-jose/go-jose/v4 v4.1.3 github.com/golang-jwt/jwt/v5 v5.3.0 github.com/google/go-cmp v0.7.0 diff --git a/go.sum b/go.sum index 0c34f05..b50d9a4 100644 --- a/go.sum +++ b/go.sum @@ -1,5 +1,7 @@ github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= +github.com/bradfitz/parentdeath v0.0.0-20260315043412-764506aeb900 h1:YTPrKVaBJlca4xkN+YU2VG3WHIDFme0vItEKuj0wrBQ= +github.com/bradfitz/parentdeath v0.0.0-20260315043412-764506aeb900/go.mod h1:nmAuQ8iUcGbqQsa847JnlFq6PQMMxpEeK0QCPAPYf3w= github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= From 9e96b495cb0e3bc39eb711ae71717a3aa2eda264 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Mon, 1 Jun 2026 21:57:29 +0000 Subject: [PATCH 42/67] gocached: bound WAL size and expose db/wal size gauges One of our cigocached servers' WAL had grown to 46 GB over ~87 days of uptime, with walFindFrame pegging CPU and stalling go-cacher clients on the CI runners. A CI vet job that normally runs in ~100s on a laptop was sitting at 30+ minutes. SQLite's autocheckpoint runs only PASSIVE checkpoints, which reuse WAL space in place but never shrink the file on disk. Fix: - Set journal_size_limit=1 GiB per connection so even passive checkpoints will truncate down to that bound when they can. - Run wal_checkpoint(TRUNCATE) once at startup (this is what shrinks the existing 46 GB WAL on the running server after deploy) and periodically (every minute) from a background goroutine, logging any partial checkpoints that hint at a long-running reader. - Run a final TRUNCATE checkpoint on Close. - Bump modernc.org/sqlite v1.45.0 -> v1.51.0 for general fixes accumulated since February. Also add two new gauges for ongoing visibility: gocached_sqlite_data_bytes size of the main .db file gocached_sqlite_wal_bytes size of the .db-wal file Updates tailscale/corp#42670 Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@097527a6aa10c3d9e1304bb0077375965440040e --- go.mod | 9 +- go.sum | 46 +++++----- gocache/gocached/gocached.go | 143 +++++++++++++++++++++++++++++- gocache/gocached/gocached_test.go | 61 +++++++++++++ 4 files changed, 229 insertions(+), 30 deletions(-) diff --git a/go.mod b/go.mod index e008369..36c75ef 100644 --- a/go.mod +++ b/go.mod @@ -1,5 +1,5 @@ module github.com/tailscale/tb -go 1.24.0 +go 1.25.0 require ( github.com/bradfitz/parentdeath v0.0.0-20260315043412-764506aeb900 @@ -9,7 +9,7 @@ require ( github.com/pierrec/lz4/v4 v4.1.25 github.com/prometheus/client_golang v1.23.0 github.com/prometheus/client_model v0.6.2 - modernc.org/sqlite v1.45.0 + modernc.org/sqlite v1.51.0 ) require ( @@ -23,10 +23,9 @@ require ( github.com/prometheus/common v0.65.0 // indirect github.com/prometheus/procfs v0.16.1 // indirect github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect - golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect - golang.org/x/sys v0.37.0 // indirect + golang.org/x/sys v0.42.0 // indirect google.golang.org/protobuf v1.36.6 // indirect - modernc.org/libc v1.67.6 // indirect + modernc.org/libc v1.72.3 // indirect modernc.org/mathutil v1.7.1 // indirect modernc.org/memory v1.11.0 // indirect ) diff --git a/go.sum b/go.sum index b50d9a4..26687c8 100644 --- a/go.sum +++ b/go.sum @@ -48,45 +48,43 @@ github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOf github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= -golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY= -golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70= -golang.org/x/mod v0.29.0 h1:HV8lRxZC4l2cr3Zq1LvtOsi/ThTgWnUk/y64QSs8GwA= -golang.org/x/mod v0.29.0/go.mod h1:NyhrlYXJ2H4eJiRy/WDBO6HMqZQ6q9nk4JzS3NuCK+w= -golang.org/x/sync v0.17.0 h1:l60nONMj9l5drqw6jlhIELNv9I0A4OFgRsG9k2oT9Ug= -golang.org/x/sync v0.17.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= +golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8= +golang.org/x/mod v0.33.0/go.mod h1:swjeQEj+6r7fODbD2cqrnje9PnziFuw4bmLbBZFrQ5w= +golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4= +golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.37.0 h1:fdNQudmxPjkdUTPnLn5mdQv7Zwvbvpaxqs831goi9kQ= -golang.org/x/sys v0.37.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= -golang.org/x/tools v0.38.0 h1:Hx2Xv8hISq8Lm16jvBZ2VQf+RLmbd7wVUsALibYI/IQ= -golang.org/x/tools v0.38.0/go.mod h1:yEsQ/d/YK8cjh0L6rZlY8tgtlKiBNTL14pGDJPJpYQs= +golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo= +golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k= +golang.org/x/tools v0.42.0/go.mod h1:Ma6lCIwGZvHK6XtgbswSoWroEkhugApmsXyrUmBhfr0= google.golang.org/protobuf v1.36.6 h1:z1NpPI8ku2WgiWnf+t9wTPsn6eP1L7ksHUlkfLvd9xY= google.golang.org/protobuf v1.36.6/go.mod h1:jduwjTPXsFjZGTmRluh+L6NjiWu7pchiJ2/5YcXBHnY= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= -modernc.org/cc/v4 v4.27.1 h1:9W30zRlYrefrDV2JE2O8VDtJ1yPGownxciz5rrbQZis= -modernc.org/cc/v4 v4.27.1/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0= -modernc.org/ccgo/v4 v4.30.1 h1:4r4U1J6Fhj98NKfSjnPUN7Ze2c6MnAdL0hWw6+LrJpc= -modernc.org/ccgo/v4 v4.30.1/go.mod h1:bIOeI1JL54Utlxn+LwrFyjCx2n2RDiYEaJVSrgdrRfM= -modernc.org/fileutil v1.3.40 h1:ZGMswMNc9JOCrcrakF1HrvmergNLAmxOPjizirpfqBA= -modernc.org/fileutil v1.3.40/go.mod h1:HxmghZSZVAz/LXcMNwZPA/DRrQZEVP9VX0V4LQGQFOc= +modernc.org/cc/v4 v4.28.2 h1:3tQ0lf2ADtoby2EtSP+J7IE2SHwEJdP8ioR59wx7XpY= +modernc.org/cc/v4 v4.28.2/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI= +modernc.org/ccgo/v4 v4.34.0 h1:yRLPFZieg532OT4rp4JFNIVcquwalMX26G95WQDqwCQ= +modernc.org/ccgo/v4 v4.34.0/go.mod h1:AS5WYMyBakQ+fhsHhtP8mWB82KTGPkNNJDGfGQCe0/A= +modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM= +modernc.org/fileutil v1.4.0/go.mod h1:EqdKFDxiByqxLk8ozOxObDSfcVOv/54xDs/DUHdvCUU= modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI= modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito= -modernc.org/gc/v3 v3.1.1 h1:k8T3gkXWY9sEiytKhcgyiZ2L0DTyCQ/nvX+LoCljoRE= -modernc.org/gc/v3 v3.1.1/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY= +modernc.org/gc/v3 v3.1.2 h1:ZtDCnhonXSZexk/AYsegNRV1lJGgaNZJuKjJSWKyEqo= +modernc.org/gc/v3 v3.1.2/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY= modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks= modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI= -modernc.org/libc v1.67.6 h1:eVOQvpModVLKOdT+LvBPjdQqfrZq+pC39BygcT+E7OI= -modernc.org/libc v1.67.6/go.mod h1:JAhxUVlolfYDErnwiqaLvUqc8nfb2r6S6slAgZOnaiE= +modernc.org/libc v1.72.3 h1:ZnDF4tXn4NBXFutMMQC4vtbTFSXhhKzR73fv0beZEAU= +modernc.org/libc v1.72.3/go.mod h1:dn0dZNnnn1clLyvRxLxYExxiKRZIRENOfqQ8XEeg4Qs= modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= -modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8= -modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns= +modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg= +modernc.org/opt v0.2.0/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns= modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w= modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE= -modernc.org/sqlite v1.45.0 h1:r51cSGzKpbptxnby+EIIz5fop4VuE4qFoVEjNvWoObs= -modernc.org/sqlite v1.45.0/go.mod h1:CzbrU2lSB1DKUusvwGz7rqEKIq+NUd8GWuBBZDs9/nA= +modernc.org/sqlite v1.51.0 h1:aH/MMSoayAIhozZ7uJbVTT9QO/VhzBf0J9tymmmuC/U= +modernc.org/sqlite v1.51.0/go.mod h1:tcNzv5p84E0skkmJn038y+hWJbLQXQqEnQfeh5r2JLM= modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0= modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A= modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y= diff --git a/gocache/gocached/gocached.go b/gocache/gocached/gocached.go index 754793e..a5d34c6 100644 --- a/gocache/gocached/gocached.go +++ b/gocache/gocached/gocached.go @@ -92,6 +92,27 @@ const ( const schemaVersion = 4 +// walJournalSizeLimit caps the on-disk WAL size after a successful +// checkpoint. It is set as a per-connection PRAGMA so that even SQLite's +// built-in PASSIVE autocheckpoint, which would otherwise leave the file at +// its high-water mark, truncates the WAL down to this size. Without it the +// WAL can grow without bound under continuous write traffic and never +// shrink, even when frames are being checkpointed in place. +const walJournalSizeLimit = 1 << 30 // 1 GiB + +// checkpointInterval is how often [Server.runCheckpointLoop] runs a TRUNCATE +// checkpoint in the background. SQLite's autocheckpoint only runs PASSIVE +// checkpoints (which reuse WAL space in place but never shrink the file on +// disk past walJournalSizeLimit); a periodic explicit TRUNCATE is what +// actually keeps the file small in steady state. +const checkpointInterval = time.Minute + +// dbSizeMetricsInterval is how often [Server.runDBSizeMetricsLoop] re-stats +// the SQLite files to update the size gauges. It is intentionally shorter +// than checkpointInterval so the gauge sees the WAL grow between checkpoints, +// not just snap back to zero each minute. +const dbSizeMetricsInterval = 15 * time.Second + const schema = ` PRAGMA journal_mode=WAL; CREATE TABLE IF NOT EXISTS Actions ( @@ -146,7 +167,9 @@ func openDB(dbDir string) (*sql.DB, error) { } } - db, err := sql.Open("sqlite", "file:"+dbPath+"?_pragma=busy_timeout(5000)") + dsn := fmt.Sprintf("file:%s?_pragma=busy_timeout(5000)&_pragma=journal_size_limit(%d)", + dbPath, walJournalSizeLimit) + db, err := sql.Open("sqlite", dsn) if err != nil { return nil, err } @@ -250,6 +273,18 @@ func (srv *Server) start() error { } srv.db = db + // Run a TRUNCATE checkpoint up front, before any other reader can pin a + // snapshot. If the WAL on disk is large (e.g. from a prior version of + // gocached that lacked the periodic checkpointer), this is what actually + // shrinks it. + ckCtx, ckCancel := context.WithTimeout(srv.shutdownCtx, 2*time.Minute) + if busy, log, ckpt, err := srv.checkpointTruncate(ckCtx); err != nil { + srv.logf("startup wal_checkpoint(TRUNCATE) error: %v", err) + } else { + srv.logf("startup wal_checkpoint(TRUNCATE): busy=%d log=%d ckpt=%d", busy, log, ckpt) + } + ckCancel() + reg := prometheus.NewRegistry() reg.MustRegister( collectors.NewGoCollector(), @@ -293,6 +328,8 @@ func (srv *Server) start() error { } go srv.runCleanLoop() + go srv.runCheckpointLoop() + go srv.runDBSizeMetricsLoop() return nil } @@ -457,6 +494,15 @@ func (srv *Server) Close() error { err = errors.Join(err, srv.writeConn.Close()) } + // Final TRUNCATE checkpoint so the WAL doesn't linger on disk past + // shutdown. Use context.Background because srv.shutdownCtx has already + // been canceled. + ckCtx, ckCancel := context.WithTimeout(context.Background(), 2*time.Minute) + if _, _, _, ckErr := srv.checkpointTruncate(ckCtx); ckErr != nil { + err = errors.Join(err, fmt.Errorf("final wal_checkpoint: %w", ckErr)) + } + ckCancel() + return errors.Join(err, srv.db.Close()) } @@ -517,6 +563,9 @@ type Server struct { Sessions expvar.Int `type:"gauge" name:"sessions" help:"number of active authenticated sessions"` Auths expvar.Int `type:"counter" name:"auth_attempts" help:"number of successful token exchanges"` AuthErrs expvar.Int `type:"counter" name:"auth_errs" help:"number of failed token exchanges"` + + SQLiteDataBytes expvar.Int `type:"gauge" name:"sqlite_data_bytes" help:"size in bytes of the SQLite main database file on disk"` + SQLiteWALBytes expvar.Int `type:"gauge" name:"sqlite_wal_bytes" help:"size in bytes of the SQLite WAL file on disk; should stay bounded near walJournalSizeLimit"` } } @@ -1587,6 +1636,98 @@ func (srv *Server) cleanOldObjects(us *usageStats) (countAndSize, error) { return ret, nil } +// checkpointTruncate runs PRAGMA wal_checkpoint(TRUNCATE) and returns SQLite's +// three result columns. A fully-applied checkpoint returns busy=0 and +// logFrames==ckptFrames; otherwise some frames remain in the WAL because of a +// concurrent reader pinning an older snapshot. +func (srv *Server) checkpointTruncate(ctx context.Context) (busy, logFrames, ckptFrames int, err error) { + err = srv.db.QueryRowContext(ctx, "PRAGMA wal_checkpoint(TRUNCATE)").Scan(&busy, &logFrames, &ckptFrames) + return busy, logFrames, ckptFrames, err +} + +// runCheckpointLoop periodically runs a TRUNCATE checkpoint to keep the WAL +// bounded on disk. SQLite's autocheckpoint only runs PASSIVE checkpoints, which +// reuse WAL space in place but never shrink the file; without this loop the WAL +// can grow without bound under continuous traffic. +func (srv *Server) runCheckpointLoop() { + t := time.NewTicker(checkpointInterval) + defer t.Stop() + for { + select { + case <-srv.shutdownCtx.Done(): + return + case <-t.C: + } + ctx, cancel := context.WithTimeout(srv.shutdownCtx, 2*time.Minute) + busy, logFrames, ckptFrames, err := srv.checkpointTruncate(ctx) + cancel() + if err != nil { + if errors.Is(err, context.Canceled) { + return + } + srv.logf("wal_checkpoint(TRUNCATE) error: %v", err) + continue + } + if busy != 0 || logFrames != ckptFrames { + // A reader is pinning frames; we'll catch up next tick. Logged + // because persistent partial checkpoints mean walJournalSizeLimit + // is the only thing keeping the file bounded, and we'd want to + // investigate. + srv.logf("wal_checkpoint(TRUNCATE) partial: busy=%d log=%d ckpt=%d", busy, logFrames, ckptFrames) + } else if srv.verbose { + srv.logf("wal_checkpoint(TRUNCATE): log=%d ckpt=%d", logFrames, ckptFrames) + } + // Refresh the size gauges immediately so dashboards see the + // post-truncate values without waiting for the next sampler tick. + srv.updateDBSizeMetrics() + } +} + +// dbPath returns the on-disk path of the SQLite main database file. +// The WAL file is at dbPath() + "-wal". +func (srv *Server) dbPath() string { + return filepath.Join(srv.dir, fmt.Sprintf("gocached-v%d.db", schemaVersion)) +} + +// updateDBSizeMetrics re-stats the SQLite files and updates the size gauges. +// A missing WAL file (e.g. on a fresh DB before the first write flushes) is +// reported as zero bytes. Other stat errors are logged but don't update the +// gauge, so a transient filesystem hiccup leaves the last-known value visible. +func (srv *Server) updateDBSizeMetrics() { + dbPath := srv.dbPath() + if fi, err := os.Stat(dbPath); err == nil { + srv.m.SQLiteDataBytes.Set(fi.Size()) + } else { + srv.logf("stat %s: %v", dbPath, err) + } + walPath := dbPath + "-wal" + switch fi, err := os.Stat(walPath); { + case err == nil: + srv.m.SQLiteWALBytes.Set(fi.Size()) + case errors.Is(err, os.ErrNotExist): + srv.m.SQLiteWALBytes.Set(0) + default: + srv.logf("stat %s: %v", walPath, err) + } +} + +// runDBSizeMetricsLoop samples the SQLite file sizes more frequently than the +// checkpoint loop runs, so the WAL gauge captures inter-checkpoint growth +// rather than only the post-truncate values. +func (srv *Server) runDBSizeMetricsLoop() { + t := time.NewTicker(dbSizeMetricsInterval) + defer t.Stop() + srv.updateDBSizeMetrics() // seed an initial sample at startup + for { + select { + case <-srv.shutdownCtx.Done(): + return + case <-t.C: + } + srv.updateDBSizeMetrics() + } +} + func (srv *Server) runCleanLoop() { for { select { diff --git a/gocache/gocached/gocached_test.go b/gocache/gocached/gocached_test.go index 956f8d9..0aad646 100644 --- a/gocache/gocached/gocached_test.go +++ b/gocache/gocached/gocached_test.go @@ -614,6 +614,67 @@ func TestLZ4Storage(t *testing.T) { } } +func TestWALCheckpoint(t *testing.T) { + st := newServerTester(t) + ctx := context.Background() + + walPath := filepath.Join(st.srv.dir, fmt.Sprintf("gocached-v%d.db-wal", schemaVersion)) + + // Generate enough write traffic for the WAL to be a few pages large. + // The exact number isn't important; we just want it big enough that + // "shrank to nearly empty" is a meaningful observation. + for i := range 500 { + if _, err := st.srv.db.ExecContext(ctx, + `INSERT INTO Actions (NamespaceID, ActionID, BlobID, AltOutputID, CreateTime, AccessTime) VALUES (0, ?, 0, '', 0, 0)`, + fmt.Sprintf("%032x", i)); err != nil { + t.Fatalf("insert %d: %v", i, err) + } + } + + before, err := os.Stat(walPath) + if err != nil { + t.Fatalf("stat WAL: %v", err) + } + t.Logf("WAL before checkpoint: %d bytes", before.Size()) + if before.Size() < 4096 { + // 500 row inserts should leave at least one full WAL page even after + // the periodic autocheckpoint reuses space in place. A tiny WAL here + // means the test isn't actually exercising the truncation path. + t.Fatalf("WAL only %d bytes before checkpoint; expected meaningful traffic", before.Size()) + } + + st.srv.updateDBSizeMetrics() + if got := st.srv.m.SQLiteWALBytes.Value(); got != before.Size() { + t.Errorf("sqlite_wal_bytes gauge before checkpoint: got %d, want %d", got, before.Size()) + } + if got := st.srv.m.SQLiteDataBytes.Value(); got <= 0 { + t.Errorf("sqlite_data_bytes gauge: got %d, want > 0", got) + } + + busy, log, ckpt, err := st.srv.checkpointTruncate(ctx) + if err != nil { + t.Fatalf("checkpointTruncate: %v", err) + } + if busy != 0 || log != ckpt { + t.Errorf("checkpoint not fully applied: busy=%d log=%d ckpt=%d", busy, log, ckpt) + } + + after, err := os.Stat(walPath) + if err != nil { + t.Fatalf("stat WAL after checkpoint: %v", err) + } + t.Logf("WAL after checkpoint: %d bytes", after.Size()) + // TRUNCATE checkpoint with no concurrent readers truncates the WAL to 0. + if after.Size() != 0 { + t.Errorf("WAL not truncated: got %d bytes, want 0", after.Size()) + } + + st.srv.updateDBSizeMetrics() + if got := st.srv.m.SQLiteWALBytes.Value(); got != 0 { + t.Errorf("sqlite_wal_bytes gauge after checkpoint: got %d, want 0", got) + } +} + func TestClientConnReuse(t *testing.T) { st := newServerTester(t) From a4f4925dcd1c1ca610537f524cae52ffaa3db8d0 Mon Sep 17 00:00:00 2001 From: Brad Fitzpatrick Date: Tue, 2 Jun 2026 14:17:42 +0000 Subject: [PATCH 43/67] gocached: do per-shard usage stats and incremental cleanup MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The usageStats() LEFT JOIN every 5 minutes didn't scale. At 272M Actions it took ~7 minutes per call, pinning a reader snapshot ~60% of the time. The WAL grew to 52 GB because SQLite's PASSIVE autocheckpoint could never advance past the snapshot; walFindFrame pegged CPU and go-cacher clients stalled. Startup also blocked on usageStats + cleanOldObjects for 10+ minutes. cleanCandidates was the same shape: GROUP BY scaling with total rows. This redoes the usage stats + LRU cleanup pipeline. Shard the histogram, persist it, refresh in the background: * New BlobShardStats(Prefix, ScannedAt, StatsJSON) table. * 256 shards keyed on a SHA256 hex prefix (--shard-prefix-len=1..4). * runShardStatsLoop picks the oldest shard, sleeps until it crosses shardStalenessTarget (10m), then rescans. * Per-shard scan is one SUM(CASE WHEN…) query scoped by SHA256 >= ? AND SHA256 < ?. No row streaming. * On restart, loadShardStats seeds lastUsage from the persisted rows. Server starts up to usable right away without a blocking step. Action-LRU cleanup driven by idx_actions_access: * evictOldestActions walks the access-time index for the N oldest stale Actions, deletes them, orphan-deletes Blobs. * INDEXED BY locks the plan; TestEvictionQueryPlan{,_atScale} asserts via EXPLAIN QUERY PLAN. * Action-LRU instead of Blob-LRU: a stale Action is evicted even if its Blob has fresher siblings. Equivalent for the common 1:1 case; stricter LRU on shared content. * One short Tx per 200-row batch on a 1s tick. No multi-minute reader snapshots. Dead-reckon PUTs and evictions between scans: * Per-shard mutex-guarded counters track (count, bytes) added or removed since the shard was last persisted. * cleanupTick includes the delta in the size pressure check so a burst of PUTs between scans triggers cleanup immediately. Observability: /usage shows cohort table, shard freshness histogram, per-shard scan duration p25/p50/p90; new gauges gocached_shard_stats_{unscanned,oldest_age_seconds}, pending blob count/bytes; new histogram gocached_shard_scan_duration_seconds replayed from persisted data at startup; new counter gocached_evicted_actions. Schema migration is additive (schemaVersion stays 4). Perf test (TestPerfQueries) always runs but defaults to 100 rows so normal `go test` is sub-second. Operators set --performance-test-rows larger for real scale; seed reused across runs via --performance-test-dir. Sample, cumulative across rows: rows DB scanShard evict legacy GROUP BY ratio usageStats 1K 396K 188µs 5.2ms 648µs 0.1x 321ns 10K 3.4M 248µs 5.5ms 6.4ms 1.2x 227ns 100K 35M 1.8ms 6.1ms 73ms 12x 328ns 1M 349M 22ms 7.5ms 763ms 102x 231ns 2M 698M 51ms 8.2ms 1.6s 196x 261ns 5M 1.8G 162ms 8.7ms 4.1s 475x 224ns 10M 3.5G 349ms 9.0ms 8.3s 922x 249ns usageStats and evictOldestActions are flat in N. scanShard scales with per-shard rows. The "legacy GROUP BY" column is the cleanCandidates query gocached used before this branch (since deleted); it scales linearly with total rows, the behavior that was wedging prod at 250M. Updates tailscale/corp#42670 Signed-off-by: Brad Fitzpatrick Migrated-from: bradfitz/go-tool-cache@d48e363773c3d03afaf3f7debe3721a30610d647 --- cmd/gocached/gocached.go | 3 + gocache/gocached/gocached.go | 1105 +++++++++++++++++++----- gocache/gocached/gocached_perf_test.go | 275 ++++++ gocache/gocached/gocached_test.go | 564 ++++++++++-- 4 files changed, 1667 insertions(+), 280 deletions(-) create mode 100644 gocache/gocached/gocached_perf_test.go diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index 018c679..79a7cfb 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -30,6 +30,8 @@ var ( maxSize = flag.Int("max-size-gb", 50, "maximum size of the cache in GiB; 0 means no limit") maxAge = flag.Int("max-age-days", 60, "maximum age of objects in the cache in days; 0 means no limit") + shardPrefixLen = flag.Int("shard-prefix-len", 2, "number of SHA256 hex characters per usage-stats shard key (valid 1..4); total shards = 16^n. Defaults to 2 (256 shards). Increase if a single shard's stats scan becomes too slow on a large DB") + jwtIssuer = flag.String("jwt-issuer", "", "the issuer to trust JWTs from; if set, all requests will require auth, and must set at least one -jwt-claim") // See example GitHub token claims for what can be available: // https://docs.github.com/en/actions/concepts/security/openid-connect @@ -67,6 +69,7 @@ func main() { gocached.WithVerbose(*verbose), gocached.WithMaxSize(int64(*maxSize) << 30), gocached.WithMaxAge(time.Duration(*maxAge) * 24 * time.Hour), + gocached.WithShardPrefixLen(*shardPrefixLen), } if *jwtIssuer != "" { diff --git a/gocache/gocached/gocached.go b/gocache/gocached/gocached.go index a5d34c6..156f11e 100644 --- a/gocache/gocached/gocached.go +++ b/gocache/gocached/gocached.go @@ -55,6 +55,7 @@ import ( "reflect" "runtime" "slices" + "strconv" "strings" "sync" "sync/atomic" @@ -113,6 +114,48 @@ const checkpointInterval = time.Minute // not just snap back to zero each minute. const dbSizeMetricsInterval = 15 * time.Second +// defaultShardPrefixLen is the [Server.shardPrefixLen] used when the option +// is not set (zero value). 2 yields 256 shards ("00".."ff"), which keeps +// individual scans fast even at hundreds of millions of rows. +const defaultShardPrefixLen = 2 + +// minShardPrefixLen is the minimum [Server.shardPrefixLen]. 1 yields 16 +// shards. Going to 0 (1 shard) would defeat the purpose of sharded +// statistics, so we reject it. +const minShardPrefixLen = 1 + +// maxShardPrefixLen is the maximum [Server.shardPrefixLen]. 4 hex chars -> +// 65,536 shards is more than we ever expect to need; higher values would +// make the stats loop spend most of its time on shards with little or no +// data, and would also blow up the size of the BlobShardStats table. +const maxShardPrefixLen = 4 + +// shardStatsMinInterval sets a defensive floor on how long [Server.runShardStatsLoop] will +// sleep between iterations to ensure it never spins in a tight loop. [shardStalenessTarget] +// dictates the interval in steady state. +const shardStatsMinInterval = 2 * time.Second + +// cleanupTickInterval is how often [Server.runCleanLoop] wakes up to check +// whether the cache is over its limits and, if so, evict one batch of +// Actions. It's tight on purpose: each tick does at most cleanupBatchSize +// deletes so the write transaction is short. +const cleanupTickInterval = 1 * time.Second + +// cleanupBatchSize is the maximum number of Actions deleted in a single +// [Server.evictOldestActions] call. Small enough that the write +// transaction (including any file removals) holds the SQLite write lock +// for only a brief moment, large enough that steady-state eviction keeps +// up with realistic write traffic at ~cleanupBatchSize/cleanupTickInterval. +const cleanupBatchSize = 200 + +// shardStalenessTarget bounds how old any single shard's cohort histogram is +// allowed to get. The stats loop sleeps until the oldest shard reaches this +// age, then scans it; on a quiet cache that means scans naturally settle to +// ~numShards/shardStalenessTarget per second instead of running continuously. +// 10m drift on the 1d cohort is well under 1% error, which is negligible +// compared to the day-granularity cohort cutoffs we record. +const shardStalenessTarget = 10 * time.Minute + const schema = ` PRAGMA journal_mode=WAL; CREATE TABLE IF NOT EXISTS Actions ( @@ -150,6 +193,16 @@ CREATE TABLE IF NOT EXISTS Namespaces ( NamespaceID INTEGER PRIMARY KEY AUTOINCREMENT, Namespace TEXT NOT NULL UNIQUE CHECK (Namespace = lower(Namespace)) ) STRICT; + +-- BlobShardStats persists per-shard usage histograms across restarts so +-- the server can serve PUTs with accurate global usage from right at startup, +-- without blocking for minutes computing stats before serving can begin. +-- There's one row per SHA256 prefix shard, refreshes one at a time in the background. +CREATE TABLE IF NOT EXISTS BlobShardStats ( + Prefix TEXT PRIMARY KEY, -- "01".."ff" by default for shardPrefixLen=2; must be lowercase hex + ScannedAt INTEGER NOT NULL, -- unix seconds + StatsJSON TEXT NOT NULL -- JSON-encoded [usageStats] of that shard +) STRICT; ` func openDB(dbDir string) (*sql.DB, error) { @@ -285,6 +338,33 @@ func (srv *Server) start() error { } ckCancel() + if srv.shardPrefixLen == 0 { + srv.shardPrefixLen = defaultShardPrefixLen + } + if srv.shardPrefixLen < minShardPrefixLen || srv.shardPrefixLen > maxShardPrefixLen { + return fmt.Errorf("shardPrefixLen %d out of range [%d, %d]", + srv.shardPrefixLen, minShardPrefixLen, maxShardPrefixLen) + } + + // Compute the cohort cutoffs once so every shard scan agrees. maxAge is + // added to the set if it's not already a standard duration so cleanup can + // read its bucket directly out of the aggregate. + srv.durs = computeDurs(srv.maxAge) + + // Allocate the dead-reckoning delta state. One entry per shard, indexed + // by the integer value of the SHA256 hex prefix. + srv.shardDeltas = make([]shardDelta, srv.numShards()) + + // Create the shardScanDuration histogram before loadShardStats so any + // QueryDuration values persisted by a previous run get replayed into + // the histogram immediately on restart, rather than waiting for the + // stats loop to repopulate it from fresh scans. + srv.shardScanDuration = prometheus.NewHistogram(prometheus.HistogramOpts{ + Name: "gocached_shard_scan_duration_seconds", + Help: "wall time of each per-shard usage-stats SQL scan", + Buckets: prometheus.DefBuckets, + }) + reg := prometheus.NewRegistry() reg.MustRegister( collectors.NewGoCollector(), @@ -292,23 +372,76 @@ func (srv *Server) start() error { collectors.NewBuildInfoCollector(), ) srv.registerMetrics(reg) + reg.MustRegister(srv.shardScanDuration) + + // Per-scrape gauges for shard stats loop health. GaugeFunc recomputes on + // every Prometheus scrape, so the values stay fresh between scans + // (an expvar.Int gauge would only update once per shard scan, and + // "oldest age" would lag the wall clock between scans). + reg.MustRegister(prometheus.NewGaugeFunc( + prometheus.GaugeOpts{ + Name: "gocached_shard_stats_unscanned", + Help: "number of shards with no cached usage stats; nonzero means the stats loop has not yet covered the full shard space (e.g. just after first deploy of this code)", + }, + func() float64 { + srv.shardStatsMu.Lock() + defer srv.shardStatsMu.Unlock() + return float64(srv.numShards() - len(srv.shardStats)) + }, + )) + reg.MustRegister(prometheus.NewGaugeFunc( + prometheus.GaugeOpts{ + Name: "gocached_shard_stats_oldest_age_seconds", + Help: "age in seconds of the oldest cached shard scan, or 0 if no shard has been scanned yet; alert when this exceeds shardStalenessTarget by some margin to catch the stats loop falling behind", + }, + func() float64 { + _, oldest, _ := srv.shardFreshness(srv.now()) + if oldest.IsZero() { + return 0 + } + return srv.now().Sub(oldest).Seconds() + }, + )) + // Dead-reckoning gauges: PUTs/evictions adjust these between shard + // scans, so they show how much the live state has drifted from the + // persisted aggregate. Signed: a sustained negative delta means evictions + // are running ahead of writes, which is fine; a sustained large positive + // delta hints the stats loop is falling behind write traffic. + reg.MustRegister(prometheus.NewGaugeFunc( + prometheus.GaugeOpts{ + Name: "gocached_pending_blob_count", + Help: "signed pending Action-count delta since the last shard scan; combined with gocached_blob_count this gives the live total", + }, + func() float64 { + c, _ := srv.sumShardDeltas() + return float64(c) + }, + )) + reg.MustRegister(prometheus.NewGaugeFunc( + prometheus.GaugeOpts{ + Name: "gocached_pending_blob_bytes", + Help: "signed pending aggregate-bytes delta since the last shard scan; combined with gocached_blob_bytes this gives the live total used to gate maxSize cleanup", + }, + func() float64 { + _, b := srv.sumShardDeltas() + return float64(b) + }, + )) srv.metricsHandler = promhttp.HandlerFor(reg, promhttp.HandlerOpts{ ErrorLog: log.Default(), }) - srv.logf("gocached: scanning usage & cleaning as needed...") - us, err := srv.usageStats() - if err != nil { - return fmt.Errorf("getting usage stats: %w", err) - } - - srv.logf("gocached: current usage: %v of limit %v", us.All(), bytesFmt(srv.maxSize)) - if res, err := srv.cleanOldObjects(us); err != nil { - return fmt.Errorf("clean old objects: %w", err) - } else if res.Count > 0 { - srv.logf("gocached: cleaned %v", res) + // Seed shardStats + the aggregate from BlobShardStats so usageStats is + // accurate from the first HTTP request after restart, without a full-table + // scan. Stale shards are refreshed by runShardStatsLoop in the + // background; the cleanup loop already runs periodically and will catch + // up any over-limit state once enough shards have been rescanned. + if err := srv.loadShardStats(srv.shutdownCtx); err != nil { + srv.logf("loading persisted shard stats: %v", err) } + srv.logf("gocached: %d/%d shards loaded; usage: %v of limit %v", + len(srv.shardStats), srv.numShards(), srv.lastUsage.Load().All(), bytesFmt(srv.maxSize)) if len(srv.jwtIssuers) > 0 { issuerURLs := make([]string, 0, len(srv.jwtIssuers)) @@ -327,9 +460,12 @@ func (srv *Server) start() error { go srv.runCleanSessionsLoop() } - go srv.runCleanLoop() - go srv.runCheckpointLoop() - go srv.runDBSizeMetricsLoop() + if !srv.disableBackgroundLoops { + go srv.runCleanLoop() + go srv.runCheckpointLoop() + go srv.runDBSizeMetricsLoop() + go srv.runShardStatsLoop() + } return nil } @@ -420,6 +556,18 @@ func WithMaxAge(maxAge time.Duration) ServerOption { } } +// WithShardPrefixLen sets the number of hex characters used as each shard key +// in BlobShardStats. The shard count is 16^n. Valid values are 1..4 inclusive +// (16, 256, 4096, or 65536 shards). Passing 0 means "use the default" (2; +// 256 shards), which keeps individual shard scans fast at hundreds of +// millions of rows. Tune up only if a single shard scan becomes too slow +// under continued growth; any other value causes [NewServer] to fail. +func WithShardPrefixLen(n int) ServerOption { + return func(srv *Server) { + srv.shardPrefixLen = n + } +} + // JWTIssuerConfig configures a single OIDC issuer for JWT-based authentication. type JWTIssuerConfig struct { // Issuer is the OIDC issuer URL. It must be a reachable HTTP(S) server @@ -538,8 +686,53 @@ type Server struct { writeConn *sql.Conn // nil until first used; single connection for writes updateAccessTimeStmt *sql.Stmt // nil until first used; for updating access times on writeConn + // lastUsage holds the most recent aggregate of per-shard stats, recomputed + // after every shard scan. It is the source of truth for /usage, + // BlobCount/BlobBytes gauges, and cleanup decisions; nothing on the hot + // path runs a full-table scan. lastUsage atomic.Pointer[usageStats] + // durs is the set of cohort cutoffs (in standardDurs order, plus maxAge if + // set and not already standard) used by all per-shard scans. Computed once + // in start. Persisted shard stats whose cohort set doesn't match these are + // discarded at load time and re-scanned by the stats loop. + durs []time.Duration + + // shardPrefixLen is the number of hex characters in each shard key. Set + // via [WithShardPrefixLen]; defaults to [defaultShardPrefixLen] when 0, + // and must be in [minShardPrefixLen, maxShardPrefixLen] inclusive once + // [Server.start] has run. The derived shard count is 16^shardPrefixLen, + // exposed via [Server.numShards]. + shardPrefixLen int + + // disableBackgroundLoops, when true, skips the start of every periodic + // goroutine (shard stats, cleanup, checkpoint, DB-size metrics). + // Test-only via withoutBackgroundLoops because the mocked clock collides + // scannedAt across concurrent scans; production never sets this. + disableBackgroundLoops bool + + // shardStatsMu guards shardStats and serializes recomputeAggregateLocked. + shardStatsMu sync.Mutex + shardStats map[shardPrefix]*shardSnapshot + + // shardDeltas is dead-reckoning state: each entry tracks (count, bytes) + // added or removed for that shard since its last persisted scan, so + // cleanupTick can react to write bursts faster than shardStalenessTarget. + // PUTs increment; evictions decrement; scanAndPersistShard subtracts the + // (newStats - oldStats) diff so the delta only carries changes that + // aren't already in the persisted shard. Allocated once in start + // (len == numShards()) so per-shard updates are an addressable struct + // with its own mutex; no allocation per PUT. + shardDeltas []shardDelta + + // shardScanDuration observes the wall time of each per-shard SQL scan. + // Combined with the count of shards, the histogram shows whether any one + // shard scan is pathologically slow (e.g. a hot range with millions of + // entries while others are nearly empty). Registered manually in start + // because it isn't an expvar.Int and so doesn't fit the m-struct + // reflection. + shardScanDuration prometheus.Histogram + // Metrics. Exported fields for reflection, but within a private struct // field to control the gocached Server API surface. m struct { @@ -558,8 +751,9 @@ type Server struct { PutsInline expvar.Int `type:"counter" name:"put_inline" help:"subset of gocached_puts that were stored inline (small objects)"` BlobCount expvar.Int `type:"gauge" name:"blob_count" help:"number of blobs currently stored in the cache"` BlobBytes expvar.Int `type:"gauge" name:"blob_bytes" help:"sum of blob sizes currently stored in the cache"` - EvictedBlobs expvar.Int `type:"counter" name:"evicted_blobs" help:"number of blobs evicted from the cache"` - EvictedBytes expvar.Int `type:"counter" name:"evicted_bytes" help:"number of bytes evicted from the cache"` + EvictedActions expvar.Int `type:"counter" name:"evicted_actions" help:"number of Actions evicted from the cache by the cleanup loop"` + EvictedBlobs expvar.Int `type:"counter" name:"evicted_blobs" help:"number of Blobs evicted from the cache; a Blob is evicted when its last referencing Action is evicted"` + EvictedBytes expvar.Int `type:"counter" name:"evicted_bytes" help:"number of bytes reclaimed by evicting Blobs from the cache"` Sessions expvar.Int `type:"gauge" name:"sessions" help:"number of active authenticated sessions"` Auths expvar.Int `type:"counter" name:"auth_attempts" help:"number of successful token exchanges"` AuthErrs expvar.Int `type:"counter" name:"auth_errs" help:"number of failed token exchanges"` @@ -1153,6 +1347,8 @@ func (s *Server) handlePut(w http.ResponseWriter, r *http.Request, stats *stats) if affected == 0 { stats.PutsDup++ + } else { + s.addBlobDelta(sha256hex, +1, +storedSize) } stats.Puts++ @@ -1367,6 +1563,12 @@ type usageStats struct { // reference a BlobID that doesn't exist in the Blobs table. // This should always be zero in a healthy system. MissingBlobRows int + + // QueryDuration is how long the per-shard SQL query took. It is set by + // [Server.scanShard] for shard snapshots and is left at the zero value + // in the aggregate returned by [Server.usageStats]. Persisted in JSON + // so /usage can surface slow shards across restarts. + QueryDuration time.Duration } func (us *usageStats) All() countAndSize { return us.ActionsLE[math.MaxInt64] } @@ -1384,256 +1586,678 @@ var standardDurs = []time.Duration{ math.MaxInt64, } -func (s *Server) usageStats() (_ *usageStats, err error) { - defer func() { - if err != nil { - s.logf("usageStats error: %v", err) +// computeDurs calculates histogram buckets for usage stats, given the +// configured maxAge. The math.MaxInt64 sentinel is always included as the +// "no upper bound" bucket. +func computeDurs(maxAge time.Duration) []time.Duration { + if maxAge == 0 || slices.Contains(standardDurs, maxAge) { + return slices.Clone(standardDurs) + } + durs := []time.Duration{maxAge, math.MaxInt64} + for _, d := range standardDurs { + if d < maxAge { + durs = append(durs, d) } - }() + } + slices.Sort(durs) + return durs +} - st := &usageStats{ - ActionsLE: make(map[time.Duration]countAndSize), +// shardPrefix is the lowercase hex prefix that identifies one SHA256 shard +// (e.g. "00", "ab"). The named type keeps shard-key strings from being +// passed where any string would do — code that means "shard key" reads as +// shardPrefix; SQL parameters and other free-form strings stay plain string. +type shardPrefix string + +// shardSnapshot is the cached state for a single SHA256-prefix shard. +type shardSnapshot struct { + stats *usageStats + scannedAt time.Time +} + +// shardDelta is the dead-reckoning state for one shard. count and bytes are +// signed: PUTs add positive values, evictions add negative ones. +// scanAndPersistShard subtracts the change in the persisted view of this +// shard (newStats - oldStats) so anything that just landed in the aggregate +// leaves the delta. +// +// The per-shard mu serializes the count+bytes pair so a reader can't catch +// a PUT mid-update with count incremented but bytes not (or vice versa). +// Storing the values inline avoids allocating a new struct on every PUT. +type shardDelta struct { + mu sync.Mutex + count int64 + bytes int64 +} + +// addBlobDelta adjusts the shardDelta for the shard owning sha256hex by the +// given signed (count, bytes) values. Callers pass +1, +storedSize on a +// successful new Action insert, and -1, -storedSize when evicting an Action. +// Bogus sha256hex (too short or non-hex prefix) is silently ignored. +func (srv *Server) addBlobDelta(sha256hex string, count, bytes int64) { + if srv.shardPrefixLen <= 0 || len(sha256hex) < srv.shardPrefixLen { + return + } + n, err := strconv.ParseUint(sha256hex[:srv.shardPrefixLen], 16, 64) + if err != nil || int(n) >= len(srv.shardDeltas) { + return } + sd := &srv.shardDeltas[n] + sd.mu.Lock() + sd.count += count + sd.bytes += bytes + sd.mu.Unlock() +} - // Build the durations to use for the histogram. - // The math.MaxInt64 value is always included. - // If s.maxAge is set, we ignore sizes above that, except - // for the math.MaxInt64 value. - var durs []time.Duration - if s.maxAge == 0 { - durs = standardDurs +// sumShardDeltas returns the sum of every shard's pending delta. The caller +// adds this to the persisted aggregate to get a live estimate of total cache +// occupancy between scans. +func (srv *Server) sumShardDeltas() (count, bytes int64) { + for i := range srv.shardDeltas { + sd := &srv.shardDeltas[i] + sd.mu.Lock() + count += sd.count + bytes += sd.bytes + sd.mu.Unlock() + } + return +} + +// numShards returns 16^shardPrefixLen, the total number of SHA256-prefix +// shards used to partition usage statistics. +func (srv *Server) numShards() int { + return 1 << (4 * srv.shardPrefixLen) +} + +// shardPrefix returns the hex prefix for shard index i, zero-padded to +// shardPrefixLen characters (e.g. 0 -> "00", 255 -> "ff"). +func (srv *Server) shardPrefix(i int) shardPrefix { + return shardPrefix(fmt.Sprintf("%0*x", srv.shardPrefixLen, i)) +} + +// shardRange returns the [lo, hi) SHA256 hex range covered by shard index i. +// For the final shard, hi is a sentinel string of 'g' characters: 'g' sorts +// after every valid hex digit, so "SHA256 < hi" excludes nothing in that +// range. Using a sentinel keeps the SQL uniform (always two parameters). +// Returns plain strings since both ends bind directly into SQL. +func (srv *Server) shardRange(i int) (lo, hi string) { + lo = string(srv.shardPrefix(i)) + if i+1 == srv.numShards() { + hi = strings.Repeat("g", srv.shardPrefixLen) } else { - durs = make([]time.Duration, 0, len(standardDurs)+1) - durs = append(durs, s.maxAge) - for _, d := range standardDurs { - if d < s.maxAge || d == math.MaxInt64 { - durs = append(durs, d) - } + hi = string(srv.shardPrefix(i + 1)) + } + return +} + +// usageStats returns the most recent aggregate of per-shard stats. It does no +// SQL work: the aggregate is maintained by the shard stats loop and seeded +// from the BlobShardStats table at startup, so callers see usage immediately +// after process restart rather than after a multi-minute full-table scan. The +// returned value is never nil; ActionsLE may be empty if no shard has been +// scanned and no persisted stats exist. +func (s *Server) usageStats() (*usageStats, error) { + if us := s.lastUsage.Load(); us != nil { + return us, nil + } + return &usageStats{ActionsLE: make(map[time.Duration]countAndSize)}, nil +} + +// scanShard runs a single SQL-aggregated query that computes the cohort +// histogram for Actions whose Blob's SHA256 falls in shard idx's range. +// The aggregation happens entirely in SQL via SUM(CASE WHEN ...) clauses, so +// the driver returns one row regardless of how many Actions fall in the shard. +// This is what makes a scan cheap enough to pin a reader snapshot for only +// milliseconds per shard rather than minutes for the whole table. +func (srv *Server) scanShard(ctx context.Context, idx int) (*usageStats, error) { + durs := srv.durs + now := srv.now().Unix() + + var sb strings.Builder + args := make([]any, 0, 2*len(durs)+2) + sb.WriteString("SELECT ") + finiteCount := 0 + for _, d := range durs { + if d == math.MaxInt64 { + continue } - slices.Sort(durs) + cutoff := now - int64(d/time.Second) + if finiteCount > 0 { + sb.WriteByte(',') + } + sb.WriteString("COALESCE(SUM(CASE WHEN a.AccessTime > ? THEN 1 ELSE 0 END),0),") + sb.WriteString("COALESCE(SUM(CASE WHEN a.AccessTime > ? THEN b.StoredSize ELSE 0 END),0)") + args = append(args, cutoff, cutoff) + finiteCount++ + } + sb.WriteString(",COUNT(*),COALESCE(SUM(b.StoredSize),0)") + sb.WriteString(" FROM Blobs b JOIN Actions a ON a.BlobID = b.BlobID") + sb.WriteString(" WHERE b.SHA256 >= ? AND b.SHA256 < ?") + lo, hi := srv.shardRange(idx) + args = append(args, lo, hi) + + vals := make([]int64, finiteCount*2+2) + dest := make([]any, len(vals)) + for i := range vals { + dest[i] = &vals[i] + } + start := time.Now() + if err := srv.db.QueryRowContext(ctx, sb.String(), args...).Scan(dest...); err != nil { + return nil, err + } + queryDuration := time.Since(start) + if srv.shardScanDuration != nil { + // Nil-check: scanShard is reachable from tests that build a bare + // Server without calling start (e.g. TestShardPrefix). + srv.shardScanDuration.Observe(queryDuration.Seconds()) } - // Flush any pending access time bumps before computing usage stats. - s.mu.RLock() - shouldFlush := len(s.accessDirty) > 0 - s.mu.RUnlock() - if shouldFlush { - // This acquires sqliteWriteMu, so we avoid that lock if there's nothing - // to do. - s.flushAccessTimeBumps() + st := &usageStats{ + ActionsLE: make(map[time.Duration]countAndSize, len(durs)), + QueryDuration: queryDuration, + } + j := 0 + for _, d := range durs { + if d == math.MaxInt64 { + continue + } + st.ActionsLE[d] = countAndSize{Count: vals[j*2], Size: vals[j*2+1]} + j++ } + st.ActionsLE[math.MaxInt64] = countAndSize{Count: vals[finiteCount*2], Size: vals[finiteCount*2+1]} + return st, nil +} + +// shardCohortsMatch reports whether the persisted ActionsLE map keys exactly +// match the current srv.durs set. A mismatch means the persisted shard was +// written with a different maxAge (or a different code version), so we +// discard it and force the stats loop to re-scan that shard. +func shardCohortsMatch(le map[time.Duration]countAndSize, durs []time.Duration) bool { + if len(le) != len(durs) { + return false + } + for _, d := range durs { + if _, ok := le[d]; !ok { + return false + } + } + return true +} - now := s.now().Unix() - rows, err := s.db.Query( - "SELECT a.BlobID, a.AccessTime, b.StoredSize FROM Actions a LEFT JOIN Blobs b ON a.BlobID = b.BlobID") +// loadShardStats reads every BlobShardStats row at startup, validates that +// the persisted cohort set matches srv.durs, and seeds shardStats + the +// aggregate. Rows with mismatched cohorts are dropped so the stats loop +// re-scans them. The aggregate becomes the source of truth for /usage and the +// cleanup loop until the first new shard scan completes. +func (srv *Server) loadShardStats(ctx context.Context) error { + rows, err := srv.db.QueryContext(ctx, "SELECT Prefix, ScannedAt, StatsJSON FROM BlobShardStats") if err != nil { - return nil, fmt.Errorf("query Actions: %w", err) + return err } - var blobID int64 - var accessTime int64 - var storedSize sql.NullInt64 + defer rows.Close() + srv.shardStatsMu.Lock() + defer srv.shardStatsMu.Unlock() + srv.shardStats = make(map[shardPrefix]*shardSnapshot, srv.numShards()) for rows.Next() { - if err := rows.Scan(&blobID, &accessTime, &storedSize); err != nil { - return nil, fmt.Errorf("rows.Scan: %w", err) + var prefix, statsJSON string + var scannedAt int64 + if err := rows.Scan(&prefix, &scannedAt, &statsJSON); err != nil { + return err } - if !storedSize.Valid { - st.MissingBlobRows++ + if len(prefix) != srv.shardPrefixLen { + // A previous run used a different shardPrefixLen. The stats loop + // won't write to these keys, so they'd otherwise sit in the + // table forever contributing stale data to the aggregate. + srv.logf("BlobShardStats[%q] prefix length %d != current %d; will rescan", + prefix, len(prefix), srv.shardPrefixLen) continue } - - dur := time.Duration(now-accessTime) * time.Second - if dur < 0 { - dur = 0 + var st usageStats + if err := json.Unmarshal([]byte(statsJSON), &st); err != nil { + srv.logf("BlobShardStats[%q] JSON decode: %v; will rescan", prefix, err) + continue } - for _, d := range durs { - if dur < d { - was := st.ActionsLE[d] - was.Count++ - was.Size += storedSize.Int64 - st.ActionsLE[d] = was - } + if !shardCohortsMatch(st.ActionsLE, srv.durs) { + srv.logf("BlobShardStats[%q] cohort mismatch; will rescan", prefix) + continue + } + srv.shardStats[shardPrefix(prefix)] = &shardSnapshot{ + stats: &st, + scannedAt: time.Unix(scannedAt, 0), + } + // Replay the persisted scan time into the histogram so the + // /metrics endpoint shows reasonable percentiles immediately after + // restart instead of waiting a full pass over every shard. The observation is + // older than "now" but it's the most accurate signal we have for + // this shard until the stats loop refreshes it. + if srv.shardScanDuration != nil && st.QueryDuration > 0 { + srv.shardScanDuration.Observe(st.QueryDuration.Seconds()) } } if err := rows.Err(); err != nil { - return nil, fmt.Errorf("rows.Next: %w", err) + return err } + srv.recomputeAggregateLocked() + return nil +} - s.lastUsage.Store(st) - all := st.All() - s.m.BlobCount.Set(all.Count) - s.m.BlobBytes.Set(all.Size) - return st, nil +// upsertShardStats persists one shard's scan result. The row is keyed on +// Prefix so subsequent scans of the same shard overwrite in place. +func (srv *Server) upsertShardStats(ctx context.Context, prefix shardPrefix, scannedAt time.Time, st *usageStats) error { + statsJSON, err := json.Marshal(st) + if err != nil { + return err + } + srv.sqliteWriteMu.Lock() + defer srv.sqliteWriteMu.Unlock() + _, err = srv.db.ExecContext(ctx, + `INSERT INTO BlobShardStats (Prefix, ScannedAt, StatsJSON) VALUES (?, ?, ?) + ON CONFLICT(Prefix) DO UPDATE SET ScannedAt=excluded.ScannedAt, StatsJSON=excluded.StatsJSON`, + prefix, scannedAt.Unix(), string(statsJSON)) + return err +} + +// recomputeAggregateLocked sums the cohort histograms across every cached +// shard, stores the result in lastUsage, and updates the BlobCount/BlobBytes +// gauges. The caller must hold shardStatsMu. +func (srv *Server) recomputeAggregateLocked() { + agg := &usageStats{ActionsLE: make(map[time.Duration]countAndSize, len(srv.durs))} + for _, d := range srv.durs { + agg.ActionsLE[d] = countAndSize{} + } + for _, sh := range srv.shardStats { + for d, cs := range sh.stats.ActionsLE { + was := agg.ActionsLE[d] + was.Count += cs.Count + was.Size += cs.Size + agg.ActionsLE[d] = was + } + } + srv.lastUsage.Store(agg) + total := agg.All() + srv.m.BlobCount.Set(total.Count) + srv.m.BlobBytes.Set(total.Size) +} + +// pickOldestShard returns the shard index whose ScannedAt is oldest, along +// with that ScannedAt. A shard that has never been scanned (no entry in +// shardStats) wins immediately, and is returned with a zero time so the +// caller treats it as infinitely stale (scan it now, no sleep). +func (srv *Server) pickOldestShard() (idx int, scannedAt time.Time) { + srv.shardStatsMu.Lock() + defer srv.shardStatsMu.Unlock() + oldestIdx := -1 + var oldestTime time.Time + for i := range srv.numShards() { + sh, ok := srv.shardStats[srv.shardPrefix(i)] + if !ok { + return i, time.Time{} + } + if oldestIdx == -1 || sh.scannedAt.Before(oldestTime) { + oldestIdx = i + oldestTime = sh.scannedAt + } + } + return oldestIdx, oldestTime +} + +// scanAndPersistShard scans shard idx, writes the result back to +// BlobShardStats, and updates the in-memory cache + aggregate. +func (srv *Server) scanAndPersistShard(ctx context.Context, idx int) error { + prefix := srv.shardPrefix(idx) + scannedAt := srv.now() + stats, err := srv.scanShard(ctx, idx) + if err != nil { + return fmt.Errorf("scan shard %q: %w", prefix, err) + } + if err := srv.upsertShardStats(ctx, prefix, scannedAt, stats); err != nil { + return fmt.Errorf("persist shard %q: %w", prefix, err) + } + + srv.shardStatsMu.Lock() + var oldAll countAndSize + if old, ok := srv.shardStats[prefix]; ok { + oldAll = old.stats.All() + } + srv.shardStats[prefix] = &shardSnapshot{stats: stats, scannedAt: scannedAt} + srv.recomputeAggregateLocked() + srv.shardStatsMu.Unlock() + + // Subtract from the delta the exact change that just landed in the + // persisted view: (newStats - oldStats). Anything in delta whose Action + // commit happened before scanShard's snapshot is now in newStats, so we + // subtract it out; anything that committed after the snapshot isn't in + // newStats, so the diff doesn't touch its delta contribution. PUTs whose + // addBlobDelta hasn't been called yet temporarily push delta negative; + // when the addBlobDelta runs it cancels out. + // + // First-scan caveat: when oldStats was missing (no prior scan for this + // shard AND no persisted row from a previous process run), oldAll is + // zero. If the DB had pre-existing rows that delta never saw (e.g. first + // deploy of this code on a populated DB whose BlobShardStats table was + // empty), subtracting newStats from delta will leave delta with a + // negative offset equal to the pre-existing row count. The aggregate + // becomes correct from then on; the negative offset stays in the delta + // until process restart (which reloads aggregate from the now-populated + // BlobShardStats and starts delta at zero). Operators deploying onto a + // populated DB can avoid this by nuking the DB first. + newAll := stats.All() + diffCount := newAll.Count - oldAll.Count + diffBytes := newAll.Size - oldAll.Size + sd := &srv.shardDeltas[idx] + sd.mu.Lock() + sd.count -= diffCount + sd.bytes -= diffBytes + sd.mu.Unlock() + return nil } -type cleanCandidate struct { - BlobID int64 - Age time.Duration - StoredSize int64 // size of the blob as stored, in bytes (compressed if lz4) +// runShardStatsLoop keeps usage stats fresh by rescanning the +// oldest-scanned shard whenever it crosses [shardStalenessTarget]. The +// per-scan cadence in steady state works out to about +// shardStalenessTarget/numShards (~2.3s at defaults) +func (srv *Server) runShardStatsLoop() { + for { + idx, scannedAt := srv.pickOldestShard() + var sleep time.Duration + if !scannedAt.IsZero() { + sleep = max(shardStalenessTarget-srv.now().Sub(scannedAt), shardStatsMinInterval) + } + if sleep > 0 { + select { + case <-srv.shutdownCtx.Done(): + return + case <-time.After(sleep): + } + } + if err := srv.scanAndPersistShard(srv.shutdownCtx, idx); err != nil { + if errors.Is(err, context.Canceled) { + return + } + srv.logf("shard stats loop: %v", err) + } + } } -func (s *Server) cleanCandidates(olderThan time.Duration, limit int64) ([]cleanCandidate, error) { - now := s.now() - nowUnix := now.Unix() - cutoff := now.Add(-olderThan).Unix() +// scanAllShards synchronously rescans every shard. Used by tests to get a +// deterministic aggregate, and by POST /usage as a manual "refresh now" +// trigger when an operator doesn't want to wait for the stats loop. +func (srv *Server) scanAllShards(ctx context.Context) error { + for i := range srv.numShards() { + if err := srv.scanAndPersistShard(ctx, i); err != nil { + return err + } + } + return nil +} - rows, err := s.db.Query(` - SELECT b.BlobID, MAX(a.AccessTime), b.StoredSize - FROM Blobs b LEFT JOIN Actions a ON b.BlobID = a.BlobID - GROUP BY b.BlobID - HAVING MAX(a.AccessTime) <= ? - ORDER BY MAX(a.AccessTime) - LIMIT ?`, cutoff, limit) - if err != nil { - return nil, fmt.Errorf("query clean candidates: %w", err) +// shardScanDurationPercentiles returns p25/p50/p90 of the QueryDuration +// across every cached shard, and the count of shards that contributed (i.e. +// have a non-zero QueryDuration). Returns zeros if no shard has been scanned +// yet. Linear interpolation is overkill given numShards is 16/256/etc.; a +// rank-index pick is good enough for a /usage page. +func (srv *Server) shardScanDurationPercentiles() (p25, p50, p90 time.Duration, n int) { + srv.shardStatsMu.Lock() + durs := make([]time.Duration, 0, len(srv.shardStats)) + for _, sh := range srv.shardStats { + if sh.stats.QueryDuration > 0 { + durs = append(durs, sh.stats.QueryDuration) + } } - defer rows.Close() + srv.shardStatsMu.Unlock() + n = len(durs) + if n == 0 { + return + } + slices.Sort(durs) + pick := func(p int) time.Duration { + i := (n * p) / 100 + if i >= n { + i = n - 1 + } + return durs[i] + } + return pick(25), pick(50), pick(90), n +} + +// shardFreshnessBuckets are the "scanned within last X" age cohorts shown on +// the /usage page. The largest cohort is intentionally larger than +// [shardStalenessTarget] so the page makes it obvious when the stats loop +// has fallen behind (the smaller cohorts drop to zero before the largest does). +var shardFreshnessBuckets = []time.Duration{ + 1 * time.Minute, + 5 * time.Minute, + 15 * time.Minute, + 1 * time.Hour, + 6 * time.Hour, + 24 * time.Hour, +} + +// shardFreshness counts how many cached shards were last scanned within each +// of [shardFreshnessBuckets], and returns the absolute time of the oldest +// scan. The buckets are cumulative ("within last X"), matching the cohort +// histogram style elsewhere in this package. +func (srv *Server) shardFreshness(now time.Time) (counts map[time.Duration]int, oldest time.Time, total int) { + counts = make(map[time.Duration]int, len(shardFreshnessBuckets)) + srv.shardStatsMu.Lock() + defer srv.shardStatsMu.Unlock() + total = len(srv.shardStats) + for _, sh := range srv.shardStats { + age := now.Sub(sh.scannedAt) + for _, b := range shardFreshnessBuckets { + if age <= b { + counts[b]++ + } + } + if oldest.IsZero() || sh.scannedAt.Before(oldest) { + oldest = sh.scannedAt + } + } + return counts, oldest, total +} + +// evictionCandidateQuery is the SQL used by [Server.evictOldestActions] to +// pick which Actions to delete next. INDEXED BY is the documented +// SQLite mechanism (https://sqlite.org/lang_indexedby.html) for locking +// down a plan so a future schema change can't silently regress this query +// from an O(log N) index range scan into an O(N) table scan. The plan is +// verified in TestEvictionQueryPlan. +const evictionCandidateQuery = ` +SELECT NamespaceID, ActionID, BlobID +FROM Actions INDEXED BY idx_actions_access +WHERE AccessTime <= ? +ORDER BY AccessTime ASC +LIMIT ?` + +// evictOldestActions deletes up to maxCount Actions whose AccessTime is at +// or before cutoff, oldest first, in a single short write transaction. For +// each deleted Action, if its Blob has no other referencing Actions +// remaining, the Blob row and its on-disk file are also removed. The batch +// stops early once at least maxBytes have been reclaimed (use math.MaxInt64 +// for no byte budget — e.g., maxAge cleanup, where we want every stale +// Action gone regardless of size). +// +// This is Action-LRU: a stale Action is evicted even if its Blob is shared +// with newer Actions (the Blob then stays alive via those newer refs). That +// differs from the prior Blob-LRU policy, where a Blob was only evicted when +// MAX(AccessTime) of all its Actions was past the cutoff. The cache is +// mostly 1:1 Action-to-Blob so the two policies are equivalent for the +// common case; on shared Blobs, Action-LRU is the stricter LRU semantics. +func (srv *Server) evictOldestActions(ctx context.Context, cutoff int64, maxCount int, maxBytes int64) (countAndSize, error) { + var ret countAndSize - var candidates []cleanCandidate - var accessTime int64 + rows, err := srv.db.QueryContext(ctx, evictionCandidateQuery, cutoff, maxCount) + if err != nil { + return ret, fmt.Errorf("query eviction candidates: %w", err) + } + type candidate struct { + NamespaceID int64 + ActionID string + BlobID int64 + } + var cands []candidate for rows.Next() { - var c cleanCandidate - if err := rows.Scan(&c.BlobID, &accessTime, &c.StoredSize); err != nil { - return nil, fmt.Errorf("rows.Scan: %w", err) + var c candidate + if err := rows.Scan(&c.NamespaceID, &c.ActionID, &c.BlobID); err != nil { + rows.Close() + return ret, fmt.Errorf("scan candidate: %w", err) } - c.Age = time.Duration(nowUnix-accessTime) * time.Second - candidates = append(candidates, c) + cands = append(cands, c) } - if err := rows.Err(); err != nil { - return nil, fmt.Errorf("rows.Next: %w", err) + if err := rows.Close(); err != nil { + return ret, fmt.Errorf("close candidate rows: %w", err) + } + if len(cands) == 0 { + return ret, nil } - return candidates, nil -} - -func (srv *Server) deleteBlobs(blobIDs ...int64) error { srv.sqliteWriteMu.Lock() defer srv.sqliteWriteMu.Unlock() - tx, err := srv.db.Begin() + tx, err := srv.db.BeginTx(ctx, nil) if err != nil { - return fmt.Errorf("delete blob Begin: %w", err) + return ret, fmt.Errorf("eviction Begin: %w", err) } defer tx.Rollback() - var sumBytes int64 - for _, blobID := range blobIDs { + // pendingDelta queues addBlobDelta updates for after Commit. Holding + // them until commit means a rolled-back tx (rare) doesn't poison the + // dead-reckoning state. + type pendingDelta struct { + sha256Hex string + storedSize int64 + } + var pending []pendingDelta + + var evictedBlobs int64 + var evictedBlobBytes int64 // bytes reclaimed (disk space + EvictedBytes metric) + var evictedActionBytes int64 // aggregate-side bytes reduction (drives maxBytes check) + var evictedActions int64 + for _, c := range cands { + // Fetch SHA256 + StoredSize up front so we can both (a) decrement + // the dead-reckoning delta for this Action and (b) remove the disk + // file if the Blob ends up orphaned. A LEFT-JOIN here is just an + // extra query; the row is keyed on the primary key so it's cheap. var sha256Hex string var storedSize int64 - if err := tx.QueryRow("SELECT SHA256, StoredSize FROM Blobs WHERE BlobID = ?", blobID).Scan(&sha256Hex, &storedSize); err != nil && !errors.Is(err, sql.ErrNoRows) { - return fmt.Errorf("querying blob SHA256: %w", err) + blobExists := true + switch err := tx.QueryRowContext(ctx, + "SELECT SHA256, StoredSize FROM Blobs WHERE BlobID = ?", + c.BlobID).Scan(&sha256Hex, &storedSize); { + case err == nil: + case errors.Is(err, sql.ErrNoRows): + blobExists = false + default: + return ret, fmt.Errorf("fetch Blob row: %w", err) } - sumBytes += storedSize - if _, err := tx.Exec("DELETE FROM Blobs WHERE BlobID = ?", blobID); err != nil { - return fmt.Errorf("deleting blob: %w", err) + + if _, err := tx.ExecContext(ctx, + "DELETE FROM Actions WHERE NamespaceID = ? AND ActionID = ?", + c.NamespaceID, c.ActionID); err != nil { + return ret, fmt.Errorf("delete action: %w", err) + } + evictedActions++ + if !blobExists { + // Orphan Action pointing at a missing Blob; delete the row and + // move on. No delta, no Blob row, no file. + continue + } + // Each Action contributes 1 count + storedSize bytes to the aggregate + // (see scanShard); removing it reduces both by exactly that. Decrement + // deferred until after Commit so a rollback doesn't leave the delta + // off. + pending = append(pending, pendingDelta{sha256Hex, storedSize}) + evictedActionBytes += storedSize + + // Is this Blob now orphaned? idx_actions_blobid makes this O(log N). + var stillReferenced int + if err := tx.QueryRowContext(ctx, + "SELECT EXISTS(SELECT 1 FROM Actions WHERE BlobID = ?)", + c.BlobID).Scan(&stillReferenced); err != nil { + return ret, fmt.Errorf("check orphan: %w", err) + } + if stillReferenced == 1 { + continue } - if _, err := tx.Exec("DELETE FROM Actions WHERE BlobID = ?", blobID); err != nil { - return fmt.Errorf("deleting actions: %w", err) + if _, err := tx.ExecContext(ctx, "DELETE FROM Blobs WHERE BlobID = ?", c.BlobID); err != nil { + return ret, fmt.Errorf("delete Blob: %w", err) } + evictedBlobs++ + evictedBlobBytes += storedSize var hash [sha256.Size]byte if _, err := hex.Decode(hash[:], []byte(sha256Hex)); err == nil { base := srv.sha256Filepath(hash) // Try removing both lz4 and plain paths; one or neither may exist. if err := os.Remove(base + ".lz4"); err != nil && !os.IsNotExist(err) { - return fmt.Errorf("removing disk file: %w", err) + return ret, fmt.Errorf("removing disk file: %w", err) } if err := os.Remove(base); err != nil && !os.IsNotExist(err) { - return fmt.Errorf("removing disk file: %w", err) + return ret, fmt.Errorf("removing disk file: %w", err) } } + // Stop once we've freed enough aggregate bytes. For maxAge cleanup + // callers pass math.MaxInt64 so this check never triggers. + if evictedActionBytes >= maxBytes { + break + } } if err := tx.Commit(); err != nil { - return err + return ret, fmt.Errorf("eviction Commit: %w", err) } - srv.m.EvictedBlobs.Add(int64(len(blobIDs))) - srv.m.EvictedBytes.Add(sumBytes) + // Apply the queued delta decrements now that the tx is durably committed. + for _, p := range pending { + srv.addBlobDelta(p.sha256Hex, -1, -p.storedSize) + } - return nil + srv.m.EvictedActions.Add(evictedActions) + srv.m.EvictedBlobs.Add(evictedBlobs) + srv.m.EvictedBytes.Add(evictedBlobBytes) + ret.Count = evictedBlobs + ret.Size = evictedBlobBytes + return ret, nil } -func (srv *Server) cleanOldObjects(us *usageStats) (countAndSize, error) { - var zero countAndSize +// cleanupTick decides whether the cache is over its configured maxAge or +// maxSize and, if so, runs one [Server.evictOldestActions] batch. The +// aggregate from the shard stats loop drives the decision; the eviction itself +// goes straight at idx_actions_access without consulting the shards. +func (srv *Server) cleanupTick(ctx context.Context) (countAndSize, error) { var ret countAndSize - - all := us.ActionsLE[math.MaxInt64] - if srv.verbose { - srv.logf("current usage stats: %v", all) - last := all - for _, d := range slices.Sorted(maps.Keys(us.ActionsLE)) { - if d == math.MaxInt64 { - continue // skip infinity - } - c := us.ActionsLE[d] - srv.logf(" <=%v: %v", durFmt(d), c) - if last == c { - break - } - last = c - } - } - - // First clean things that are just too old. + us := srv.lastUsage.Load() + if us == nil { + return ret, nil + } + all := us.All() + // Include the dead-reckoning delta in the size pressure check so a + // burst of PUTs between scans actually triggers cleanup. The maxAge + // check intentionally ignores the delta: fresh PUTs are by definition + // within maxAge, so they neither add to nor subtract from the count of + // over-age Actions. + dCount, dBytes := srv.sumShardDeltas() + all.Count += dCount + all.Size += dBytes + + // Default cutoff to "any age" for size-only cleanup. evictOldestActions + // still uses idx_actions_access (the WHERE/ORDER BY drive the plan), so + // the very-large cutoff just means "no upper bound". + cutoff := int64(math.MaxInt64) + overAge := false if srv.maxAge > 0 { - if toDelete := all.Count - us.ActionsLE[srv.maxAge].Count; toDelete > 0 { - srv.logf("Cleaning %d objects older than %v ...", toDelete, durFmt(srv.maxAge)) - candidates, err := srv.cleanCandidates(srv.maxAge, toDelete+1) - if err != nil { - return zero, fmt.Errorf("getting clean candidates: %v", err) - } - blobIDs := make([]int64, 0, len(candidates)) - var sumSize int64 - for _, c := range candidates { - blobIDs = append(blobIDs, c.BlobID) - sumSize += c.StoredSize - } - if err := srv.deleteBlobs(blobIDs...); err != nil { - return zero, fmt.Errorf("deleting old blobs: %v", err) - } - all.Count -= int64(len(candidates)) - all.Size -= sumSize - ret.Count += int64(len(candidates)) - ret.Size += sumSize - } + cutoff = srv.now().Add(-srv.maxAge).Unix() + overAge = us.All().Count-us.ActionsLE[srv.maxAge].Count > 0 } - - for srv.maxSize > 0 && all.Size > srv.maxSize { - toClean := all.Size - srv.maxSize - if srv.verbose { - srv.logf("need to clean %v to get under max size of %v ...", - bytesFmt(toClean), bytesFmt(srv.maxSize)) - } - - var batchBytes int64 - var blobIDs []int64 - candidates, err := srv.cleanCandidates(0, 10000) - if err != nil { - return zero, fmt.Errorf("getting clean candidates: %v", err) - } - for _, c := range candidates { - blobIDs = append(blobIDs, c.BlobID) - batchBytes += c.StoredSize - if batchBytes >= toClean { - break - } - } - if err := srv.deleteBlobs(blobIDs...); err != nil { - return zero, fmt.Errorf("deleting old blobs: %v", err) - } - - ret.Count += int64(len(blobIDs)) - ret.Size += batchBytes - all.Count -= int64(len(blobIDs)) - all.Size -= batchBytes - - if len(blobIDs) == len(candidates) { - // We didn't find enough candidates to delete. - // Just stop here. - srv.logf("[unexpected] didn't find enough candidates to delete") - break - } + overSize := srv.maxSize > 0 && all.Size > srv.maxSize + if !overAge && !overSize { + return ret, nil } - - return ret, nil + // maxAge cleanup runs to count limit (every stale Action must go); + // size-only cleanup runs until we've reclaimed enough bytes. + maxBytes := int64(math.MaxInt64) + if !overAge && overSize { + maxBytes = all.Size - srv.maxSize + } + return srv.evictOldestActions(ctx, cutoff, cleanupBatchSize, maxBytes) } // checkpointTruncate runs PRAGMA wal_checkpoint(TRUNCATE) and returns SQLite's @@ -1733,23 +2357,13 @@ func (srv *Server) runCleanLoop() { select { case <-srv.shutdownCtx.Done(): return - case <-time.After(5 * time.Minute): + case <-time.After(cleanupTickInterval): } - - us, err := srv.usageStats() - if err != nil { - srv.logf("error getting usage stats: %v", err) - continue - } - - res, err := srv.cleanOldObjects(us) - if err != nil { - srv.logf("error cleaning old objects: %v", err) - continue - } - if res.Count > 0 { - srv.logf("cleaned %v", res) - srv.usageStats() // for side effect of updating lastUsage + if _, err := srv.cleanupTick(srv.shutdownCtx); err != nil { + if errors.Is(err, context.Canceled) { + return + } + srv.logf("cleanup: %v", err) } } } @@ -1815,10 +2429,11 @@ func bytesFmt(n int64) string { func (srv *Server) serveUsage(w http.ResponseWriter, r *http.Request) { if r.Method == "POST" { - // For side effect of updating lastUsage. - _, err := srv.usageStats() - if err != nil { - http.Error(w, "error getting usage stats: "+err.Error(), http.StatusInternalServerError) + // Manual "refresh now" trigger: a synchronous scan of every shard. + // Slow on a large DB (numShards * per-shard query time), but + // operators may want an immediate aggregate after a config change. + if err := srv.scanAllShards(r.Context()); err != nil { + http.Error(w, "scanAllShards: "+err.Error(), http.StatusInternalServerError) return } } @@ -1829,12 +2444,23 @@ func (srv *Server) serveUsage(w http.ResponseWriter, r *http.Request) { return } - // Print out an HTML table of the usage stats, sorted by age. + // Dead-reckon the totals so PUT/eviction bursts since the last shard + // scan show up on this page immediately. The persisted vs pending split + // is implementation noise for /usage readers; if anyone needs the + // breakdown they can read gocached_{blob,pending_blob}_{count,bytes} + // from /metrics. + dCount, dBytes := srv.sumShardDeltas() + live := us.All() + live.Count += dCount + live.Size += dBytes + w.Header().Set("Content-Type", "text/html; charset=utf-8") fmt.Fprintf(w, "

gocached usage stats

\n") fmt.Fprintf(w, "

Current usage: %v of limit %v

\n", - us.All(), bytesFmt(srv.maxSize)) + live, bytesFmt(srv.maxSize)) + // Cohort histogram of stored Actions, by access-time age. + fmt.Fprintf(w, "

Actions by access-time age

\n") fmt.Fprintf(w, "
\n") fmt.Fprintf(w, "\n") for _, d := range slices.Sorted(maps.Keys(us.ActionsLE)) { @@ -1849,6 +2475,41 @@ func (srv *Server) serveUsage(w http.ResponseWriter, r *http.Request) { title, c.Count, bytesFmt(c.Size)) } fmt.Fprintf(w, "
AgeCountSize
\n") + + // Shard scan freshness — how many shards were scanned within each + // cohort. Watching the smaller-bucket counts drop is the easiest way to + // spot the stats loop falling behind. + now := srv.now() + freshness, oldest, totalShards := srv.shardFreshness(now) + fmt.Fprintf(w, "

Shard scan freshness

\n") + fmt.Fprintf(w, "

%d of %d shards have a cached scan; staleness target %v.

\n", + totalShards, srv.numShards(), shardStalenessTarget) + if !oldest.IsZero() { + fmt.Fprintf(w, "

Oldest scan: %v ago (%v).

\n", + durFmt(now.Sub(oldest).Round(time.Second)), oldest.UTC().Format(time.RFC3339)) + } + fmt.Fprintf(w, "\n") + fmt.Fprintf(w, "\n") + for _, b := range shardFreshnessBuckets { + fmt.Fprintf(w, "\n", durFmt(b), freshness[b]) + } + fmt.Fprintf(w, "
Scanned withinShards
<= %v%d
\n") + + // Per-shard query duration percentiles, snapshotted from the cached + // QueryDuration of every persisted shard. The Prometheus histogram + // gocached_shard_scan_duration_seconds carries the same data over time. + p25, p50, p90, nDur := srv.shardScanDurationPercentiles() + fmt.Fprintf(w, "

Per-shard scan duration

\n") + if nDur == 0 { + fmt.Fprintf(w, "

No shard has been scanned yet.

\n") + } else { + fmt.Fprintf(w, "\n") + fmt.Fprintf(w, "\n") + fmt.Fprintf(w, "\n", nDur, p25.Round(time.Millisecond)) + fmt.Fprintf(w, "\n", p50.Round(time.Millisecond)) + fmt.Fprintf(w, "\n", p90.Round(time.Millisecond)) + fmt.Fprintf(w, "
PercentileDuration
p25 (n=%d)%v
p50%v
p90%v
\n") + } } func (srv *Server) serveSessions(w http.ResponseWriter, r *http.Request) { diff --git a/gocache/gocached/gocached_perf_test.go b/gocache/gocached/gocached_perf_test.go new file mode 100644 index 0000000..bbe2911 --- /dev/null +++ b/gocache/gocached/gocached_perf_test.go @@ -0,0 +1,275 @@ +// Copyright (c) Tailscale Inc & AUTHORS +// SPDX-License-Identifier: BSD-3-Clause + +package gocached + +// This file holds the perf regression test. It always runs in CI (so +// refactors that change query shape are caught by the assertions immediately), +// but defaults to a tiny row count so normal CI takes a fraction of a +// second. Operators investigating real-world scale set +// -performance-test-rows to e.g. 250_000_000 and -performance-test-dir to a +// path on a fast disk; the seed is reused across runs. + +import ( + "context" + "crypto/sha256" + "database/sql" + "encoding/hex" + "errors" + "flag" + "fmt" + "math" + "os" + "strconv" + "strings" + "testing" + "time" +) + +var ( + perfTestRows = flag.Int("performance-test-rows", 100, + "target number of (Blob, Action) pairs to populate before measuring; default 100 so plain `go test` finishes in a fraction of a second. Bump to e.g. 250_000_000 for real-scale measurements; only the deficit is inserted on each run, so re-runs with the same value are instant.") + perfTestDir = flag.String("performance-test-dir", "", + "directory holding the pre-seeded perf test DB. If empty (the default), a fresh per-run t.TempDir() is used. Set to a persistent path (with $GOCACHED_PERFTEST_DIR as an alternative) when -performance-test-rows is large enough that you don't want to re-seed every run.") +) + +// perfSeedBatchSize is the number of rows inserted per transaction during +// seeding. 5000 keeps each multi-row INSERT comfortably below SQLite's +// default 32766 host-parameter limit (5000 * 4 = 20000 < 32766) and gives a +// reasonable bytes-per-syscall ratio for fsync-free seed inserts. +const perfSeedBatchSize = 5000 + +func TestPerfQueries(t *testing.T) { + dir := *perfTestDir + if dir == "" { + if v := os.Getenv("GOCACHED_PERFTEST_DIR"); v != "" { + dir = v + } else { + // Default: per-test temp dir so a normal `go test` doesn't + // leave a seeded DB hanging around. Operators running real-scale + // measurements set -performance-test-dir explicitly to get + // seed reuse across runs. + dir = t.TempDir() + } + } + if err := os.MkdirAll(dir, 0o750); err != nil { + t.Fatal(err) + } + t.Logf("perf test dir: %s (rows=%d)", dir, *perfTestRows) + + srv, err := NewServer( + WithDir(dir), + WithLogf(t.Logf), + withoutBackgroundLoops(), + WithShardPrefixLen(2), + ) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { srv.Close() }) + + n := *perfTestRows + if err := perfSeed(t, srv, n); err != nil { + t.Fatalf("seed: %v", err) + } + + // Confirm the eviction query plan still uses idx_actions_access even on + // a 250M-row table — that's the whole point of INDEXED BY, but it's + // worth re-asserting at scale because the planner picks differently as + // table cardinality grows. + t.Run("EvictionQueryPlan_atScale", func(t *testing.T) { + row := srv.db.QueryRow("EXPLAIN QUERY PLAN "+evictionCandidateQuery, 0, cleanupBatchSize) + var id, parent, notused int + var detail string + if err := row.Scan(&id, &parent, ¬used, &detail); err != nil { + t.Fatalf("EXPLAIN QUERY PLAN: %v", err) + } + t.Logf("plan: %s", detail) + if !strings.Contains(detail, "idx_actions_access") { + t.Errorf("plan does not use idx_actions_access: %q", detail) + } + if strings.Contains(detail, "SCAN Actions") || strings.Contains(detail, "SCAN TABLE Actions") { + t.Errorf("plan does a table scan of Actions: %q", detail) + } + }) + + timed := func(name string, f func() error) time.Duration { + t.Helper() + start := time.Now() + if err := f(); err != nil { + t.Errorf("%s failed: %v", name, err) + return 0 + } + d := time.Since(start) + t.Logf("%-44s %v", name, d) + return d + } + + // usageStats is a plain in-memory aggregate read; it doesn't touch + // SQLite at all. The benchmark is here to lock in that property: any + // future regression that re-introduces a per-call DB scan would blow + // past the sub-millisecond threshold. + aggRead := timed("usageStats (in-memory aggregate read)", func() error { + _, err := srv.usageStats() + return err + }) + + perShard := timed(fmt.Sprintf("scanShard (1 shard, ~%dM rows)", n/256/1_000_000), func() error { + _, err := srv.scanShard(context.Background(), 0) + return err + }) + + cleanup := timed(fmt.Sprintf("evictOldestActions (%d-row batch)", cleanupBatchSize), func() error { + _, err := srv.evictOldestActions(context.Background(), math.MaxInt64, cleanupBatchSize, math.MaxInt64) + return err + }) + + // Legacy world-scoped GROUP BY: the query gocached used to run every + // 5 minutes. It's here as a baseline so reviewers can see how much + // faster the per-shard approach is and as an early-warning if SQLite + // ever gets dramatically faster at this shape. + legacy := timed(fmt.Sprintf("legacy GROUP BY cleanCandidates (%d-row LIMIT)", cleanupBatchSize), func() error { + rows, err := srv.db.Query(` + SELECT b.BlobID, MAX(a.AccessTime), b.StoredSize + FROM Blobs b LEFT JOIN Actions a ON b.BlobID = a.BlobID + GROUP BY b.BlobID + HAVING MAX(a.AccessTime) <= ? + ORDER BY MAX(a.AccessTime) + LIMIT ?`, int64(math.MaxInt64), cleanupBatchSize) + if err != nil { + return err + } + defer rows.Close() + var blobID, accessTime, storedSize int64 + for rows.Next() { + if err := rows.Scan(&blobID, &accessTime, &storedSize); err != nil { + return err + } + } + return rows.Err() + }) + + if cleanup > 0 { + t.Logf("legacy/incremental cleanup ratio: %.1fx slower", legacy.Seconds()/cleanup.Seconds()) + } + + // Wall-clock regression thresholds. These are deliberately loose for + // laptop/dev hardware; tighten on the CI box once we have a baseline + // run there. The intent is "scream loudly if any of these queries + // quietly grows by an order of magnitude," not "lock in a tight SLO." + if aggRead > 1*time.Millisecond { + t.Errorf("usageStats took %v; want sub-millisecond (it should not hit SQLite)", aggRead) + } + if perShard > 30*time.Second { + t.Errorf("scanShard took %v on ~%dM rows; want < 30s", perShard, n/256/1_000_000) + } + if cleanup > 2*time.Second { + t.Errorf("evictOldestActions took %v on %d-row batch; want < 2s (idx_actions_access range scan)", cleanup, cleanupBatchSize) + } +} + +// perfSeed brings the DB up to at least n (Blob, Action) pairs. If the DB +// already has that many it returns immediately so re-runs are fast; if it's +// short, only the deficit is inserted. Per-batch work is wrapped in one +// transaction so a Ctrl-C resumes cleanly at the last complete batch. +func perfSeed(t *testing.T, srv *Server, n int) error { + t.Helper() + + var seq sql.NullInt64 + err := srv.db.QueryRow(`SELECT seq FROM sqlite_sequence WHERE name = 'Blobs'`).Scan(&seq) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("query sqlite_sequence: %w", err) + } + cur := int(seq.Int64) + if cur >= n { + t.Logf("DB already has %d Blobs (>= %d target); reusing seed", cur, n) + return nil + } + + t.Logf("seeding rows %d..%d", cur, n) + // PRAGMA synchronous=OFF makes the seed dramatically faster at the + // cost of crash safety. Acceptable for a synthetic test DB; the + // process explicitly runs with no other writers (background loops are + // disabled via withoutBackgroundLoops). The pragma resets on Close. + if _, err := srv.db.Exec(`PRAGMA synchronous = OFF`); err != nil { + return fmt.Errorf("PRAGMA synchronous: %w", err) + } + // 512 MiB page cache so the bulk inserts thrash less. + if _, err := srv.db.Exec(`PRAGMA cache_size = -524288`); err != nil { + return fmt.Errorf("PRAGMA cache_size: %w", err) + } + + start := time.Now() + lastLog := start + for i := cur; i < n; i += perfSeedBatchSize { + end := min(i+perfSeedBatchSize, n) + if err := perfSeedBatch(srv, i, end); err != nil { + return fmt.Errorf("seed batch [%d, %d): %w", i, end, err) + } + if time.Since(lastLog) > 30*time.Second { + elapsed := time.Since(start) + rate := float64(end-cur) / elapsed.Seconds() + remaining := time.Duration(float64(n-end) / rate * float64(time.Second)) + t.Logf("seed progress: %d / %d (%.1f%%, elapsed %v, ETA %v)", + end, n, 100*float64(end)/float64(n), + elapsed.Round(time.Second), remaining.Round(time.Second)) + lastLog = time.Now() + } + } + t.Logf("seeded %d rows in %v", n-cur, time.Since(start).Round(time.Second)) + return nil +} + +// perfSeedBatch inserts (Blob, Action) pairs for indices [start, end) in a +// single transaction. Synthetic SHA256s come from sha256(itoa(i)) so the +// resulting hashes are uniformly distributed across all shards. BlobIDs are +// the AUTOINCREMENT-assigned sequential values; we know them in advance +// because seeding always continues from sqlite_sequence's max. +func perfSeedBatch(srv *Server, start, end int) error { + tx, err := srv.db.Begin() + if err != nil { + return err + } + defer tx.Rollback() + + rowCount := end - start + + var sb strings.Builder + sb.Grow(64 + rowCount*8) + sb.WriteString("INSERT INTO Blobs (SHA256, StoredSize, UncompressedSize) VALUES ") + args := make([]any, 0, rowCount*3) + for i := start; i < end; i++ { + if i > start { + sb.WriteByte(',') + } + sb.WriteString("(?,?,?)") + h := sha256.Sum256([]byte(strconv.Itoa(i))) + args = append(args, hex.EncodeToString(h[:]), int64(100), int64(100)) + } + if _, err := tx.Exec(sb.String(), args...); err != nil { + return fmt.Errorf("insert Blobs: %w", err) + } + + sb.Reset() + sb.Grow(96 + rowCount*16) + sb.WriteString("INSERT INTO Actions (NamespaceID, ActionID, BlobID, AltOutputID, CreateTime, AccessTime) VALUES ") + args = args[:0] + for i := start; i < end; i++ { + if i > start { + sb.WriteByte(',') + } + sb.WriteString("(0,?,?,'',?,?)") + h := sha256.Sum256([]byte(strconv.Itoa(i))) + args = append(args, + hex.EncodeToString(h[:]), + int64(i+1), // BlobID; matches the AUTOINCREMENT sequence + int64(1234+i), // CreateTime + int64(1234+i), // AccessTime — monotonically increasing makes cleanup deterministic + ) + } + if _, err := tx.Exec(sb.String(), args...); err != nil { + return fmt.Errorf("insert Actions: %w", err) + } + + return tx.Commit() +} diff --git a/gocache/gocached/gocached_test.go b/gocache/gocached/gocached_test.go index 0aad646..bcb0b89 100644 --- a/gocache/gocached/gocached_test.go +++ b/gocache/gocached/gocached_test.go @@ -78,6 +78,13 @@ func (t *tester) mkClient() *cachers.HTTPClient { func (st *tester) usageStats() *usageStats { st.t.Helper() + // Stats are maintained by an async per-shard stats loop in production. Tests + // need a deterministic aggregate, so flush pending access-time bumps and + // rotate every shard before reading. + st.srv.flushAccessTimeBumps() + if err := st.srv.scanAllShards(context.Background()); err != nil { + st.t.Fatalf("scanAllShards: %v", err) + } stats, err := st.srv.usageStats() if err != nil { st.t.Fatalf("usageStats: %v", err) @@ -85,13 +92,26 @@ func (st *tester) usageStats() *usageStats { return stats } +// cleanOldObjects runs cleanup until it's a no-op, so tests get a fully +// drained eviction in one call. Updates the aggregate (via a forced shard +// rescan) between batches so the cleanupTick's "are we over budget" check +// sees the post-eviction state, not stale pre-eviction shard data. func (st *tester) cleanOldObjects() countAndSize { st.t.Helper() - stats, err := st.srv.cleanOldObjects(st.usageStats()) - if err != nil { - st.t.Fatalf("cleanOldObjects: %v", err) + var total countAndSize + ctx := context.Background() + for { + _ = st.usageStats() + res, err := st.srv.cleanupTick(ctx) + if err != nil { + st.t.Fatalf("cleanupTick: %v", err) + } + if res.Count == 0 { + return total + } + total.Count += res.Count + total.Size += res.Size } - return stats } func (st *tester) diskFiles() []string { @@ -205,6 +225,18 @@ func withClock(clk func() time.Time) ServerOption { } } +// withoutBackgroundLoops disables every periodic goroutine that start would +// otherwise launch (stats loop, cleanup, checkpoint, DB-size metrics). +// Tests need this because the mocked clock makes scannedAt collide across +// concurrent scans, so a background scan can overwrite a freshly-PUT +// shard with a stale-snapshot empty result. Tests drive scanShard / +// cleanupTick / etc. synchronously instead. +func withoutBackgroundLoops() ServerOption { + return func(cfg *Server) { + cfg.disableBackgroundLoops = true + } +} + func newServerTester(t testing.TB, extraOpts ...ServerOption) *tester { st := &tester{ t: t, @@ -216,6 +248,13 @@ func newServerTester(t testing.TB, extraOpts ...ServerOption) *tester { WithLogf(t.Logf), WithVerbose(true), withClock(st.now), + withoutBackgroundLoops(), + // Default to 16 shards instead of production's 256 to keep + // scanAllShards fast in tests (each shard is a separate SQL + // round-trip; on near-empty tables overhead dominates). Tests + // that need a specific shardPrefixLen pass WithShardPrefixLen + // in extraOpts to override. + WithShardPrefixLen(1), } srv, err := NewServer(append(opts, extraOpts...)...) if err != nil { @@ -336,11 +375,9 @@ func TestServer(t *testing.T) { st.wantGet(c3, testActionID, testOutputID, testObjectValue) st.wantMetric(&st.srv.m.GetAccessBumps, 1) - // Get usage stats. - stats, err := st.srv.usageStats() - if err != nil { - t.Fatalf("usageStats: %v", err) - } + // Get usage stats. The tester helper flushes pending access-time bumps + // and rotates every shard so the aggregate is deterministic. + stats := st.usageStats() bigStored := int64(len(testObjectValueBig)) // below lz4CompressThreshold, stored uncompressed totalSize := int64(9) + bigStored // 9 (inline "test data") + big + 0 (empty) want := &usageStats{ @@ -363,64 +400,204 @@ func TestServer(t *testing.T) { st.advanceClock(relAtimeSeconds * 2 * time.Second) // advance clock by 2 days } -func TestCleanCandidates(t *testing.T) { +// TestEvictionQueryPlan is the lockdown test for the eviction query: it +// runs EXPLAIN QUERY PLAN and asserts SQLite uses idx_actions_access and +// never does a SCAN of the Actions table. INDEXED BY in the query already +// forces this at runtime, but the test catches schema or query drift early +// (and gives a concrete failure message instead of a runtime SQL error). +func TestEvictionQueryPlan(t *testing.T) { st := newServerTester(t) + row := st.srv.db.QueryRow("EXPLAIN QUERY PLAN "+evictionCandidateQuery, 0, 1) + var id, parent, notused int + var detail string + if err := row.Scan(&id, &parent, ¬used, &detail); err != nil { + t.Fatalf("EXPLAIN QUERY PLAN: %v", err) + } + t.Logf("plan: %s", detail) + if !strings.Contains(detail, "idx_actions_access") { + t.Errorf("plan does not use idx_actions_access: %q", detail) + } + if strings.Contains(detail, "SCAN Actions") || strings.Contains(detail, "SCAN TABLE Actions") { + t.Errorf("plan does a table scan of Actions: %q", detail) + } +} - // Populate some data. - c1 := st.mkClient() - st.wantPut(c1, "0001", "9901", "1") +func TestEvictOldestActions(t *testing.T) { + st := newServerTester(t) + ctx := context.Background() + + // Four unique blobs (different content => different SHA256 => different + // BlobID), inserted 24h apart so AccessTime order matches insert order. + c := st.mkClient() + st.wantPut(c, "0001", "9901", "1") st.advanceClock(24 * time.Hour) - st.wantPut(c1, "0002", "9902", "22") + st.wantPut(c, "0002", "9902", "22") st.advanceClock(24 * time.Hour) - st.wantPut(c1, "0003", "9903", "333") + st.wantPut(c, "0003", "9903", "333") st.advanceClock(24 * time.Hour) - st.wantPut(c1, "0004", "9904", strings.Repeat("x", smallObjectSize+1)) + st.wantPut(c, "0004", "9904", "4444") - const day = 24 * time.Hour + // All four should still be on disk. + if got, want := len(st.diskFiles()), 0; got != want { + // All four are <= smallObjectSize, so they're inline in the DB, + // not on disk. Confirms the test isn't accidentally on the disk path. + t.Errorf("diskFiles before evict: %d, want %d (all should be inline)", got, want) + } - tests := []struct { - maxAge time.Duration - limit int64 - want []cleanCandidate - }{ - { - maxAge: 0, - limit: 100, - want: []cleanCandidate{ - {BlobID: 1, Age: 3 * day, StoredSize: 1}, - {BlobID: 2, Age: 2 * day, StoredSize: 2}, - {BlobID: 3, Age: 1 * day, StoredSize: 3}, - {BlobID: 4, Age: 0, StoredSize: smallObjectSize + 1}, // below lz4CompressThreshold, stored uncompressed - }, - }, - { - maxAge: 25 * time.Hour, - limit: 100, - want: []cleanCandidate{ - {BlobID: 1, Age: 3 * day, StoredSize: 1}, - {BlobID: 2, Age: 2 * day, StoredSize: 2}, - }, - }, - { - maxAge: 0, - limit: 2, - want: []cleanCandidate{ - {BlobID: 1, Age: 3 * day, StoredSize: 1}, - {BlobID: 2, Age: 2 * day, StoredSize: 2}, - }, - }, + // Cutoff lets only the two oldest (1 and 22) through. limit=10 is generous. + res, err := st.srv.evictOldestActions(ctx, st.srv.now().Add(-25*time.Hour).Unix(), 10, math.MaxInt64) + if err != nil { + t.Fatalf("evictOldestActions: %v", err) + } + // 1:1 blobs => evicting Actions evicts Blobs too. + if want := (countAndSize{Count: 2, Size: 3}); res != want { + t.Errorf("evict result = %+v, want %+v", res, want) } - for _, tt := range tests { - t.Run(fmt.Sprintf("maxAge=%v,limit=%d", tt.maxAge, tt.limit), func(t *testing.T) { - candidates, err := st.srv.cleanCandidates(tt.maxAge, tt.limit) - if err != nil { - t.Fatal(err) - } - if diff := cmp.Diff(candidates, tt.want); diff != "" { - t.Errorf("cleanCandidates mismatch (-got +want):\n%s", diff) - } - }) + // Now no cutoff (math.MaxInt64) but limit=1: just the next-oldest ("333"). + res, err = st.srv.evictOldestActions(ctx, math.MaxInt64, 1, math.MaxInt64) + if err != nil { + t.Fatalf("evictOldestActions: %v", err) + } + if want := (countAndSize{Count: 1, Size: 3}); res != want { + t.Errorf("evict result = %+v, want %+v", res, want) + } +} + +// TestActionLRUSharedBlob locks in the Action-LRU semantic on shared blobs: +// when two Actions point at the same Blob, evicting the old Action leaves +// the Blob alive (because the fresh Action still references it). +func TestActionLRUSharedBlob(t *testing.T) { + st := newServerTester(t) + ctx := context.Background() + + c := st.mkClient() + const content = "shared-blob-content" + // Two distinct ActionIDs but identical content => identical SHA256 => + // the same BlobID. Action 0001 is inserted "two days ago", 0002 today. + st.wantPut(c, "0001", "9901", content) + st.advanceClock(48 * time.Hour) + st.wantPut(c, "0002", "9902", content) + + // Evict everything older than 25h. Action-LRU: the old 0001 Action + // must go, but the Blob stays because 0002 still references it. + res, err := st.srv.evictOldestActions(ctx, st.srv.now().Add(-25*time.Hour).Unix(), 10, math.MaxInt64) + if err != nil { + t.Fatalf("evictOldestActions: %v", err) + } + if got, want := res.Count, int64(0); got != want { + t.Errorf("evicted Blob count = %d, want %d (Blob still has a fresher Action)", got, want) + } + var actionCount int + if err := st.srv.db.QueryRow(`SELECT COUNT(*) FROM Actions`).Scan(&actionCount); err != nil { + t.Fatal(err) + } + if actionCount != 1 { + t.Errorf("Actions count after evict = %d, want 1 (0001 evicted, 0002 alive)", actionCount) + } + var blobCount int + if err := st.srv.db.QueryRow(`SELECT COUNT(*) FROM Blobs`).Scan(&blobCount); err != nil { + t.Fatal(err) + } + if blobCount != 1 { + t.Errorf("Blobs count after evict = %d, want 1 (Blob shared with 0002, stays)", blobCount) + } +} + +func TestShardDeltaPUT(t *testing.T) { + st := newServerTester(t) + + // Before any PUT: delta is zero. + if c, b := st.srv.sumShardDeltas(); c != 0 || b != 0 { + t.Errorf("delta before PUT = (%d, %d), want (0, 0)", c, b) + } + + // Two PUTs of unique content. Each is small enough to be stored inline + // so StoredSize == len(content). + c := st.mkClient() + st.wantPut(c, "0001", "9901", "hello") + st.wantPut(c, "0002", "9902", "world!!") + + if got, want := func() int64 { c, _ := st.srv.sumShardDeltas(); return c }(), int64(2); got != want { + t.Errorf("delta count after 2 PUTs = %d, want %d", got, want) + } + if got, want := func() int64 { _, b := st.srv.sumShardDeltas(); return b }(), int64(5+7); got != want { + t.Errorf("delta bytes after 2 PUTs = %d, want %d", got, want) + } + + // Duplicate PUT (same actionID/outputID/content) goes through INSERT OR + // IGNORE: affected=0, so the delta must NOT advance. + st.wantPut(c, "0001", "9901", "hello") + if got, _ := st.srv.sumShardDeltas(); got != 2 { + t.Errorf("delta count after duplicate PUT = %d, want 2 (dup must not double-count)", got) + } +} + +func TestShardDeltaZeroedAfterScan(t *testing.T) { + // The whole point of the preDelta dance in scanAndPersistShard is that + // once a shard is rescanned, anything that was reflected in the new + // persisted snapshot leaves the delta. Verify it: PUT, scanAllShards, + // delta == 0. + st := newServerTester(t) + c := st.mkClient() + st.wantPut(c, "0001", "9901", "hello") + st.wantPut(c, "0002", "9902", "world!!") + + if cnt, b := st.srv.sumShardDeltas(); cnt == 0 && b == 0 { + t.Fatalf("delta is already zero before scan; test isn't exercising the post-scan zeroing") + } + if err := st.srv.scanAllShards(context.Background()); err != nil { + t.Fatal(err) + } + if cnt, b := st.srv.sumShardDeltas(); cnt != 0 || b != 0 { + t.Errorf("delta after scanAllShards = (%d, %d), want (0, 0)", cnt, b) + } +} + +func TestShardDeltaDrivesSizeCleanup(t *testing.T) { + // Dead-reckoning's reason for existing: cleanup must react to PUTs that + // happen between scans. Scenario: scan first so the persisted + // aggregate is at zero, then PUT enough to go over the limit *without* + // a scan, then assert cleanupTick fires. + st := newServerTester(t) + st.srv.maxSize = 8 + ctx := context.Background() + + // Empty cache: aggregate scanned, delta zero, cleanupTick no-op. + if err := st.srv.scanAllShards(ctx); err != nil { + t.Fatal(err) + } + if res, err := st.srv.cleanupTick(ctx); err != nil || res.Count != 0 { + t.Fatalf("cleanupTick on empty: %+v, err=%v; want zero", res, err) + } + + // PUT 10 bytes total. NO scanAllShards — the persisted aggregate still + // says zero. cleanupTick must consult the delta and decide to evict. + c := st.mkClient() + st.wantPut(c, "0001", "9901", "1") + st.advanceClock(time.Second) + st.wantPut(c, "0002", "9902", "22") + st.advanceClock(time.Second) + st.wantPut(c, "0003", "9903", "333") + st.advanceClock(time.Second) + st.wantPut(c, "0004", "9904", "4444") + st.advanceClock(time.Second) + + if us := st.srv.lastUsage.Load(); us.All().Size != 0 { + t.Fatalf("persisted aggregate is non-zero (%v); the test premise (no scan yet) is broken", us.All()) + } + if _, b := st.srv.sumShardDeltas(); b != 10 { + t.Fatalf("delta bytes = %d, want 10 (10 = 1+2+3+4)", b) + } + res, err := st.srv.cleanupTick(ctx) + if err != nil { + t.Fatalf("cleanupTick: %v", err) + } + // maxSize=8, delta-only total=10, so we need at least 2 bytes freed. + // The oldest are "1" (1B) and "22" (2B), summing to 3 ≥ 2 — stop after + // evicting both. Same as TestCleanOldObjectsBySize. + if want := (countAndSize{Count: 2, Size: 3}); res != want { + t.Errorf("cleanupTick result = %+v, want %+v", res, want) } } @@ -614,6 +791,277 @@ func TestLZ4Storage(t *testing.T) { } } +func TestShardPrefix(t *testing.T) { + // Verify the %0*x zero-padding contract: the width comes from the * + // argument and pads with leading zeros to that width. A typo (%0x, + // %x, %2x) would silently break shard keys. + cases := []struct { + prefixLen int + idx int + want shardPrefix + }{ + {1, 0, "0"}, + {1, 15, "f"}, + {2, 0, "00"}, + {2, 1, "01"}, + {2, 15, "0f"}, + {2, 16, "10"}, + {2, 255, "ff"}, + {3, 0, "000"}, + {3, 16, "010"}, + {3, 4095, "fff"}, + {4, 0, "0000"}, + {4, 65535, "ffff"}, + } + for _, c := range cases { + srv := &Server{shardPrefixLen: c.prefixLen} + if got := srv.shardPrefix(c.idx); got != c.want { + t.Errorf("shardPrefix(prefixLen=%d, idx=%d) = %q, want %q", + c.prefixLen, c.idx, got, c.want) + } + } +} + +func TestShardRange(t *testing.T) { + // shardRange's last-shard sentinel ("g" chars) is the easiest place for a + // future refactor to silently drop the final shard from the stats loop. + cases := []struct { + prefixLen int + idx int + wantLo, wantHi string + }{ + // First shard: lo = all zeros prefix, hi = next prefix. + {1, 0, "0", "1"}, + {2, 0, "00", "01"}, + // Middle shard. + {2, 15, "0f", "10"}, + {2, 16, "10", "11"}, + // Last shard: hi is the "g"-sentinel that sorts after all hex. + {1, 15, "f", "g"}, + {2, 255, "ff", "gg"}, + {3, 4095, "fff", "ggg"}, + {4, 65535, "ffff", "gggg"}, + } + for _, c := range cases { + srv := &Server{shardPrefixLen: c.prefixLen} + gotLo, gotHi := srv.shardRange(c.idx) + if gotLo != c.wantLo || gotHi != c.wantHi { + t.Errorf("shardRange(prefixLen=%d, idx=%d) = (%q, %q), want (%q, %q)", + c.prefixLen, c.idx, gotLo, gotHi, c.wantLo, c.wantHi) + } + } +} + +func TestLoadShardStats(t *testing.T) { + // Exercise the three branches that decide whether a persisted row gets + // kept or dropped: prefix-length mismatch, cohort mismatch, and a clean + // match. A dropped row stays in the table on disk but doesn't contribute + // to the in-memory aggregate. + // Override the tester default (shardPrefixLen=1) because this test + // asserts the wrong-prefix-length and wrong-cohort branches using + // 2-character prefixes. + st := newServerTester(t, WithShardPrefixLen(2)) + ctx := context.Background() + + if got := st.srv.numShards(); got != 256 { + t.Fatalf("numShards = %d, want 256", got) + } + + // One clean shard scan so we have a row in BlobShardStats whose schema + // matches the current server. + if err := st.srv.scanAndPersistShard(ctx, 0); err != nil { + t.Fatalf("scanAndPersistShard: %v", err) + } + + // Hand-write two more rows that should be rejected on the next load. + mustExec := func(query string, args ...any) { + t.Helper() + if _, err := st.srv.db.ExecContext(ctx, query, args...); err != nil { + t.Fatalf("exec %q: %v", query, err) + } + } + mustExec(`INSERT INTO BlobShardStats (Prefix, ScannedAt, StatsJSON) VALUES (?, ?, ?)`, + "abc", 1, `{"ActionsLE":{}}`) // wrong prefix length + mustExec(`INSERT INTO BlobShardStats (Prefix, ScannedAt, StatsJSON) VALUES (?, ?, ?)`, + "01", 1, `{"ActionsLE":{"60000000000":{"Count":1,"Size":1}}}`) // wrong cohort set + + if err := st.srv.loadShardStats(ctx); err != nil { + t.Fatalf("loadShardStats: %v", err) + } + if got, want := len(st.srv.shardStats), 1; got != want { + t.Errorf("shardStats has %d entries, want %d (only the clean prefix=00 row should survive)", got, want) + } + if _, ok := st.srv.shardStats["00"]; !ok { + t.Errorf("shardStats missing prefix=00") + } +} + +func TestShardScanDurationPercentiles(t *testing.T) { + srv := &Server{ + shardStats: map[shardPrefix]*shardSnapshot{ + "00": {stats: &usageStats{QueryDuration: 10 * time.Millisecond}}, + "01": {stats: &usageStats{QueryDuration: 20 * time.Millisecond}}, + "02": {stats: &usageStats{QueryDuration: 30 * time.Millisecond}}, + "03": {stats: &usageStats{QueryDuration: 40 * time.Millisecond}}, + // Zero-duration shard is excluded so freshly-loaded-but-not-yet- + // observed entries don't skew the low percentile to zero. + "04": {stats: &usageStats{QueryDuration: 0}}, + }, + } + p25, p50, p90, n := srv.shardScanDurationPercentiles() + if n != 4 { + t.Errorf("n = %d, want 4 (zero-duration shard should be excluded)", n) + } + // With sorted = [10,20,30,40] ms and rank-index picks, p25 -> idx 1, + // p50 -> idx 2, p90 -> idx 3. + if p25 != 20*time.Millisecond { + t.Errorf("p25 = %v, want 20ms", p25) + } + if p50 != 30*time.Millisecond { + t.Errorf("p50 = %v, want 30ms", p50) + } + if p90 != 40*time.Millisecond { + t.Errorf("p90 = %v, want 40ms", p90) + } + + empty := &Server{shardStats: map[shardPrefix]*shardSnapshot{}} + if p25, p50, p90, n := empty.shardScanDurationPercentiles(); p25 != 0 || p50 != 0 || p90 != 0 || n != 0 { + t.Errorf("empty server: got (%v, %v, %v, n=%d), want all zero", p25, p50, p90, n) + } +} + +func TestShardFreshness(t *testing.T) { + now := time.Unix(1_000_000, 0) + at := func(secsAgo int64) time.Time { return now.Add(-time.Duration(secsAgo) * time.Second) } + + srv := &Server{ + shardStats: map[shardPrefix]*shardSnapshot{ + "00": {stats: &usageStats{}, scannedAt: at(10)}, // <= 1m + "01": {stats: &usageStats{}, scannedAt: at(200)}, // <= 5m + "02": {stats: &usageStats{}, scannedAt: at(800)}, // <= 15m + "03": {stats: &usageStats{}, scannedAt: at(3000)}, // <= 1h + "04": {stats: &usageStats{}, scannedAt: at(80_000)}, // <= 24h + "05": {stats: &usageStats{}, scannedAt: at(200_000)}, // older than 24h + }, + } + counts, oldest, total := srv.shardFreshness(now) + if total != 6 { + t.Errorf("total = %d, want 6", total) + } + if got, want := oldest, at(200_000); !got.Equal(want) { + t.Errorf("oldest = %v, want %v", got, want) + } + want := map[time.Duration]int{ + 1 * time.Minute: 1, // just the 10s-old shard + 5 * time.Minute: 2, // adds the 200s-old + 15 * time.Minute: 3, // adds the 800s-old + 1 * time.Hour: 4, // adds the 3000s-old + 6 * time.Hour: 4, // unchanged (next is 80000s = ~22h) + 24 * time.Hour: 5, // adds the 22h-old; the 200000s (~2.3d) one is still excluded + } + for d, c := range want { + if got := counts[d]; got != c { + t.Errorf("counts[%v] = %d, want %d", d, got, c) + } + } +} + +func TestShardHealthGauges(t *testing.T) { + st := newServerTester(t) + ctx := context.Background() + + // Before any scan: unscanned should equal numShards, oldest age 0. + body := scrapeMetrics(t, st) + if got, want := promGauge(t, body, "gocached_shard_stats_unscanned"), float64(st.srv.numShards()); got != want { + t.Errorf("unscanned before scan = %v, want %v", got, want) + } + if got := promGauge(t, body, "gocached_shard_stats_oldest_age_seconds"); got != 0 { + t.Errorf("oldest_age_seconds before scan = %v, want 0", got) + } + + // Scan one shard at t=0, then advance the clock and re-scrape. The + // oldest-age gauge should reflect the advance; unscanned should drop by 1. + if err := st.srv.scanAndPersistShard(ctx, 0); err != nil { + t.Fatalf("scanAndPersistShard: %v", err) + } + st.advanceClock(90 * time.Second) + + body = scrapeMetrics(t, st) + if got, want := promGauge(t, body, "gocached_shard_stats_unscanned"), float64(st.srv.numShards()-1); got != want { + t.Errorf("unscanned after 1 scan = %v, want %v", got, want) + } + if got, want := promGauge(t, body, "gocached_shard_stats_oldest_age_seconds"), float64(90); got != want { + t.Errorf("oldest_age_seconds after 90s advance = %v, want %v", got, want) + } +} + +// scrapeMetrics returns the /metrics body served by the test server's +// metrics handler. +func scrapeMetrics(t *testing.T, st *tester) string { + t.Helper() + req := httptest.NewRequest("GET", "/metrics", nil) + rec := httptest.NewRecorder() + st.srv.metricsHandler.ServeHTTP(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("/metrics status = %d, body=%s", rec.Code, rec.Body.String()) + } + return rec.Body.String() +} + +// promGauge returns the value of the named gauge in a Prometheus text +// exposition body, fatal-ing if the line isn't present. Matches lines like +// " " but not commented HELP/TYPE lines. +func promGauge(t *testing.T, body, name string) float64 { + t.Helper() + for line := range strings.Lines(body) { + line = strings.TrimRight(line, "\n") + if strings.HasPrefix(line, "#") { + continue + } + k, v, ok := strings.Cut(line, " ") + if !ok || k != name { + continue + } + var f float64 + if _, err := fmt.Sscanf(v, "%g", &f); err != nil { + t.Fatalf("parsing %q value %q: %v", name, v, err) + } + return f + } + t.Fatalf("metric %q not found in /metrics body:\n%s", name, body) + return 0 +} + +func TestServeUsage(t *testing.T) { + st := newServerTester(t) + c := st.mkClient() + st.wantPut(c, "0001", "9901", "hello") + // Force a scan pass so the freshness section has data to render. + if err := st.srv.scanAllShards(context.Background()); err != nil { + t.Fatalf("scanAllShards: %v", err) + } + + req := httptest.NewRequest("GET", "/usage", nil) + rec := httptest.NewRecorder() + st.srv.serveUsage(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rec.Code) + } + body := rec.Body.String() + // Each new section should render at least its heading. Asserting on the + // headings (rather than exact bytes) lets the formatting tweak without + // breaking the test. + for _, want := range []string{ + "Actions by access-time age", + "Shard scan freshness", + "Per-shard scan duration", + } { + if !strings.Contains(body, want) { + t.Errorf("/usage body missing %q\nbody:\n%s", want, body) + } + } +} + func TestWALCheckpoint(t *testing.T) { st := newServerTester(t) ctx := context.Background() From e4be1fd1ed5226836b767c9e99501044801a1c34 Mon Sep 17 00:00:00 2001 From: Tom Proctor Date: Fri, 12 Jun 2026 10:08:52 +0100 Subject: [PATCH 44/67] gocached: support namespaces (bradfitz/go-tool-cache#34) Remove the special-case concept of "global writes" and instead allow callers to provide a namespace mapping function that makes policy decisions about which clients are globally trusted and which clients should be isolated from each other. As a result, all clients are now able to write, but not necessarily into the shared global namespace. All clients can still read from the global namespace, as well as their own. The Namespaces table already existed in the schema, but we drop the lowercase constraint. However, there are some breaking changes in the package API, WithJWTAuth now takes issuer URLs only and policy moves into the new WithNamespaceMapping option. cmd/gocached implements the spirit of the old API in terms of a namespace mapping function, with the main difference that it now allows writes if you don't have the global claims, but just into your own isolated namespace. Updates tailscale/corp#38092 Signed-off-by: Tom Proctor Migrated-from: bradfitz/go-tool-cache@acf1bfd17e4fcc08c9cfb0570616ee07b2f2b829 --- cmd/gocached/gocached.go | 31 ++- gocache/gocached/gocached.go | 321 +++++++++++++++++++---------- gocache/gocached/gocached_test.go | 331 ++++++++++++++++++------------ 3 files changed, 434 insertions(+), 249 deletions(-) diff --git a/cmd/gocached/gocached.go b/cmd/gocached/gocached.go index 79a7cfb..f9c191b 100644 --- a/cmd/gocached/gocached.go +++ b/cmd/gocached/gocached.go @@ -9,7 +9,6 @@ import ( "flag" "fmt" "log" - "maps" "net" "net/http" "os" @@ -77,15 +76,27 @@ func main() { log.Fatal("must specify --jwt-claim at least once when --jwt-issuer is set") } - globalClaims := map[string]string{} - maps.Copy(globalClaims, jwtClaims) - maps.Copy(globalClaims, globalJWTClaims) - - opts = append(opts, gocached.WithJWTAuth(gocached.JWTIssuerConfig{ - Issuer: *jwtIssuer, - RequiredClaims: jwtClaims, - GlobalWriteClaims: globalClaims, - })) + opts = append(opts, + gocached.WithJWTAuth(*jwtIssuer), + gocached.WithNamespaceMapping(func(claims map[string]any) (gocached.Namespace, error) { + var ns gocached.Namespace + for k, want := range jwtClaims { + if got := claims[k]; got != want { + return "", fmt.Errorf("claim %q = %v, want %v", k, got, want) + } + if ns != "" { + ns += "," + } + ns += gocached.Namespace(fmt.Sprintf("%s=%s", k, want)) + } + for k, want := range globalJWTClaims { + if got := claims[k]; got != want { + return ns, nil + } + } + return gocached.GlobalNamespace, nil + }), + ) } srv, err := gocached.NewServer(opts...) diff --git a/gocache/gocached/gocached.go b/gocache/gocached/gocached.go index 156f11e..86f3aae 100644 --- a/gocache/gocached/gocached.go +++ b/gocache/gocached/gocached.go @@ -191,7 +191,7 @@ CREATE UNIQUE INDEX IF NOT EXISTS idx_blobs_sha256 ON Blobs(SHA256); CREATE TABLE IF NOT EXISTS Namespaces ( NamespaceID INTEGER PRIMARY KEY AUTOINCREMENT, - Namespace TEXT NOT NULL UNIQUE CHECK (Namespace = lower(Namespace)) + Namespace TEXT NOT NULL UNIQUE ) STRICT; -- BlobShardStats persists per-shard usage histograms across restarts so @@ -230,6 +230,15 @@ func openDB(dbDir string) (*sql.DB, error) { db.SetMaxOpenConns(numConns) db.SetMaxIdleConns(numConns) db.SetConnMaxLifetime(0) // no limit + var ddl string + err = db.QueryRow(`SELECT sql FROM sqlite_master WHERE type='table' AND name='Namespaces'`).Scan(&ddl) + if err == nil && strings.Contains(ddl, "lower(Namespace)") { + // Drop the Namespaces table _only_ if it has the lowercase constraint so it + // can be recreated without it. + if _, err := db.Exec(`DROP TABLE Namespaces`); err != nil { + return nil, fmt.Errorf("dropping Namespaces lowercase constraint: %w", err) + } + } if _, err := db.Exec(schema); err != nil { return nil, err } @@ -365,6 +374,24 @@ func (srv *Server) start() error { Buckets: prometheus.DefBuckets, }) + // Fill the namespace ID cache. + rows, err := srv.db.Query("SELECT NamespaceID, Namespace FROM Namespaces") + if err != nil { + return err + } + defer rows.Close() + for rows.Next() { + var id int64 + var ns string + if err := rows.Scan(&id, &ns); err != nil { + return err + } + srv.namespaces[Namespace(ns)] = id + } + if err := rows.Err(); err != nil { + return err + } + reg := prometheus.NewRegistry() reg.MustRegister( collectors.NewGoCollector(), @@ -444,17 +471,13 @@ func (srv *Server) start() error { len(srv.shardStats), srv.numShards(), srv.lastUsage.Load().All(), bytesFmt(srv.maxSize)) if len(srv.jwtIssuers) > 0 { - issuerURLs := make([]string, 0, len(srv.jwtIssuers)) - for iss := range srv.jwtIssuers { - issuerURLs = append(issuerURLs, iss) - } - srv.jwtValidator = ijwt.NewJWTValidator(srv.logf, gocachedAudience, issuerURLs) + srv.jwtValidator = ijwt.NewJWTValidator(srv.logf, gocachedAudience, srv.jwtIssuers) if err := srv.jwtValidator.RunUpdateJWKSLoop(srv.shutdownCtx); err != nil { return fmt.Errorf("failed to fetch JWKS for JWT validator: %w", err) } - for iss, entry := range srv.jwtIssuers { - srv.logf("gocached: using JWT issuer %q with required claims %v, global write claims %v", iss, entry.requiredClaims, entry.globalWriteClaims) + for _, iss := range srv.jwtIssuers { + srv.logf("gocached: using JWT issuer %q", iss) } go srv.runCleanSessionsLoop() @@ -568,40 +591,44 @@ func WithShardPrefixLen(n int) ServerOption { } } -// JWTIssuerConfig configures a single OIDC issuer for JWT-based authentication. -type JWTIssuerConfig struct { - // Issuer is the OIDC issuer URL. It must be a reachable HTTP(S) server - // that serves its JWKS via a URL discoverable at - // /.well-known/openid-configuration. - Issuer string +// Namespace identifies a logical partition of the cache where each peer is +// equally trusted. Every session is associated with exactly one Namespace to +// which it can read and write; sessions for non-global namespaces also read +// from [GlobalNamespace]. See [WithNamespaceMapping]. It may only contain +// characters from the set [a-zA-Z0-9._~:/@+|=-]. +type Namespace string - // RequiredClaims are claims that any JWT from this issuer must have to - // start a session. All key-value pairs must match exactly. - RequiredClaims map[string]string +// GlobalNamespace is a trusted namespace that all sessions can read from. Only +// sessions explicitly mapped to GlobalNamespace can write to it. +const GlobalNamespace Namespace = "" - // GlobalWriteClaims are claims that a JWT from this issuer must have to - // write to the cache's global namespace. It should be a superset of - // RequiredClaims. - GlobalWriteClaims map[string]string +// WithJWTAuth enables JWT-based authentication for the server. Each issuer +// must be a reachable HTTP(S) server that serves its JWKS via a URL +// discoverable at /.well-known/openid-configuration. JWTs presented for token +// exchange must pass the standard signature/issuer/audience/expiry checks +// against one of these issuers. If [WithNamespaceMapping] is provided, then +// it may still be rejected if the mapping function returns an error for its +// claims. No requests other than token exchange are allowed without +// authentication. It may be called multiple times; issuers accumulate. +func WithJWTAuth(issuers ...string) ServerOption { + return func(srv *Server) { + srv.jwtIssuers = append(srv.jwtIssuers, issuers...) + } } -// WithJWTAuth enables JWT-based authentication for the server. Each issuer must -// be a reachable HTTP(S) server that serves its JWKS via a URL discoverable at -// /.well-known/openid-configuration, and any JWT presented to the server must -// exactly match the issuer's required claims to start a session. No requests are -// allowed without authentication if JWT auth is enabled. It can be called multiple -// times; configs accumulate. -func WithJWTAuth(issuers ...JWTIssuerConfig) ServerOption { +// WithNamespaceMapping sets the function that makes policy decisions based on +// a JWT's claims. It is called once per token exchange after the JWT's +// signature and standard claims have been validated. It should return an error +// if the claims are not authorized, and otherwise return which [Namespace] the +// session is allowed to read and write in. See [Namespace] for character set +// constraints. All authorized sessions are allowed to read from the +// [GlobalNamespace] regardless of the namespace returned. Check claims["iss"] +// to switch on per-issuer rules. If JWT auth is enabled but no mapping +// function is provided, all sessions will read and write in the +// [GlobalNamespace]. +func WithNamespaceMapping(fn func(claims map[string]any) (Namespace, error)) ServerOption { return func(srv *Server) { - if srv.jwtIssuers == nil { - srv.jwtIssuers = make(map[string]*jwtIssuerConfig) - } - for _, ic := range issuers { - srv.jwtIssuers[ic.Issuer] = &jwtIssuerConfig{ - requiredClaims: ic.RequiredClaims, - globalWriteClaims: ic.GlobalWriteClaims, - } - } + srv.namespaceMapping = fn } } @@ -613,12 +640,21 @@ func NewServer(opts ...ServerOption) (*Server, error) { shutdownCtx: context.Background(), logf: log.Printf, sessions: make(map[string]*sessionData), + namespaces: make(map[Namespace]int64), clock: time.Now, } for _, opt := range opts { opt(srv) } + if len(srv.jwtIssuers) > 0 && srv.namespaceMapping == nil { + // If JWT auth is enabled, but not namespace mapping, every session is in + // the global namespace. + srv.namespaceMapping = func(claims map[string]any) (Namespace, error) { + return GlobalNamespace, nil + } + } + err := srv.start() if err != nil { return nil, err @@ -668,13 +704,15 @@ type Server struct { shutdownCtx context.Context shutdownCancel context.CancelFunc - jwtValidator *ijwt.Validator // nil unless jwtIssuers is non-empty - jwtIssuers map[string]*jwtIssuerConfig // keyed by issuer URL + jwtValidator *ijwt.Validator // nil unless jwtIssuers is non-empty + jwtIssuers []string // accepted issuer URLs + namespaceMapping func(claims map[string]any) (Namespace, error) // required when jwtIssuers is non-empty mu sync.RWMutex // guards following fields in this block sessions map[string]*sessionData // maps access token -> session data. accessDirty map[actionKey]int64 // action -> accessTime accessFlushTimer *time.Timer // nil if no flush is scheduled + namespaces map[Namespace]int64 // cached namespace string -> NamespaceID // sqliteWriteMu serializes access to SQLite. In theory the SQLite driver // should serialize access with our 5000ms busy timeout, but empirically we @@ -763,18 +801,13 @@ type Server struct { } } -// jwtIssuerConfig holds per-issuer claim requirements for JWT auth. -type jwtIssuerConfig struct { - requiredClaims map[string]string - globalWriteClaims map[string]string -} - // sessionData corresponds to a specific access token, and is only used if JWT // auth is enabled. type sessionData struct { - expiry time.Time // Session valid until. - globalNSWrite bool // Whether this session can write to the cache's global namespace. - claims map[string]any // Claims from the JWT used to create this session, stored for debug. + expiry time.Time // Session valid until. + namespaceID int64 // The namespace this session writes to. 0 means GlobalNamespace; non-zero sessions also read from 0. + namespace Namespace // Namespace this session writes to, stored for debug. + claims map[string]any // Claims from the JWT used to create this session, stored for debug. mu sync.Mutex // Guards stats. stats stats @@ -884,12 +917,11 @@ func (srv *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { } if r.Method == "PUT" { - if sessionData != nil && !sessionData.globalNSWrite { - // TODO(tomhjp): support per-namespace writes. - http.Error(w, "forbidden", http.StatusForbidden) - return + var writeNS int64 + if sessionData != nil { + writeNS = sessionData.namespaceID } - srv.handlePut(w, r, reqStats) + srv.handlePut(w, r, reqStats, writeNS) return } if r.Method != "GET" && r.Method != "HEAD" { @@ -897,7 +929,7 @@ func (srv *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { return } if strings.HasPrefix(r.URL.Path, "/action/") { - srv.handleGetAction(w, r, reqStats) + srv.handleGetAction(w, r, reqStats, sessionData) return } if sessionData != nil && r.URL.Path == "/session/stats" { @@ -966,7 +998,7 @@ func getHexSuffix(r *http.Request, prefix string) (hexSuffix string, ok bool) { // actionKey is the comparable value type for the (NamespaceID, ActionID) // primary key tuple used in the SQLite Actions table. type actionKey struct { - NamespaceID int // 0 for global + NamespaceID int64 // 0 for global ActionID string } @@ -988,7 +1020,27 @@ func validHex(x string) bool { // we do a DB write to update it. const relAtimeSeconds = 60 * 60 * 24 // 1 day -func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats *stats) { +const getFromGlobalNamespace = ` +SELECT b.SHA256, b.StoredSize, b.UncompressedSize, b.SmallData, a.AltOutputID, a.AccessTime, a.NamespaceID +FROM Actions a, Blobs b +WHERE a.NamespaceID = 0 + AND a.ActionID = ? + AND a.BlobID = b.BlobID +` + +// Get hits from the global namespace first so shared cache mtime is bumped +// with higher priority than namespaced cache. +const getFromSessionNamespace = ` +SELECT b.SHA256, b.StoredSize, b.UncompressedSize, b.SmallData, a.AltOutputID, a.AccessTime, a.NamespaceID +FROM Actions a, Blobs b +WHERE a.NamespaceID IN (0, ?) + AND a.ActionID = ? + AND a.BlobID = b.BlobID +ORDER BY CASE a.NamespaceID WHEN 0 THEN 0 ELSE 1 END +LIMIT 1 +` + +func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats *stats, sessionData *sessionData) { srv.m.ActiveGets.Add(1) defer srv.m.ActiveGets.Add(-1) @@ -1014,19 +1066,24 @@ func (srv *Server) handleGetAction(w http.ResponseWriter, r *http.Request, stats return } - var sha256hex string - var storedSize, uncompressedSize int64 - var smallData sql.NullString - var altObjectID string - var accessTime int64 - var actionKey = actionKey{ - NamespaceID: 0, // global for now; TODO(bradfitz): support namespac - ActionID: actionID, + var ( + sha256hex string + storedSize, uncompressedSize int64 + smallData sql.NullString + altObjectID string + accessTime int64 + err error + actionKey = actionKey{ + ActionID: actionID, + } + ) + if sessionData != nil && sessionData.namespaceID != 0 { + err = srv.db.QueryRow(getFromSessionNamespace, sessionData.namespaceID, actionID).Scan( + &sha256hex, &storedSize, &uncompressedSize, &smallData, &altObjectID, &accessTime, &actionKey.NamespaceID) + } else { + err = srv.db.QueryRow(getFromGlobalNamespace, actionID).Scan( + &sha256hex, &storedSize, &uncompressedSize, &smallData, &altObjectID, &accessTime, &actionKey.NamespaceID) } - err := srv.db.QueryRow( - "SELECT b.SHA256, b.StoredSize, b.UncompressedSize, b.SmallData, a.AltOutputID, a.AccessTime FROM Actions a, Blobs b WHERE a.NameSpaceID = ? AND a.ActionID = ? AND a.BlobID = b.BlobID", - actionKey.NamespaceID, actionKey.ActionID).Scan( - &sha256hex, &storedSize, &uncompressedSize, &smallData, &altObjectID, &accessTime) if err != nil { if errors.Is(err, sql.ErrNoRows) { http.Error(w, "not found", http.StatusNotFound) @@ -1240,7 +1297,7 @@ func (srv *Server) getObjectFromDiskOrPeer(_ context.Context, sha256hex string, return f, nil } -func (s *Server) handlePut(w http.ResponseWriter, r *http.Request, stats *stats) { +func (s *Server) handlePut(w http.ResponseWriter, r *http.Request, stats *stats, namespaceID int64) { s.m.ActivePuts.Add(1) defer s.m.ActivePuts.Add(-1) @@ -1317,13 +1374,12 @@ func (s *Server) handlePut(w http.ResponseWriter, r *http.Request, stats *stats) // Insert or update the action in the database. nowUnix := s.now().Unix() altObjectID := "" - namespace := 0 // global for now; TODO(bradfitz): support namespaces if sha256hex != outputID { altObjectID = outputID } res, err := s.db.Exec(`INSERT OR IGNORE INTO Actions (NamespaceID, ActionID, BlobID, AltOutputID, CreateTime, AccessTime) VALUES (?, ?, ?, ?, ?, ?)`, - namespace, + namespaceID, actionID, blobID, altObjectID, @@ -1384,23 +1440,44 @@ func (srv *Server) handleTokenExchange(w http.ResponseWriter, r *http.Request) { return } - globalNSWrite, err := srv.evaluateClaims(jwtClaims) + ns, err := srv.namespaceMapping(jwtClaims) if err != nil { srv.m.AuthErrs.Add(1) if srv.verbose { - srv.logf("token exchange: %v", err) + srv.logf("token exchange: namespace func error: %v", err) } http.Error(w, "unauthorized", http.StatusUnauthorized) return } + if err := validateNamespace(ns); err != nil { + srv.m.AuthErrs.Add(1) + if srv.verbose { + srv.logf("token exchange: invalid namespace from claims: %v", err) + } + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + + var namespaceID int64 + if ns != GlobalNamespace { + namespaceID, err = srv.resolveNamespaceID(ns) + if err != nil { + srv.m.AuthErrs.Add(1) + srv.logf("token exchange: %v", err) + http.Error(w, "internal error", http.StatusInternalServerError) + return + } + } + const ttl = time.Hour // 52 base32 characters, 256 bits of entropy. accessToken := tokenPrefix + strings.ToLower(rand.Text()+rand.Text()) srv.addSessionData(accessToken, &sessionData{ - expiry: srv.now().UTC().Add(ttl), - globalNSWrite: globalNSWrite, - claims: jwtClaims, + expiry: srv.now().UTC().Add(ttl), + namespaceID: namespaceID, + namespace: ns, + claims: jwtClaims, }) resp := map[string]any{ @@ -1418,38 +1495,61 @@ func (srv *Server) handleTokenExchange(w http.ResponseWriter, r *http.Request) { srv.m.Auths.Add(1) } -func (srv *Server) evaluateClaims(claims map[string]any) (globalNSWrite bool, _ error) { - iss, _ := claims["iss"].(string) - cfg, ok := srv.jwtIssuers[iss] - if !ok { - return false, fmt.Errorf("got claims %v; unknown issuer %q", claims, iss) - } +// namespaceAllowedBytes is the set of non-alphanumeric bytes permitted in a +// namespace. It is chosen to cover characters common in JWT identity claims: +// issuer URLs, emails, and provider-structured "sub" values such as +// "repo:org/repo:environment:prod" and "auth0|abc123". It is deliberately +// ASCII-only so SQLite BINARY comparison matches Go string equality byte for +// byte, avoiding Unicode normalization divergence. +const namespaceAllowedBytes = "._~:/@+|=-" - if missing := findMissingClaims(cfg.requiredClaims, claims); len(missing) > 0 { - return false, fmt.Errorf("got claims %v; missing required claims: %v", claims, missing) +func validateNamespace(ns Namespace) error { + if ns == GlobalNamespace { + return nil + } + for i := 0; i < len(ns); i++ { + c := ns[i] + switch { + case c >= 'A' && c <= 'Z', c >= 'a' && c <= 'z', c >= '0' && c <= '9': + case strings.IndexByte(namespaceAllowedBytes, c) >= 0: + default: + return fmt.Errorf("namespace contains disallowed byte %#x at index %d", c, i) + } } + return nil +} - if missing := findMissingClaims(cfg.globalWriteClaims, claims); len(missing) == 0 { - return true, nil - } else if srv.verbose { - srv.logf("token exchange: missing global namespace write claims: %v", missing) +// resolveNamespaceID returns the integer ID for the given Namespace, +// inserting a row in the Namespaces table if one doesn't already exist. +func (srv *Server) resolveNamespaceID(ns Namespace) (int64, error) { + // If it's not a new namespace, we only need to consult our cache of IDs. + srv.mu.Lock() + id, ok := srv.namespaces[ns] + srv.mu.Unlock() + if ok { + return id, nil } - return false, nil -} + srv.sqliteWriteMu.Lock() + defer srv.sqliteWriteMu.Unlock() -func findMissingClaims(wantClaims map[string]string, gotClaims map[string]any) map[string]any { - if wantClaims == nil { - return nil + srv.mu.Lock() + defer srv.mu.Unlock() + + // Check if we lost a race now that we have both locks. + if id, ok = srv.namespaces[ns]; ok { + return id, nil } - missing := make(map[string]any) - for k, want := range wantClaims { - if got, ok := gotClaims[k]; !ok || got != want { - missing[k] = want - } + err := srv.db.QueryRow(`INSERT INTO Namespaces (Namespace) VALUES (?) + RETURNING NamespaceID;`, ns).Scan(&id) + if err != nil { + return 0, fmt.Errorf("resolving namespace %q: %w", ns, err) } - return missing + + srv.namespaces[ns] = id + + return id, nil } func (srv *Server) handleSessionStats(w http.ResponseWriter, sessionData *sessionData) { @@ -2524,10 +2624,11 @@ func (srv *Server) serveSessions(w http.ResponseWriter, r *http.Request) { for _, v := range srv.sessions { v.mu.Lock() sessions = append(sessions, &sessionData{ - expiry: v.expiry, - globalNSWrite: v.globalNSWrite, - claims: v.claims, - stats: v.stats, + expiry: v.expiry, + namespaceID: v.namespaceID, + namespace: v.namespace, + claims: v.claims, + stats: v.stats, }) v.mu.Unlock() } @@ -2535,15 +2636,13 @@ func (srv *Server) serveSessions(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/html; charset=utf-8") fmt.Fprintf(w, "

gocached sessions

\n") - for iss, cfg := range srv.jwtIssuers { + for _, iss := range srv.jwtIssuers { fmt.Fprintf(w, "

JWT issuer: %s

\n", iss) - fmt.Fprintf(w, "

JWT claims required: %v

\n", cfg.requiredClaims) - fmt.Fprintf(w, "

JWT global write claims required: %v

\n", cfg.globalWriteClaims) } fmt.Fprintf(w, "

Number of sessions: %d

\n", len(sessions)) fmt.Fprintf(w, "\n") - fmt.Fprintf(w, "\n") + fmt.Fprintf(w, "\n") slices.SortFunc(sessions, func(a, b *sessionData) int { return a.stats.LastUsed.Compare(b.stats.LastUsed) }) @@ -2552,10 +2651,14 @@ func (srv *Server) serveSessions(w http.ResponseWriter, r *http.Request) { if !d.stats.LastUsed.IsZero() { lastUsed = durFmt(time.Since(d.stats.LastUsed)) + " ago" } + nsLabel := "(global)" + if d.namespaceID != 0 { + nsLabel = fmt.Sprintf("%q (id=%d)", d.namespace, d.namespaceID) + } statsJSON, _ := json.MarshalIndent(d.stats, "", " ") claimsJSON, _ := json.MarshalIndent(d.claims, "", " ") - fmt.Fprintf(w, "\n", - lastUsed, d.expiry.Format(time.RFC3339), d.globalNSWrite, statsJSON, claimsJSON) + fmt.Fprintf(w, "\n", + lastUsed, d.expiry.Format(time.RFC3339), nsLabel, statsJSON, claimsJSON) } fmt.Fprintf(w, "
Last usedExpiry timeGlobal writeStatsClaims
Last usedExpiry timeNamespaceStatsClaims
%s%s%v
%s
%s
%s%s%s
%s
%s
\n") } diff --git a/gocache/gocached/gocached_test.go b/gocache/gocached/gocached_test.go index bcb0b89..bd00d4b 100644 --- a/gocache/gocached/gocached_test.go +++ b/gocache/gocached/gocached_test.go @@ -38,6 +38,23 @@ import ( // value in SQLite to store bytes, as it's common. const sha256OfEmpty = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855" +// Package-level ECDSA P-256 keys reused across tests that need OIDC signing +// keys. Generating these is the single most expensive step in the JWT tests, +// so we share them rather than regenerating per test. +var ( + testKey1 = mustGenerateTestKey() + testKey2 = mustGenerateTestKey() + testKey3 = mustGenerateTestKey() +) + +func mustGenerateTestKey() *ecdsa.PrivateKey { + k, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + panic(fmt.Sprintf("generating test ECDSA key: %v", err)) + } + return k +} + type tester struct { t testing.TB srv *Server @@ -1150,41 +1167,17 @@ func TestClientConnReuse(t *testing.T) { } func TestExchangeToken(t *testing.T) { - // Generate private keys outside of the loop for speed. - privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatalf("error generating OIDC server private key: %v", err) - } - otherPrivateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatalf("error generating OIDC server private key: %v", err) - } - wantClaims := map[string]string{ - "sub": "user123", - } - wantGlobalClaims := map[string]string{ - "sub": "user123", - "ref": "refs/heads/main", - } + privateKey := testKey1 + otherPrivateKey := testKey2 for name, tc := range map[string]struct { mutateClaims func(jwt.MapClaims) signingKey *ecdsa.PrivateKey wantStatusCode int - wantWrite bool }{ // Base case: no mutation. - "valid_read": { + "valid": { wantStatusCode: http.StatusOK, - wantWrite: false, - }, - // Additional claim needed for write scope. - "valid_write": { - mutateClaims: func(cl jwt.MapClaims) { - cl["ref"] = "refs/heads/main" - }, - wantStatusCode: http.StatusOK, - wantWrite: true, }, // Every other test makes one mutation from the base case that should cause failure. "missing_sub": { @@ -1231,10 +1224,12 @@ func TestExchangeToken(t *testing.T) { t.Run(name, func(t *testing.T) { issuer, createJWT := startOIDCServer(t, privateKey.Public()) st := newServerTester(t, - WithJWTAuth(JWTIssuerConfig{ - Issuer: issuer, - RequiredClaims: wantClaims, - GlobalWriteClaims: wantGlobalClaims, + WithJWTAuth(issuer), + WithNamespaceMapping(func(claims map[string]any) (Namespace, error) { + if claims["sub"] != "user123" { + return "", fmt.Errorf("sub = %v, want user123", claims["sub"]) + } + return GlobalNamespace, nil }), ) @@ -1307,14 +1302,8 @@ func TestExchangeToken(t *testing.T) { cl.AccessToken = d.AccessToken st.wantGetMiss(cl, "abc123") - if tc.wantWrite { - st.wantPut(cl, "abc123", "def456", "data789") - st.wantGet(cl, "abc123", "def456", "data789") - } else { - if _, err := cl.Put(t.Context(), "abc123", "def456", 0, nil); err == nil { - t.Fatalf("Put without write scope succeeded unexpectedly") - } - } + st.wantPut(cl, "abc123", "def456", "data789") + st.wantGet(cl, "abc123", "def456", "data789") // Check session stats. reqStats, err := http.NewRequest("GET", st.hs.URL+"/session/stats", nil) @@ -1342,102 +1331,91 @@ func TestExchangeToken(t *testing.T) { if stats.Gets == 0 { t.Errorf("expected non-zero gets in session stats") } - if stats.Puts == 0 && tc.wantWrite { + if stats.Puts == 0 { t.Errorf("expected non-zero puts in session stats") } }) } } -func TestMultiIssuerAuth(t *testing.T) { - // Generate separate keys for each issuer. - keyA, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatalf("error generating key A: %v", err) - } - keyB, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatalf("error generating key B: %v", err) - } - keyC, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) - if err != nil { - t.Fatalf("error generating key C: %v", err) +func TestExchangeTokenNamespaceValidation(t *testing.T) { + for name, tc := range map[string]struct { + namespace string + wantStatusCode int + }{ + "simple": {"user123", http.StatusOK}, + "structured_sub": {"repo:octo-org/octo-repo:environment:prod", http.StatusOK}, + "auth0_sub": {"auth0|507f1f77bcf86cd799439020", http.StatusOK}, + "email": {"alice+ci@example.com", http.StatusOK}, + "space": {"has space", http.StatusUnauthorized}, + "non_ascii": {"héllo", http.StatusUnauthorized}, + "control_char": {"bad\tns", http.StatusUnauthorized}, + "html_meta": {"