[优化] 引如缓存库ristretto

This commit is contained in:
ryan
2026-03-14 23:54:31 +08:00
parent 37deb84986
commit 3d52ddc933
8 changed files with 388 additions and 36 deletions
+1
View File
@@ -144,6 +144,7 @@ func HeartbeatNode(node *model.Node, payload AgentNodePayload) (*HeartbeatRespon
return nil, err
}
}
refreshAgentTokenCache(node)
persistHeartbeatObservability(node.NodeID, payload, node.LastSeenAt)
activeConfig, err := GetActiveConfigMetaForAgent()
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
+2
View File
@@ -368,11 +368,13 @@ func TestOpenRestyProxyRequestBufferingDefaultsToOff(t *testing.T) {
func setupServiceTestDB(t *testing.T) {
t.Helper()
nodeAgentTokenCache.reset()
common.SQLitePath = filepath.Join(t.TempDir(), "service.db")
if err := model.InitDB(); err != nil {
t.Fatalf("failed to init db: %v", err)
}
t.Cleanup(func() {
nodeAgentTokenCache.reset()
if err := model.CloseDB(); err != nil {
t.Fatalf("failed to close db: %v", err)
}
+12 -2
View File
@@ -87,6 +87,7 @@ func CreateNode(input NodeInput) (*NodeView, error) {
}
return nil, err
}
refreshAgentTokenCache(node)
slog.Info("node created", "name", node.Name, "node_id", node.NodeID)
return buildNodeView(node), nil
}
@@ -113,6 +114,7 @@ func UpdateNode(id uint, input NodeInput) (*NodeView, error) {
if err = node.Update(); err != nil {
return nil, err
}
refreshAgentTokenCache(node)
slog.Info("node updated", "name", node.Name, "node_id", node.NodeID)
return buildNodeView(node), nil
}
@@ -123,7 +125,11 @@ func DeleteNode(id uint) error {
return err
}
slog.Info("node deleted", "name", node.Name, "node_id", node.NodeID)
return node.Delete()
if err := node.Delete(); err != nil {
return err
}
invalidateAgentTokenCache(node.AgentToken)
return nil
}
func GetNodeAgentRelease(ctx context.Context, id uint, channel string) (*NodeAgentReleaseInfo, error) {
@@ -163,6 +169,7 @@ func RequestNodeAgentUpdate(id uint, input NodeAgentUpdateInput) (*NodeView, err
if err = model.DB.Model(node).Select("update_requested", "update_channel", "update_tag").Updates(node).Error; err != nil {
return nil, err
}
refreshAgentTokenCache(node)
slog.Info("agent manual update requested", "node_id", node.NodeID, "name", node.Name, "channel", channel.String(), "tag", tagName)
return buildNodeView(node), nil
}
@@ -176,6 +183,7 @@ func RequestNodeOpenrestyRestart(id uint) (*NodeView, error) {
if err = model.DB.Model(node).Select("restart_openresty_requested").Updates(node).Error; err != nil {
return nil, err
}
refreshAgentTokenCache(node)
slog.Info("openresty restart requested", "node_id", node.NodeID, "name", node.Name)
return buildNodeView(node), nil
}
@@ -185,7 +193,7 @@ func AuthenticateAgentToken(token string) (*model.Node, error) {
if token == "" {
return nil, errors.New("缺少 Agent Token")
}
return model.GetNodeByAgentToken(token)
return authenticateAgentTokenWithCache(token)
}
func ValidateDiscoveryToken(token string) error {
@@ -357,6 +365,7 @@ func RegisterNodeWithAgentToken(node *model.Node, payload AgentNodePayload) (*Ag
if err := node.Update(); err != nil {
return nil, err
}
refreshAgentTokenCache(node)
slog.Info("agent register succeeded on reserved node", "node_id", node.NodeID, "name", node.Name)
return &AgentRegistrationResponse{
NodeID: node.NodeID,
@@ -394,6 +403,7 @@ func RegisterNodeWithDiscovery(payload AgentNodePayload) (*AgentRegistrationResp
}
return nil, err
}
refreshAgentTokenCache(node)
slog.Info("agent discovery register succeeded", "node_id", node.NodeID, "name", node.Name)
return &AgentRegistrationResponse{
NodeID: node.NodeID,
@@ -0,0 +1,177 @@
package service
import (
"atsflare/model"
"errors"
"time"
ristretto "github.com/dgraph-io/ristretto/v2"
"gorm.io/gorm"
)
const (
agentTokenPositiveCacheTTL = 2 * time.Minute
agentTokenNegativeCacheTTL = 10 * time.Minute
agentTokenNegativeCacheCap = 10000
)
type cachedAgentNode struct {
node *model.Node
expiresAt time.Time
}
type cachedMissingAgentToken struct {
expiresAt time.Time
}
type agentTokenAuthCache struct {
positive *ristretto.Cache[string, cachedAgentNode]
negative *ristretto.Cache[string, cachedMissingAgentToken]
now func() time.Time
loadNodeByToken func(string) (*model.Node, error)
}
var nodeAgentTokenCache = newAgentTokenAuthCache()
func newAgentTokenAuthCache() *agentTokenAuthCache {
return &agentTokenAuthCache{
positive: mustNewAgentTokenPositiveCache(),
negative: mustNewAgentTokenNegativeCache(),
now: time.Now,
loadNodeByToken: func(token string) (*model.Node, error) {
return model.GetNodeByAgentToken(token)
},
}
}
func mustNewAgentTokenPositiveCache() *ristretto.Cache[string, cachedAgentNode] {
cache, err := ristretto.NewCache(&ristretto.Config[string, cachedAgentNode]{
NumCounters: 1e5,
MaxCost: 2e4,
BufferItems: 64,
})
if err != nil {
panic(err)
}
return cache
}
func mustNewAgentTokenNegativeCache() *ristretto.Cache[string, cachedMissingAgentToken] {
cache, err := ristretto.NewCache(&ristretto.Config[string, cachedMissingAgentToken]{
NumCounters: 1e5,
MaxCost: agentTokenNegativeCacheCap,
BufferItems: 64,
})
if err != nil {
panic(err)
}
return cache
}
func (c *agentTokenAuthCache) authenticate(token string) (*model.Node, error) {
now := c.now()
if node, ok := c.getNode(token, now); ok {
return node, nil
}
if c.isMissing(token, now) {
return nil, gorm.ErrRecordNotFound
}
node, err := c.loadNodeByToken(token)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
c.storeMissing(token, now.Add(agentTokenNegativeCacheTTL))
}
return nil, err
}
c.storeNode(token, node, now.Add(agentTokenPositiveCacheTTL))
return cloneCachedNode(node), nil
}
func (c *agentTokenAuthCache) getNode(token string, now time.Time) (*model.Node, bool) {
entry, ok := c.positive.Get(token)
if !ok {
return nil, false
}
if now.After(entry.expiresAt) {
c.positive.Del(token)
return nil, false
}
return cloneCachedNode(entry.node), true
}
func (c *agentTokenAuthCache) isMissing(token string, now time.Time) bool {
entry, ok := c.negative.Get(token)
if !ok {
return false
}
if now.After(entry.expiresAt) {
c.negative.Del(token)
return false
}
return true
}
func (c *agentTokenAuthCache) storeNode(token string, node *model.Node, expiresAt time.Time) {
if token == "" || node == nil {
return
}
c.negative.Del(token)
c.positive.Set(token, cachedAgentNode{
node: cloneCachedNode(node),
expiresAt: expiresAt,
}, 1)
c.positive.Wait()
}
func (c *agentTokenAuthCache) storeMissing(token string, expiresAt time.Time) {
if token == "" {
return
}
c.positive.Del(token)
c.negative.Set(token, cachedMissingAgentToken{
expiresAt: expiresAt,
}, 1)
c.negative.Wait()
}
func (c *agentTokenAuthCache) invalidate(token string) {
if token == "" {
return
}
c.positive.Del(token)
c.negative.Del(token)
}
func (c *agentTokenAuthCache) reset() {
c.positive.Clear()
c.negative.Clear()
}
func cloneCachedNode(node *model.Node) *model.Node {
if node == nil {
return nil
}
cloned := *node
return &cloned
}
func authenticateAgentTokenWithCache(token string) (*model.Node, error) {
return nodeAgentTokenCache.authenticate(token)
}
func refreshAgentTokenCache(node *model.Node) {
if node == nil {
return
}
nodeAgentTokenCache.storeNode(
node.AgentToken,
node,
nodeAgentTokenCache.now().Add(agentTokenPositiveCacheTTL),
)
}
func invalidateAgentTokenCache(token string) {
nodeAgentTokenCache.invalidate(token)
}
@@ -0,0 +1,113 @@
package service
import (
"atsflare/model"
"errors"
"fmt"
"testing"
"time"
"gorm.io/gorm"
)
func TestAgentTokenAuthCacheUsesPositiveCacheUntilLogicalExpiry(t *testing.T) {
cache := newAgentTokenAuthCache()
cache.reset()
baseTime := time.Date(2026, 3, 14, 16, 0, 0, 0, time.UTC)
currentTime := baseTime
cache.now = func() time.Time {
return currentTime
}
loadCount := 0
cache.loadNodeByToken = func(token string) (*model.Node, error) {
loadCount++
return &model.Node{
NodeID: fmt.Sprintf("node-%d", loadCount),
Name: "edge",
AgentToken: token,
}, nil
}
first, err := cache.authenticate("token-a")
if err != nil {
t.Fatalf("expected first auth to succeed: %v", err)
}
if loadCount != 1 {
t.Fatalf("expected one db load, got %d", loadCount)
}
second, err := cache.authenticate("token-a")
if err != nil {
t.Fatalf("expected cached auth to succeed: %v", err)
}
if loadCount != 1 {
t.Fatalf("expected cache hit without db load, got %d", loadCount)
}
if first.NodeID != second.NodeID {
t.Fatalf("expected cached node to match original, got %s and %s", first.NodeID, second.NodeID)
}
currentTime = baseTime.Add(agentTokenPositiveCacheTTL + time.Second)
third, err := cache.authenticate("token-a")
if err != nil {
t.Fatalf("expected auth after expiry to succeed: %v", err)
}
if loadCount != 2 {
t.Fatalf("expected reload after logical expiry, got %d loads", loadCount)
}
if third.NodeID == second.NodeID {
t.Fatalf("expected refreshed cache entry after expiry, got unchanged node id %s", third.NodeID)
}
}
func TestAgentTokenAuthCacheRefreshesAfterMissingEntryExpires(t *testing.T) {
cache := newAgentTokenAuthCache()
cache.reset()
baseTime := time.Date(2026, 3, 14, 16, 30, 0, 0, time.UTC)
currentTime := baseTime
cache.now = func() time.Time {
return currentTime
}
loadCount := 0
cache.loadNodeByToken = func(token string) (*model.Node, error) {
loadCount++
if loadCount == 1 {
return nil, gorm.ErrRecordNotFound
}
return &model.Node{
NodeID: "node-recovered",
Name: "edge",
AgentToken: token,
}, nil
}
_, err := cache.authenticate("token-missing")
if !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("expected first lookup to miss, got %v", err)
}
if loadCount != 1 {
t.Fatalf("expected one db load for first miss, got %d", loadCount)
}
_, err = cache.authenticate("token-missing")
if !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("expected cached missing lookup to miss, got %v", err)
}
if loadCount != 1 {
t.Fatalf("expected missing cache hit without db load, got %d", loadCount)
}
currentTime = baseTime.Add(agentTokenNegativeCacheTTL + time.Second)
node, err := cache.authenticate("token-missing")
if err != nil {
t.Fatalf("expected lookup after missing expiry to reload successfully: %v", err)
}
if loadCount != 2 {
t.Fatalf("expected db reload after missing cache expiry, got %d", loadCount)
}
if node.NodeID != "node-recovered" {
t.Fatalf("unexpected recovered node: %+v", node)
}
}