mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat(coderd): batch agent stats inserts (#8875)
This PR adds support for batching inserts to the workspace_agents_stats table. Up to 1024 stats are batched, and flushed every second in a batch.
This commit is contained in:
@@ -0,0 +1,289 @@
|
||||
package batchstats
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/sloghuman"
|
||||
|
||||
"github.com/coder/coder/coderd/database"
|
||||
"github.com/coder/coder/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/codersdk/agentsdk"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultBufferSize = 1024
|
||||
defaultFlushInterval = time.Second
|
||||
)
|
||||
|
||||
// Batcher holds a buffer of agent stats and periodically flushes them to
|
||||
// its configured store. It also updates the workspace's last used time.
|
||||
type Batcher struct {
|
||||
store database.Store
|
||||
log slog.Logger
|
||||
|
||||
mu sync.Mutex
|
||||
// TODO: make this a buffered chan instead?
|
||||
buf *database.InsertWorkspaceAgentStatsParams
|
||||
// NOTE: we batch this separately as it's a jsonb field and
|
||||
// pq.Array + unnest doesn't play nicely with this.
|
||||
connectionsByProto []map[string]int64
|
||||
batchSize int
|
||||
|
||||
// tickCh is used to periodically flush the buffer.
|
||||
tickCh <-chan time.Time
|
||||
ticker *time.Ticker
|
||||
interval time.Duration
|
||||
// flushLever is used to signal the flusher to flush the buffer immediately.
|
||||
flushLever chan struct{}
|
||||
flushForced atomic.Bool
|
||||
// flushed is used during testing to signal that a flush has completed.
|
||||
flushed chan<- int
|
||||
}
|
||||
|
||||
// Option is a functional option for configuring a Batcher.
|
||||
type Option func(b *Batcher)
|
||||
|
||||
// WithStore sets the store to use for storing stats.
|
||||
func WithStore(store database.Store) Option {
|
||||
return func(b *Batcher) {
|
||||
b.store = store
|
||||
}
|
||||
}
|
||||
|
||||
// WithBatchSize sets the number of stats to store in a batch.
|
||||
func WithBatchSize(size int) Option {
|
||||
return func(b *Batcher) {
|
||||
b.batchSize = size
|
||||
}
|
||||
}
|
||||
|
||||
// WithInterval sets the interval for flushes.
|
||||
func WithInterval(d time.Duration) Option {
|
||||
return func(b *Batcher) {
|
||||
b.interval = d
|
||||
}
|
||||
}
|
||||
|
||||
// WithLogger sets the logger to use for logging.
|
||||
func WithLogger(log slog.Logger) Option {
|
||||
return func(b *Batcher) {
|
||||
b.log = log
|
||||
}
|
||||
}
|
||||
|
||||
// New creates a new Batcher and starts it.
|
||||
func New(ctx context.Context, opts ...Option) (*Batcher, func(), error) {
|
||||
b := &Batcher{}
|
||||
b.log = slog.Make(sloghuman.Sink(os.Stderr))
|
||||
b.flushLever = make(chan struct{}, 1) // Buffered so that it doesn't block.
|
||||
for _, opt := range opts {
|
||||
opt(b)
|
||||
}
|
||||
|
||||
if b.store == nil {
|
||||
return nil, nil, xerrors.Errorf("no store configured for batcher")
|
||||
}
|
||||
|
||||
if b.interval == 0 {
|
||||
b.interval = defaultFlushInterval
|
||||
}
|
||||
|
||||
if b.batchSize == 0 {
|
||||
b.batchSize = defaultBufferSize
|
||||
}
|
||||
|
||||
if b.tickCh == nil {
|
||||
b.ticker = time.NewTicker(b.interval)
|
||||
b.tickCh = b.ticker.C
|
||||
}
|
||||
|
||||
cancelCtx, cancelFunc := context.WithCancel(ctx)
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
b.run(cancelCtx)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
closer := func() {
|
||||
cancelFunc()
|
||||
if b.ticker != nil {
|
||||
b.ticker.Stop()
|
||||
}
|
||||
<-done
|
||||
}
|
||||
|
||||
return b, closer, nil
|
||||
}
|
||||
|
||||
// Add adds a stat to the batcher for the given workspace and agent.
|
||||
func (b *Batcher) Add(
|
||||
agentID uuid.UUID,
|
||||
templateID uuid.UUID,
|
||||
userID uuid.UUID,
|
||||
workspaceID uuid.UUID,
|
||||
st agentsdk.Stats,
|
||||
) error {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
|
||||
now := database.Now()
|
||||
|
||||
b.buf.ID = append(b.buf.ID, uuid.New())
|
||||
b.buf.CreatedAt = append(b.buf.CreatedAt, now)
|
||||
b.buf.AgentID = append(b.buf.AgentID, agentID)
|
||||
b.buf.UserID = append(b.buf.UserID, userID)
|
||||
b.buf.TemplateID = append(b.buf.TemplateID, templateID)
|
||||
b.buf.WorkspaceID = append(b.buf.WorkspaceID, workspaceID)
|
||||
|
||||
// Store the connections by proto separately as it's a jsonb field. We marshal on flush.
|
||||
// b.buf.ConnectionsByProto = append(b.buf.ConnectionsByProto, st.ConnectionsByProto)
|
||||
b.connectionsByProto = append(b.connectionsByProto, st.ConnectionsByProto)
|
||||
|
||||
b.buf.ConnectionCount = append(b.buf.ConnectionCount, st.ConnectionCount)
|
||||
b.buf.RxPackets = append(b.buf.RxPackets, st.RxPackets)
|
||||
b.buf.RxBytes = append(b.buf.RxBytes, st.RxBytes)
|
||||
b.buf.TxPackets = append(b.buf.TxPackets, st.TxPackets)
|
||||
b.buf.TxBytes = append(b.buf.TxBytes, st.TxBytes)
|
||||
b.buf.SessionCountVSCode = append(b.buf.SessionCountVSCode, st.SessionCountVSCode)
|
||||
b.buf.SessionCountJetBrains = append(b.buf.SessionCountJetBrains, st.SessionCountJetBrains)
|
||||
b.buf.SessionCountReconnectingPTY = append(b.buf.SessionCountReconnectingPTY, st.SessionCountReconnectingPTY)
|
||||
b.buf.SessionCountSSH = append(b.buf.SessionCountSSH, st.SessionCountSSH)
|
||||
b.buf.ConnectionMedianLatencyMS = append(b.buf.ConnectionMedianLatencyMS, st.ConnectionMedianLatencyMS)
|
||||
|
||||
// If the buffer is over 80% full, signal the flusher to flush immediately.
|
||||
// We want to trigger flushes early to reduce the likelihood of
|
||||
// accidentally growing the buffer over batchSize.
|
||||
filled := float64(len(b.buf.ID)) / float64(b.batchSize)
|
||||
if filled >= 0.8 && !b.flushForced.Load() {
|
||||
b.flushLever <- struct{}{}
|
||||
b.flushForced.Store(true)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Run runs the batcher.
|
||||
func (b *Batcher) run(ctx context.Context) {
|
||||
b.initBuf(b.batchSize)
|
||||
// nolint:gocritic // This is only ever used for one thing - inserting agent stats.
|
||||
authCtx := dbauthz.AsSystemRestricted(ctx)
|
||||
for {
|
||||
select {
|
||||
case <-b.tickCh:
|
||||
b.flush(authCtx, false, "scheduled")
|
||||
case <-b.flushLever:
|
||||
// If the flush lever is depressed, flush the buffer immediately.
|
||||
b.flush(authCtx, true, "reaching capacity")
|
||||
case <-ctx.Done():
|
||||
b.log.Warn(ctx, "context done, flushing before exit")
|
||||
b.flush(authCtx, true, "exit")
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// flush flushes the batcher's buffer.
|
||||
func (b *Batcher) flush(ctx context.Context, forced bool, reason string) {
|
||||
b.mu.Lock()
|
||||
b.flushForced.Store(true)
|
||||
start := time.Now()
|
||||
count := len(b.buf.ID)
|
||||
defer func() {
|
||||
b.flushForced.Store(false)
|
||||
b.mu.Unlock()
|
||||
// Notify that a flush has completed. This only happens in tests.
|
||||
if b.flushed != nil {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
close(b.flushed)
|
||||
default:
|
||||
b.flushed <- count
|
||||
}
|
||||
}
|
||||
if count > 0 {
|
||||
elapsed := time.Since(start)
|
||||
b.log.Debug(ctx, "flush complete",
|
||||
slog.F("count", count),
|
||||
slog.F("elapsed", elapsed),
|
||||
slog.F("forced", forced),
|
||||
slog.F("reason", reason),
|
||||
)
|
||||
}
|
||||
}()
|
||||
|
||||
if len(b.buf.ID) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// marshal connections by proto
|
||||
payload, err := json.Marshal(b.connectionsByProto)
|
||||
if err != nil {
|
||||
b.log.Error(ctx, "unable to marshal agent connections by proto, dropping data", slog.Error(err))
|
||||
b.buf.ConnectionsByProto = json.RawMessage(`[]`)
|
||||
} else {
|
||||
b.buf.ConnectionsByProto = payload
|
||||
}
|
||||
|
||||
err = b.store.InsertWorkspaceAgentStats(ctx, *b.buf)
|
||||
elapsed := time.Since(start)
|
||||
if err != nil {
|
||||
b.log.Error(ctx, "error inserting workspace agent stats", slog.Error(err), slog.F("elapsed", elapsed))
|
||||
return
|
||||
}
|
||||
|
||||
b.resetBuf()
|
||||
}
|
||||
|
||||
// initBuf resets the buffer. b MUST be locked.
|
||||
func (b *Batcher) initBuf(size int) {
|
||||
b.buf = &database.InsertWorkspaceAgentStatsParams{
|
||||
ID: make([]uuid.UUID, 0, b.batchSize),
|
||||
CreatedAt: make([]time.Time, 0, b.batchSize),
|
||||
UserID: make([]uuid.UUID, 0, b.batchSize),
|
||||
WorkspaceID: make([]uuid.UUID, 0, b.batchSize),
|
||||
TemplateID: make([]uuid.UUID, 0, b.batchSize),
|
||||
AgentID: make([]uuid.UUID, 0, b.batchSize),
|
||||
ConnectionsByProto: json.RawMessage("[]"),
|
||||
ConnectionCount: make([]int64, 0, b.batchSize),
|
||||
RxPackets: make([]int64, 0, b.batchSize),
|
||||
RxBytes: make([]int64, 0, b.batchSize),
|
||||
TxPackets: make([]int64, 0, b.batchSize),
|
||||
TxBytes: make([]int64, 0, b.batchSize),
|
||||
SessionCountVSCode: make([]int64, 0, b.batchSize),
|
||||
SessionCountJetBrains: make([]int64, 0, b.batchSize),
|
||||
SessionCountReconnectingPTY: make([]int64, 0, b.batchSize),
|
||||
SessionCountSSH: make([]int64, 0, b.batchSize),
|
||||
ConnectionMedianLatencyMS: make([]float64, 0, b.batchSize),
|
||||
}
|
||||
|
||||
b.connectionsByProto = make([]map[string]int64, 0, size)
|
||||
}
|
||||
|
||||
func (b *Batcher) resetBuf() {
|
||||
b.buf.ID = b.buf.ID[:0]
|
||||
b.buf.CreatedAt = b.buf.CreatedAt[:0]
|
||||
b.buf.UserID = b.buf.UserID[:0]
|
||||
b.buf.WorkspaceID = b.buf.WorkspaceID[:0]
|
||||
b.buf.TemplateID = b.buf.TemplateID[:0]
|
||||
b.buf.AgentID = b.buf.AgentID[:0]
|
||||
b.buf.ConnectionsByProto = json.RawMessage(`[]`)
|
||||
b.buf.ConnectionCount = b.buf.ConnectionCount[:0]
|
||||
b.buf.RxPackets = b.buf.RxPackets[:0]
|
||||
b.buf.RxBytes = b.buf.RxBytes[:0]
|
||||
b.buf.TxPackets = b.buf.TxPackets[:0]
|
||||
b.buf.TxBytes = b.buf.TxBytes[:0]
|
||||
b.buf.SessionCountVSCode = b.buf.SessionCountVSCode[:0]
|
||||
b.buf.SessionCountJetBrains = b.buf.SessionCountJetBrains[:0]
|
||||
b.buf.SessionCountReconnectingPTY = b.buf.SessionCountReconnectingPTY[:0]
|
||||
b.buf.SessionCountSSH = b.buf.SessionCountSSH[:0]
|
||||
b.buf.ConnectionMedianLatencyMS = b.buf.ConnectionMedianLatencyMS[:0]
|
||||
b.connectionsByProto = b.connectionsByProto[:0]
|
||||
}
|
||||
@@ -0,0 +1,226 @@
|
||||
package batchstats
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
|
||||
"github.com/coder/coder/coderd/database"
|
||||
"github.com/coder/coder/coderd/database/dbgen"
|
||||
"github.com/coder/coder/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/coderd/rbac"
|
||||
"github.com/coder/coder/codersdk/agentsdk"
|
||||
"github.com/coder/coder/cryptorand"
|
||||
)
|
||||
|
||||
func TestBatchStats(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Given: a fresh batcher with no data
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
log := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
|
||||
store, _ := dbtestutil.NewDB(t)
|
||||
|
||||
// Set up some test dependencies.
|
||||
deps1 := setupDeps(t, store)
|
||||
deps2 := setupDeps(t, store)
|
||||
tick := make(chan time.Time)
|
||||
flushed := make(chan int)
|
||||
|
||||
b, closer, err := New(ctx,
|
||||
WithStore(store),
|
||||
WithLogger(log),
|
||||
func(b *Batcher) {
|
||||
b.tickCh = tick
|
||||
b.flushed = flushed
|
||||
},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(closer)
|
||||
|
||||
// Given: no data points are added for workspace
|
||||
// When: it becomes time to report stats
|
||||
t1 := time.Now()
|
||||
// Signal a tick and wait for a flush to complete.
|
||||
tick <- t1
|
||||
f := <-flushed
|
||||
require.Equal(t, 0, f, "expected no data to be flushed")
|
||||
t.Logf("flush 1 completed")
|
||||
|
||||
// Then: it should report no stats.
|
||||
stats, err := store.GetWorkspaceAgentStats(ctx, t1)
|
||||
require.NoError(t, err, "should not error getting stats")
|
||||
require.Empty(t, stats, "should have no stats for workspace")
|
||||
|
||||
// Given: a single data point is added for workspace
|
||||
t2 := time.Now()
|
||||
t.Logf("inserting 1 stat")
|
||||
require.NoError(t, b.Add(deps1.Agent.ID, deps1.User.ID, deps1.Template.ID, deps1.Workspace.ID, randAgentSDKStats(t)))
|
||||
|
||||
// When: it becomes time to report stats
|
||||
// Signal a tick and wait for a flush to complete.
|
||||
tick <- t2
|
||||
f = <-flushed // Wait for a flush to complete.
|
||||
require.Equal(t, 1, f, "expected one stat to be flushed")
|
||||
t.Logf("flush 2 completed")
|
||||
|
||||
// Then: it should report a single stat.
|
||||
stats, err = store.GetWorkspaceAgentStats(ctx, t2)
|
||||
require.NoError(t, err, "should not error getting stats")
|
||||
require.Len(t, stats, 1, "should have stats for workspace")
|
||||
|
||||
// Given: a lot of data points are added for both workspaces
|
||||
// (equal to batch size)
|
||||
t3 := time.Now()
|
||||
done := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
defer close(done)
|
||||
t.Logf("inserting %d stats", defaultBufferSize)
|
||||
for i := 0; i < defaultBufferSize; i++ {
|
||||
if i%2 == 0 {
|
||||
require.NoError(t, b.Add(deps1.Agent.ID, deps1.User.ID, deps1.Template.ID, deps1.Workspace.ID, randAgentSDKStats(t)))
|
||||
} else {
|
||||
require.NoError(t, b.Add(deps2.Agent.ID, deps2.User.ID, deps2.Template.ID, deps2.Workspace.ID, randAgentSDKStats(t)))
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
// When: the buffer comes close to capacity
|
||||
// Then: The buffer will force-flush once.
|
||||
f = <-flushed
|
||||
t.Logf("flush 3 completed")
|
||||
require.Greater(t, f, 819, "expected at least 819 stats to be flushed (>=80% of buffer)")
|
||||
// And we should finish inserting the stats
|
||||
<-done
|
||||
|
||||
stats, err = store.GetWorkspaceAgentStats(ctx, t3)
|
||||
require.NoError(t, err, "should not error getting stats")
|
||||
require.Len(t, stats, 2, "should have stats for both workspaces")
|
||||
|
||||
// Ensures that a subsequent flush pushes all the remaining data
|
||||
t4 := time.Now()
|
||||
tick <- t4
|
||||
f2 := <-flushed
|
||||
t.Logf("flush 4 completed")
|
||||
expectedCount := defaultBufferSize - f
|
||||
require.Equal(t, expectedCount, f2, "did not flush expected remaining rows")
|
||||
|
||||
// Ensure that a subsequent flush does not push stale data.
|
||||
t5 := time.Now()
|
||||
tick <- t5
|
||||
f = <-flushed
|
||||
require.Zero(t, f, "expected zero stats to have been flushed")
|
||||
t.Logf("flush 5 completed")
|
||||
|
||||
stats, err = store.GetWorkspaceAgentStats(ctx, t5)
|
||||
require.NoError(t, err, "should not error getting stats")
|
||||
require.Len(t, stats, 0, "should have no stats for workspace")
|
||||
|
||||
// Ensure that buf never grew beyond what we expect
|
||||
require.Equal(t, defaultBufferSize, cap(b.buf.ID), "buffer grew beyond expected capacity")
|
||||
}
|
||||
|
||||
// randAgentSDKStats returns a random agentsdk.Stats
|
||||
func randAgentSDKStats(t *testing.T, opts ...func(*agentsdk.Stats)) agentsdk.Stats {
|
||||
t.Helper()
|
||||
s := agentsdk.Stats{
|
||||
ConnectionsByProto: map[string]int64{
|
||||
"ssh": mustRandInt64n(t, 9) + 1,
|
||||
"vscode": mustRandInt64n(t, 9) + 1,
|
||||
"jetbrains": mustRandInt64n(t, 9) + 1,
|
||||
"reconnecting_pty": mustRandInt64n(t, 9) + 1,
|
||||
},
|
||||
ConnectionCount: mustRandInt64n(t, 99) + 1,
|
||||
ConnectionMedianLatencyMS: float64(mustRandInt64n(t, 99) + 1),
|
||||
RxPackets: mustRandInt64n(t, 99) + 1,
|
||||
RxBytes: mustRandInt64n(t, 99) + 1,
|
||||
TxPackets: mustRandInt64n(t, 99) + 1,
|
||||
TxBytes: mustRandInt64n(t, 99) + 1,
|
||||
SessionCountVSCode: mustRandInt64n(t, 9) + 1,
|
||||
SessionCountJetBrains: mustRandInt64n(t, 9) + 1,
|
||||
SessionCountReconnectingPTY: mustRandInt64n(t, 9) + 1,
|
||||
SessionCountSSH: mustRandInt64n(t, 9) + 1,
|
||||
Metrics: []agentsdk.AgentMetric{},
|
||||
}
|
||||
for _, opt := range opts {
|
||||
opt(&s)
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// deps is a set of test dependencies.
|
||||
type deps struct {
|
||||
Agent database.WorkspaceAgent
|
||||
Template database.Template
|
||||
User database.User
|
||||
Workspace database.Workspace
|
||||
}
|
||||
|
||||
// setupDeps sets up a set of test dependencies.
|
||||
// It creates an organization, user, template, workspace, and agent
|
||||
// along with all the other miscellaneous plumbing required to link
|
||||
// them together.
|
||||
func setupDeps(t *testing.T, store database.Store) deps {
|
||||
t.Helper()
|
||||
|
||||
org := dbgen.Organization(t, store, database.Organization{})
|
||||
user := dbgen.User(t, store, database.User{})
|
||||
_, err := store.InsertOrganizationMember(context.Background(), database.InsertOrganizationMemberParams{
|
||||
OrganizationID: org.ID,
|
||||
UserID: user.ID,
|
||||
Roles: []string{rbac.RoleOrgMember(org.ID)},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
tv := dbgen.TemplateVersion(t, store, database.TemplateVersion{
|
||||
OrganizationID: org.ID,
|
||||
CreatedBy: user.ID,
|
||||
})
|
||||
tpl := dbgen.Template(t, store, database.Template{
|
||||
CreatedBy: user.ID,
|
||||
OrganizationID: org.ID,
|
||||
ActiveVersionID: tv.ID,
|
||||
})
|
||||
ws := dbgen.Workspace(t, store, database.Workspace{
|
||||
TemplateID: tpl.ID,
|
||||
OwnerID: user.ID,
|
||||
OrganizationID: org.ID,
|
||||
LastUsedAt: time.Now().Add(-time.Hour),
|
||||
})
|
||||
pj := dbgen.ProvisionerJob(t, store, database.ProvisionerJob{
|
||||
InitiatorID: user.ID,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
_ = dbgen.WorkspaceBuild(t, store, database.WorkspaceBuild{
|
||||
TemplateVersionID: tv.ID,
|
||||
WorkspaceID: ws.ID,
|
||||
JobID: pj.ID,
|
||||
})
|
||||
res := dbgen.WorkspaceResource(t, store, database.WorkspaceResource{
|
||||
Transition: database.WorkspaceTransitionStart,
|
||||
JobID: pj.ID,
|
||||
})
|
||||
agt := dbgen.WorkspaceAgent(t, store, database.WorkspaceAgent{
|
||||
ResourceID: res.ID,
|
||||
})
|
||||
return deps{
|
||||
Agent: agt,
|
||||
Template: tpl,
|
||||
User: user,
|
||||
Workspace: ws,
|
||||
}
|
||||
}
|
||||
|
||||
// mustRandInt64n returns a random int64 in the range [0, n).
|
||||
func mustRandInt64n(t *testing.T, n int64) int64 {
|
||||
t.Helper()
|
||||
i, err := cryptorand.Intn(int(n))
|
||||
require.NoError(t, err)
|
||||
return int64(i)
|
||||
}
|
||||
Reference in New Issue
Block a user