From ac9ae5a748d6f09d772f0f4d0c9256ba87af77bf Mon Sep 17 00:00:00 2001 From: lyon Date: Mon, 22 Jun 2026 13:07:18 +0800 Subject: [PATCH] feat: add workbench redis cache abstraction --- go.mod | 7 +- go.sum | 8 + internal/workbenchruntime/cache.go | 591 ++++++++++++++++++++++++ internal/workbenchruntime/cache_test.go | 102 ++++ internal/workbenchruntime/service.go | 24 +- 5 files changed, 727 insertions(+), 5 deletions(-) create mode 100644 internal/workbenchruntime/cache.go create mode 100644 internal/workbenchruntime/cache_test.go diff --git a/go.mod b/go.mod index ae410d40..0ac27844 100644 --- a/go.mod +++ b/go.mod @@ -2,9 +2,14 @@ module github.com/pikasTech/HWLAB go 1.18 -require github.com/jackc/pgx/v4 v4.18.3 +require ( + github.com/jackc/pgx/v4 v4.18.3 + github.com/redis/go-redis/v9 v9.7.0 +) require ( + github.com/cespare/xxhash/v2 v2.2.0 // indirect + github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect github.com/jackc/chunkreader/v2 v2.0.1 // indirect github.com/jackc/pgconn v1.14.3 // indirect github.com/jackc/pgio v1.0.0 // indirect diff --git a/go.sum b/go.sum index 359d9b85..9a5d04d7 100644 --- a/go.sum +++ b/go.sum @@ -1,6 +1,10 @@ github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= github.com/Masterminds/semver/v3 v3.1.1 h1:hLg3sBzpNErnxhQtUy/mmLR2I9foDujNK030IGemrRc= github.com/Masterminds/semver/v3 v3.1.1/go.mod h1:VPu/7SZ7ePZ3QOrcuXROw5FAcLl4a0cBrbBpGY/8hQs= +github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs= +github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA= +github.com/cespare/xxhash/v2 v2.2.0 h1:DC2CZ1Ep5Y4k3ZQ899DldepgrayRUGE6BBZ/cd9Cj44= +github.com/cespare/xxhash/v2 v2.2.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/cockroachdb/apd v1.1.0 h1:3LFP3629v+1aKXU5Q37mxmRxX/pIu1nijXydLShEq5I= github.com/cockroachdb/apd v1.1.0/go.mod h1:8Sl8LxpKi29FqWXR16WEFZRNSz3SoPzUzeMeY4+DwBQ= github.com/coreos/go-systemd v0.0.0-20190321100706-95778dfbb74e/go.mod h1:F5haX7vjVVG0kc13fIWeqUViNPyEJxv/OmvnBo0Yme4= @@ -9,6 +13,8 @@ github.com/creack/pty v1.1.7/go.mod h1:lj5s0c3V2DBrqTV7llrYr5NG6My20zk30Fl46Y7Do github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= 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/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78= +github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc= github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY= github.com/go-logfmt/logfmt v0.5.0/go.mod h1:wCYkCAKZfumFQihp8CzCvQ3paCTfi41vtzG1KdI/P7A= github.com/go-stack/stack v1.8.0/go.mod h1:v0f6uXyyMGvRgIKkXu+yp6POWl0qKG85gN/melR3HDY= @@ -83,6 +89,8 @@ github.com/pkg/errors v0.8.1 h1:iURUrRGxPUNPdy5/HRSm+Yj6okJ6UtLINN0Q9M4+h3I= github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= 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/redis/go-redis/v9 v9.7.0 h1:HhLSs+B6O021gwzl+locl0zEDnyNkxMtf/Z3NNBMa9E= +github.com/redis/go-redis/v9 v9.7.0/go.mod h1:f6zhXITC7JUJIlPEiBOTXxJgPLdZcA93GewI7inzyWw= github.com/rogpeppe/go-internal v1.3.0/go.mod h1:M8bDsm7K2OlrFYOpmOWEs/qY81heoFRclV5y23lUDJ4= github.com/rs/xid v1.2.1/go.mod h1:+uKXf+4Djp6Md1KODXJxgGQPKngRmWyn10oCKFzNHOQ= github.com/rs/zerolog v1.13.0/go.mod h1:YbFCdg8HfsridGWAh22vktObvhZbQsZXe4/zB0OKkWU= diff --git a/internal/workbenchruntime/cache.go b/internal/workbenchruntime/cache.go new file mode 100644 index 00000000..24664fba --- /dev/null +++ b/internal/workbenchruntime/cache.go @@ -0,0 +1,591 @@ +// Redis derived cache boundary follows PJ2026-01060505, PJ2026-0104010803, +// PJ2026-010403 and implementation version draft-2026-06-22-p1-workbench-redis-derived-cache. +package workbenchruntime + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "net" + "os" + "strconv" + "strings" + "sync" + "time" + + "github.com/redis/go-redis/v9" +) + +const cacheKeySchemaVersion = "wb-derived-v1" + +var errCacheDisabled = errors.New("workbench redis cache disabled") + +type cacheConfig struct { + Enabled bool + RedisURL string + ConnectTimeout time.Duration + OperationTimeout time.Duration + MaxKeyBytes int + MaxPayloadBytes int + SessionsSummaryTTL time.Duration + TerminalTurnTTL time.Duration + TerminalTracePageTTL time.Duration + RunningPageTTL time.Duration + UnavailableLivenessFailure bool + UnavailableReadMode string + UnavailableFallbackStateSource bool +} + +type cacheKeyParts struct { + Class string + ActorID string + SessionID string + TraceID string + Cursor string + ProjectionRevision string + ProjectionSeq int64 + Authority string +} + +type derivedCache struct { + config cacheConfig + client *redis.Client + configured bool + initErr error + metrics cacheMetrics + createdTime time.Time +} + +type cacheMetrics struct { + mu sync.Mutex + Hits uint64 + Misses uint64 + Sets uint64 + Deletes uint64 + Errors uint64 + LastErrorKind string + LastErrorOp string + LastErrorAt time.Time +} + +type cacheError struct { + Op string + Kind string + Err error +} + +func (e cacheError) Error() string { + if e.Err == nil { + return e.Kind + } + return e.Kind + ": " + e.Err.Error() +} + +func (e cacheError) Unwrap() error { return e.Err } + +func newDerivedCache(config cacheConfig) *derivedCache { + cache := &derivedCache{config: config, configured: strings.TrimSpace(config.RedisURL) != "", createdTime: time.Now().UTC()} + if !config.Enabled { + return cache + } + if strings.TrimSpace(config.RedisURL) == "" { + cache.initErr = cacheError{Op: "init", Kind: "cache_config_missing_url", Err: errors.New("redis url is empty")} + return cache + } + options, err := redis.ParseURL(config.RedisURL) + if err != nil { + cache.initErr = cacheError{Op: "init", Kind: "cache_config_invalid_url", Err: err} + return cache + } + if config.ConnectTimeout > 0 { + options.DialTimeout = config.ConnectTimeout + } + if config.OperationTimeout > 0 { + options.ReadTimeout = config.OperationTimeout + options.WriteTimeout = config.OperationTimeout + } + cache.client = redis.NewClient(options) + return cache +} + +func cacheConfigFromEnv() cacheConfig { + return cacheConfig{ + Enabled: envBool("HWLAB_WORKBENCH_CACHE_REDIS_ENABLED", false), + RedisURL: first(os.Getenv("HWLAB_WORKBENCH_CACHE_REDIS_URL")), + ConnectTimeout: time.Duration(envInt("HWLAB_WORKBENCH_CACHE_REDIS_CONNECT_TIMEOUT_MS", 200)) * time.Millisecond, + OperationTimeout: time.Duration(envInt("HWLAB_WORKBENCH_CACHE_REDIS_OPERATION_TIMEOUT_MS", 100)) * time.Millisecond, + MaxKeyBytes: envInt("HWLAB_WORKBENCH_CACHE_REDIS_MAX_KEY_BYTES", 512), + MaxPayloadBytes: envInt("HWLAB_WORKBENCH_CACHE_REDIS_MAX_PAYLOAD_BYTES", 262144), + SessionsSummaryTTL: time.Duration(envInt("HWLAB_WORKBENCH_CACHE_REDIS_TTL_SESSIONS_SUMMARY_MS", 30000)) * time.Millisecond, + TerminalTurnTTL: time.Duration(envInt("HWLAB_WORKBENCH_CACHE_REDIS_TTL_TERMINAL_TURN_MS", 300000)) * time.Millisecond, + TerminalTracePageTTL: time.Duration(envInt("HWLAB_WORKBENCH_CACHE_REDIS_TTL_TERMINAL_TRACE_PAGE_MS", 300000)) * time.Millisecond, + RunningPageTTL: time.Duration(envInt("HWLAB_WORKBENCH_CACHE_REDIS_TTL_RUNNING_PAGE_MS", 5000)) * time.Millisecond, + UnavailableLivenessFailure: envBool("HWLAB_WORKBENCH_CACHE_REDIS_UNAVAILABLE_LIVENESS_FAILURE", false), + UnavailableReadMode: first(os.Getenv("HWLAB_WORKBENCH_CACHE_REDIS_UNAVAILABLE_READ_MODE"), "db-or-degraded"), + UnavailableFallbackStateSource: envBool("HWLAB_WORKBENCH_CACHE_REDIS_UNAVAILABLE_FALLBACK_STATE_SOURCE", false), + } +} + +func (c *derivedCache) Close() error { + if c == nil || c.client == nil { + return nil + } + return c.client.Close() +} + +func (c *derivedCache) BuildKey(parts cacheKeyParts) (string, error) { + if c == nil { + return "", errCacheDisabled + } + return buildCacheKey(c.config.MaxKeyBytes, parts) +} + +func (c *derivedCache) Get(ctx context.Context, key string, target any) (bool, error) { + if !c.usable() { + return false, nil + } + if err := validateCacheKey(c.config.MaxKeyBytes, key); err != nil { + wrapped := cacheError{Op: "get", Kind: classifyCacheError(err), Err: err} + c.metrics.recordError(wrapped) + return false, wrapped + } + var payload []byte + err := c.withTimeout(ctx, func(opCtx context.Context) error { + var err error + payload, err = c.client.Get(opCtx, key).Bytes() + return err + }) + if errors.Is(err, redis.Nil) { + c.metrics.recordMiss() + return false, nil + } + if err != nil { + wrapped := cacheError{Op: "get", Kind: classifyCacheError(err), Err: err} + c.metrics.recordError(wrapped) + return false, wrapped + } + if len(payload) > c.config.MaxPayloadBytes { + wrapped := cacheError{Op: "get", Kind: "cache_payload_too_large", Err: fmt.Errorf("payload bytes %d exceeds %d", len(payload), c.config.MaxPayloadBytes)} + c.metrics.recordError(wrapped) + return false, wrapped + } + if target != nil { + if err := json.Unmarshal(payload, target); err != nil { + wrapped := cacheError{Op: "get", Kind: "cache_payload_decode_failed", Err: err} + c.metrics.recordError(wrapped) + return false, wrapped + } + } + c.metrics.recordHit() + return true, nil +} + +func (c *derivedCache) Set(ctx context.Context, key string, value any, ttl time.Duration) error { + if !c.usable() { + return nil + } + if err := validateCacheKey(c.config.MaxKeyBytes, key); err != nil { + wrapped := cacheError{Op: "set", Kind: classifyCacheError(err), Err: err} + c.metrics.recordError(wrapped) + return wrapped + } + payload, err := cachePayloadBytes(value, c.config.MaxPayloadBytes) + if err != nil { + wrapped := cacheError{Op: "set", Kind: classifyCacheError(err), Err: err} + c.metrics.recordError(wrapped) + return wrapped + } + err = c.withTimeout(ctx, func(opCtx context.Context) error { + return c.client.Set(opCtx, key, payload, ttl).Err() + }) + if err != nil { + wrapped := cacheError{Op: "set", Kind: classifyCacheError(err), Err: err} + c.metrics.recordError(wrapped) + return wrapped + } + c.metrics.recordSet() + return nil +} + +func (c *derivedCache) Delete(ctx context.Context, keys ...string) error { + if !c.usable() || len(keys) == 0 { + return nil + } + for _, key := range keys { + if err := validateCacheKey(c.config.MaxKeyBytes, key); err != nil { + wrapped := cacheError{Op: "delete", Kind: classifyCacheError(err), Err: err} + c.metrics.recordError(wrapped) + return wrapped + } + } + err := c.withTimeout(ctx, func(opCtx context.Context) error { + return c.client.Del(opCtx, keys...).Err() + }) + if err != nil { + wrapped := cacheError{Op: "delete", Kind: classifyCacheError(err), Err: err} + c.metrics.recordError(wrapped) + return wrapped + } + c.metrics.recordDelete(uint64(len(keys))) + return nil +} + +func (c *derivedCache) TTL(ctx context.Context, key string) (time.Duration, error) { + if !c.usable() { + return 0, nil + } + if err := validateCacheKey(c.config.MaxKeyBytes, key); err != nil { + wrapped := cacheError{Op: "ttl", Kind: classifyCacheError(err), Err: err} + c.metrics.recordError(wrapped) + return 0, wrapped + } + var ttl time.Duration + err := c.withTimeout(ctx, func(opCtx context.Context) error { + var err error + ttl, err = c.client.TTL(opCtx, key).Result() + return err + }) + if err != nil { + wrapped := cacheError{Op: "ttl", Kind: classifyCacheError(err), Err: err} + c.metrics.recordError(wrapped) + return 0, wrapped + } + return ttl, nil +} + +func (c *derivedCache) Summary() map[string]any { + if c == nil { + return map[string]any{"enabled": false, "configured": false, "role": "derived-read-cache", "valuesRedacted": true} + } + result := map[string]any{ + "enabled": c.config.Enabled, + "configured": c.config.RedisURL != "", + "role": "derived-read-cache", + "adapter": "redis", + "schemaVersion": cacheKeySchemaVersion, + "maxKeyBytes": c.config.MaxKeyBytes, + "maxPayloadBytes": c.config.MaxPayloadBytes, + "unavailableReadMode": c.config.UnavailableReadMode, + "unavailableFallbackStateSource": c.config.UnavailableFallbackStateSource, + "secretMaterialStored": false, + "valuesRedacted": true, + "metrics": c.metrics.snapshot(), + } + if c.initErr != nil { + result["degraded"] = true + result["errorKind"] = classifyCacheError(c.initErr) + } + return result +} + +func (c *derivedCache) Diagnostic(ctx context.Context) map[string]any { + result := c.Summary() + if c == nil || !c.config.Enabled { + result["available"] = false + return result + } + if c.initErr != nil || c.client == nil { + result["available"] = false + result["degraded"] = true + return result + } + err := c.withTimeout(ctx, func(opCtx context.Context) error { + return c.client.Ping(opCtx).Err() + }) + if err != nil { + wrapped := cacheError{Op: "ping", Kind: classifyCacheError(err), Err: err} + c.metrics.recordError(wrapped) + result["available"] = false + result["degraded"] = true + result["errorKind"] = wrapped.Kind + result["metrics"] = c.metrics.snapshot() + return result + } + result["available"] = true + result["degraded"] = false + return result +} + +func (c *derivedCache) usable() bool { + return c != nil && c.config.Enabled && c.client != nil && c.initErr == nil +} + +func (c *derivedCache) withTimeout(ctx context.Context, fn func(context.Context) error) error { + if ctx == nil { + ctx = context.Background() + } + if c == nil || c.config.OperationTimeout <= 0 { + return fn(ctx) + } + opCtx, cancel := context.WithTimeout(ctx, c.config.OperationTimeout) + defer cancel() + return fn(opCtx) +} + +func buildCacheKey(maxBytes int, parts cacheKeyParts) (string, error) { + if strings.TrimSpace(parts.Class) == "" { + return "", cacheError{Op: "key", Kind: "cache_key_invalid", Err: errors.New("class is required")} + } + if strings.TrimSpace(parts.ActorID) == "" { + return "", cacheError{Op: "key", Kind: "cache_key_invalid", Err: errors.New("actor id is required")} + } + if strings.TrimSpace(parts.Authority) == "" && strings.TrimSpace(parts.ProjectionRevision) == "" && parts.ProjectionSeq <= 0 { + return "", cacheError{Op: "key", Kind: "cache_key_invalid", Err: errors.New("projection authority is required")} + } + components := []string{ + "hwlab", + "workbench", + cacheKeySchemaVersion, + "class=" + safeCacheToken(parts.Class), + "actor=" + cacheHash(parts.ActorID), + "session=" + optionalCacheHash(parts.SessionID), + "trace=" + optionalCacheHash(parts.TraceID), + "cursor=" + optionalCacheHash(parts.Cursor), + } + if parts.ProjectionSeq > 0 { + components = append(components, "seq="+strconv.FormatInt(parts.ProjectionSeq, 10)) + } + if strings.TrimSpace(parts.ProjectionRevision) != "" { + components = append(components, "rev="+cacheHash(parts.ProjectionRevision)) + } + if strings.TrimSpace(parts.Authority) != "" { + components = append(components, "auth="+cacheHash(parts.Authority)) + } + key := strings.Join(components, ":") + if maxBytes <= 0 { + maxBytes = 512 + } + if len([]byte(key)) <= maxBytes { + return key, nil + } + compact := strings.Join(components[:4], ":") + ":hash=" + cacheHash(key) + if len([]byte(compact)) > maxBytes { + return "", cacheError{Op: "key", Kind: "cache_key_too_large", Err: fmt.Errorf("key bytes exceed %d", maxBytes)} + } + return compact, nil +} + +func validateCacheKey(maxBytes int, key string) error { + if strings.TrimSpace(key) == "" { + return cacheError{Op: "key", Kind: "cache_key_invalid", Err: errors.New("cache key is empty")} + } + if maxBytes <= 0 { + maxBytes = 512 + } + if len([]byte(key)) > maxBytes { + return cacheError{Op: "key", Kind: "cache_key_too_large", Err: fmt.Errorf("key bytes exceed %d", maxBytes)} + } + return nil +} + +func cachePayloadBytes(value any, maxBytes int) ([]byte, error) { + payload, err := json.Marshal(value) + if err != nil { + return nil, cacheError{Op: "payload", Kind: "cache_payload_encode_failed", Err: err} + } + var decoded any + if err := json.Unmarshal(payload, &decoded); err != nil { + return nil, cacheError{Op: "payload", Kind: "cache_payload_decode_failed", Err: err} + } + if key, ok := cachePayloadSensitiveKey(decoded); ok { + return nil, cacheError{Op: "payload", Kind: "cache_payload_rejected", Err: fmt.Errorf("sensitive payload key %q is not cacheable", key)} + } + if maxBytes <= 0 { + maxBytes = 262144 + } + if len(payload) > maxBytes { + return nil, cacheError{Op: "payload", Kind: "cache_payload_too_large", Err: fmt.Errorf("payload bytes %d exceeds %d", len(payload), maxBytes)} + } + return payload, nil +} + +func cachePayloadSensitiveKey(value any) (string, bool) { + switch item := value.(type) { + case map[string]any: + for key, child := range item { + if cachePayloadKeyRejected(key) { + return key, true + } + if nested, ok := cachePayloadSensitiveKey(child); ok { + return nested, true + } + } + case []any: + for _, child := range item { + if nested, ok := cachePayloadSensitiveKey(child); ok { + return nested, true + } + } + case []map[string]any: + for _, child := range item { + if nested, ok := cachePayloadSensitiveKey(child); ok { + return nested, true + } + } + } + return "", false +} + +func cachePayloadKeyRejected(key string) bool { + normalized := normalizedCachePayloadKey(key) + for _, fragment := range []string{"authorization", "cookie", "secret", "password", "token", "apikey", "databaseurl", "dbdsn", "dsn", "prompt", "providerpayload", "stdout", "stderr"} { + if strings.Contains(normalized, fragment) { + return true + } + } + return false +} + +func normalizedCachePayloadKey(key string) string { + lower := strings.ToLower(strings.TrimSpace(key)) + var builder strings.Builder + for _, ch := range lower { + if ch >= 'a' && ch <= 'z' || ch >= '0' && ch <= '9' { + builder.WriteRune(ch) + } + } + return builder.String() +} + +func classifyCacheError(err error) string { + if err == nil { + return "" + } + var cacheErr cacheError + if errors.As(err, &cacheErr) && cacheErr.Kind != "" { + return cacheErr.Kind + } + if errors.Is(err, redis.Nil) { + return "cache_miss" + } + if errors.Is(err, context.Canceled) { + return "cache_context_canceled" + } + if errors.Is(err, context.DeadlineExceeded) { + return "cache_timeout" + } + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + return "cache_timeout" + } + message := strings.ToLower(err.Error()) + switch { + case strings.Contains(message, "connection refused") || strings.Contains(message, "no such host") || strings.Contains(message, "network is unreachable") || strings.Contains(message, "i/o timeout"): + return "cache_unavailable" + case strings.Contains(message, "payload") && strings.Contains(message, "large"): + return "cache_payload_too_large" + case strings.Contains(message, "sensitive payload"): + return "cache_payload_rejected" + case strings.Contains(message, "key") && strings.Contains(message, "exceed"): + return "cache_key_too_large" + case strings.Contains(message, "key"): + return "cache_key_invalid" + default: + return "cache_error" + } +} + +func safeCacheToken(value string) string { + text := strings.ToLower(strings.TrimSpace(value)) + var builder strings.Builder + for _, ch := range text { + switch { + case ch >= 'a' && ch <= 'z' || ch >= '0' && ch <= '9': + builder.WriteRune(ch) + case ch == '-' || ch == '_' || ch == '.': + builder.WriteRune(ch) + default: + builder.WriteRune('-') + } + } + result := strings.Trim(builder.String(), "-._") + if result == "" { + return "unknown" + } + return result +} + +func optionalCacheHash(value string) string { + if strings.TrimSpace(value) == "" { + return "none" + } + return cacheHash(value) +} + +func cacheHash(value string) string { + digest := sha256.Sum256([]byte(strings.TrimSpace(value))) + return hex.EncodeToString(digest[:])[:24] +} + +func (m *cacheMetrics) recordHit() { + m.mu.Lock() + defer m.mu.Unlock() + m.Hits++ +} + +func (m *cacheMetrics) recordMiss() { + m.mu.Lock() + defer m.mu.Unlock() + m.Misses++ +} + +func (m *cacheMetrics) recordSet() { + m.mu.Lock() + defer m.mu.Unlock() + m.Sets++ +} + +func (m *cacheMetrics) recordDelete(count uint64) { + m.mu.Lock() + defer m.mu.Unlock() + m.Deletes += count +} + +func (m *cacheMetrics) recordError(err error) { + m.mu.Lock() + defer m.mu.Unlock() + m.Errors++ + m.LastErrorKind = classifyCacheError(err) + var cacheErr cacheError + if errors.As(err, &cacheErr) { + m.LastErrorOp = cacheErr.Op + } + m.LastErrorAt = time.Now().UTC() +} + +func (m *cacheMetrics) snapshot() map[string]any { + m.mu.Lock() + defer m.mu.Unlock() + result := map[string]any{ + "hits": int64(m.Hits), + "misses": int64(m.Misses), + "sets": int64(m.Sets), + "deletes": int64(m.Deletes), + "errors": int64(m.Errors), + "valuesRedacted": true, + } + if m.LastErrorKind != "" { + result["lastErrorKind"] = m.LastErrorKind + result["lastErrorOp"] = m.LastErrorOp + result["lastErrorAt"] = m.LastErrorAt.Format(time.RFC3339Nano) + } + return result +} + +func envBool(name string, fallback bool) bool { + value := strings.ToLower(strings.TrimSpace(os.Getenv(name))) + if value == "" { + return fallback + } + switch value { + case "1", "true", "yes", "y", "on", "enabled": + return true + case "0", "false", "no", "n", "off", "disabled": + return false + default: + return fallback + } +} diff --git a/internal/workbenchruntime/cache_test.go b/internal/workbenchruntime/cache_test.go new file mode 100644 index 00000000..0770409f --- /dev/null +++ b/internal/workbenchruntime/cache_test.go @@ -0,0 +1,102 @@ +package workbenchruntime + +import ( + "context" + "errors" + "strings" + "testing" + + "github.com/redis/go-redis/v9" +) + +func TestBuildCacheKeyActorIsolationAndAuthority(t *testing.T) { + base := cacheKeyParts{Class: "sessions.summary", ActorID: "actor-a", SessionID: "ses-1", TraceID: "trace-1", Cursor: "idx:0", ProjectionSeq: 42} + keyA, err := buildCacheKey(512, base) + if err != nil { + t.Fatalf("build key actor a: %v", err) + } + base.ActorID = "actor-b" + keyB, err := buildCacheKey(512, base) + if err != nil { + t.Fatalf("build key actor b: %v", err) + } + if keyA == keyB { + t.Fatalf("expected actor-isolated keys, got identical key %q", keyA) + } + if !strings.Contains(keyA, cacheKeySchemaVersion) || !strings.Contains(keyA, "class=sessions.summary") { + t.Fatalf("key missing schema/class: %s", keyA) + } + if strings.Contains(keyA, "actor-a") || strings.Contains(keyB, "actor-b") { + t.Fatalf("cache key should hash raw actor ids: %s / %s", keyA, keyB) + } + + _, err = buildCacheKey(512, cacheKeyParts{Class: "sessions.summary", ActorID: "actor-a"}) + if cacheErrKind(err) != "cache_key_invalid" { + t.Fatalf("expected missing authority to be rejected, got %v", err) + } + _, err = buildCacheKey(512, cacheKeyParts{Class: "sessions.summary", ProjectionSeq: 1}) + if cacheErrKind(err) != "cache_key_invalid" { + t.Fatalf("expected missing actor to be rejected, got %v", err) + } +} + +func TestCachePayloadBoundary(t *testing.T) { + _, err := cachePayloadBytes(map[string]any{"sessionId": "ses-1", "summary": map[string]any{"messageCount": 2}}, 1024) + if err != nil { + t.Fatalf("safe payload rejected: %v", err) + } + + for _, payload := range []map[string]any{ + {"authorization": "bearer secret"}, + {"session": map[string]any{"prompt": "full prompt text"}}, + {"providerPayload": map[string]any{"raw": true}}, + {"stdout": strings.Repeat("x", 16)}, + } { + _, err := cachePayloadBytes(payload, 1024) + if cacheErrKind(err) != "cache_payload_rejected" { + t.Fatalf("expected payload rejection for %#v, got %v", payload, err) + } + } + + _, err = cachePayloadBytes(map[string]any{"sessionId": strings.Repeat("x", 64)}, 16) + if cacheErrKind(err) != "cache_payload_too_large" { + t.Fatalf("expected payload too large, got %v", err) + } +} + +func TestCacheErrorClassification(t *testing.T) { + cases := []struct { + err error + kind string + }{ + {redis.Nil, "cache_miss"}, + {context.DeadlineExceeded, "cache_timeout"}, + {context.Canceled, "cache_context_canceled"}, + {errors.New("dial tcp 127.0.0.1:6379: connect: connection refused"), "cache_unavailable"}, + } + for _, tc := range cases { + if got := classifyCacheError(tc.err); got != tc.kind { + t.Fatalf("classifyCacheError(%v)=%s, want %s", tc.err, got, tc.kind) + } + } +} + +func TestDisabledCacheDoesNotFailReadPath(t *testing.T) { + cache := newDerivedCache(cacheConfig{Enabled: false}) + hit, err := cache.Get(context.Background(), "ignored", &map[string]any{}) + if err != nil || hit { + t.Fatalf("disabled cache should be transparent, hit=%v err=%v", hit, err) + } + diag := cache.Diagnostic(context.Background()) + if diag["enabled"] != false || diag["available"] != false { + t.Fatalf("unexpected disabled diagnostic: %#v", diag) + } +} + +func cacheErrKind(err error) string { + var cacheErr cacheError + if errors.As(err, &cacheErr) { + return cacheErr.Kind + } + return classifyCacheError(err) +} diff --git a/internal/workbenchruntime/service.go b/internal/workbenchruntime/service.go index 34112738..30cb49e9 100644 --- a/internal/workbenchruntime/service.go +++ b/internal/workbenchruntime/service.go @@ -32,11 +32,13 @@ type Config struct { QueryTimeout time.Duration ReadinessDelay time.Duration PoolMax int + Cache cacheConfig } type Server struct { config Config db *sql.DB + cache *derivedCache mux *http.ServeMux } @@ -140,6 +142,7 @@ func Run(ctx context.Context) error { return err } server := NewServer(config, db) + defer server.Close() httpServer := &http.Server{Addr: config.Addr, Handler: server.otelHTTPHandler(server.mux), ReadHeaderTimeout: 5 * time.Second} errCh := make(chan error, 1) go func() { @@ -173,11 +176,18 @@ func ensureWorkbenchRuntimeReadIndexes(ctx context.Context, db *sql.DB) error { } func NewServer(config Config, db *sql.DB) *Server { - server := &Server{config: config, db: db, mux: http.NewServeMux()} + server := &Server{config: config, db: db, cache: newDerivedCache(config.Cache), mux: http.NewServeMux()} server.routes() return server } +func (s *Server) Close() error { + if s == nil || s.cache == nil { + return nil + } + return s.cache.Close() +} + func (s *Server) routes() { s.mux.HandleFunc("/health/live", s.handleLive) s.mux.HandleFunc("/health/ready", s.handleReady) @@ -188,7 +198,12 @@ func (s *Server) routes() { } func (s *Server) handleLive(w http.ResponseWriter, r *http.Request) { - writeJSON(w, http.StatusOK, map[string]any{"ok": true, "serviceId": serviceID, "status": "live", "valuesRedacted": true}) + cache := s.cache.Diagnostic(r.Context()) + if s.config.Cache.UnavailableLivenessFailure && cache["enabled"] == true && cache["available"] != true { + writeJSON(w, http.StatusServiceUnavailable, map[string]any{"ok": false, "serviceId": serviceID, "status": "cache-unavailable", "cache": cache, "valuesRedacted": true}) + return + } + writeJSON(w, http.StatusOK, map[string]any{"ok": true, "serviceId": serviceID, "status": "live", "cache": cache, "valuesRedacted": true}) } func (s *Server) handleReady(w http.ResponseWriter, r *http.Request) { @@ -198,7 +213,7 @@ func (s *Server) handleReady(w http.ResponseWriter, r *http.Request) { writeAPIError(w, http.StatusServiceUnavailable, "workbench_runtime_db_unavailable", "workbench runtime database is unavailable") return } - writeJSON(w, http.StatusOK, map[string]any{"ok": true, "serviceId": serviceID, "status": "ready", "valuesRedacted": true}) + writeJSON(w, http.StatusOK, map[string]any{"ok": true, "serviceId": serviceID, "status": "ready", "cache": s.cache.Diagnostic(r.Context()), "valuesRedacted": true}) } func (s *Server) handleFacts(w http.ResponseWriter, r *http.Request) { @@ -1422,11 +1437,12 @@ func configFromEnv() Config { QueryTimeout: time.Duration(envInt("HWLAB_WORKBENCH_RUNTIME_QUERY_TIMEOUT_MS", 10000)) * time.Millisecond, ReadinessDelay: time.Duration(envInt("HWLAB_WORKBENCH_RUNTIME_READY_TIMEOUT_MS", 2000)) * time.Millisecond, PoolMax: envInt("HWLAB_WORKBENCH_RUNTIME_DB_POOL_MAX", 8), + Cache: cacheConfigFromEnv(), } } func (s *Server) persistenceSummary() map[string]any { - return map[string]any{"ready": true, "durable": true, "adapter": "postgres", "serviceId": serviceID, "valuesRedacted": true} + return map[string]any{"ready": true, "durable": true, "adapter": "postgres", "serviceId": serviceID, "cache": s.cache.Summary(), "valuesRedacted": true} } func decodeJSON(w http.ResponseWriter, r *http.Request, target any) bool {