diff --git a/cmd/ax-controller/main.go b/cmd/ax-controller/main.go index 9d1020ec..35f7dcf5 100644 --- a/cmd/ax-controller/main.go +++ b/cmd/ax-controller/main.go @@ -42,12 +42,14 @@ func main() { redisPassword string redisGroup string redisConsumer string + concurrency int ) flag.StringVar(&redisAddr, "redis-addr", "localhost:6379", "Redis server address (e.g. localhost:6379)") flag.StringVar(&redisPassword, "redis-password", "", "Redis password") flag.StringVar(&redisGroup, "redis-group", "ax-controllers", "Redis stream consumer group") flag.StringVar(&redisConsumer, "redis-consumer", "", "Redis stream consumer ID (defaults to hostname)") + flag.IntVar(&concurrency, "concurrency", 16, "Number of tasks reconciled in parallel by this controller") flag.StringVar(&substrateEndpoint, "substrate-endpoint", "api.ate-system.svc.cluster.local:443", "Agent Substrate Control API endpoint") flag.StringVar(&substrateAuthority, "substrate-authority", "api.ate-system.svc", "Authority / TLS ServerName for Substrate endpoint") flag.StringVar(&substrateTokenFile, "substrate-token-file", "", "Path to bearer token file for Substrate auth") @@ -104,6 +106,7 @@ func main() { rStore := redis.NewStore(rClient, redis.Options{}) worker := controller.NewWorker(rStore, reconciler, redisGroup, redisConsumer) + worker.Concurrency = concurrency if err := worker.Run(ctx); err != nil && err != context.Canceled { slog.Error("redis worker stopped with error", "error", err) os.Exit(1) diff --git a/internal/controller/worker.go b/internal/controller/worker.go index 8af08e50..80a3ab5f 100644 --- a/internal/controller/worker.go +++ b/internal/controller/worker.go @@ -18,8 +18,10 @@ import ( "context" "errors" "fmt" + "hash/fnv" "log/slog" "os" + "sync" "time" "github.com/google/ax/internal/store" @@ -31,6 +33,10 @@ const ( // readRetryDelay is how long the worker waits after a transient error from the // event queue before trying again. readRetryDelay = time.Second + // slotQueueSize is how many events may wait for one busy slot before the worker + // stops reading. It keeps a task with several queued events from stalling the + // other slots while bounding how much the worker holds unacknowledged. + slotQueueSize = 64 ) // Worker consumes task events from the store's event queue and reconciles each @@ -41,6 +47,11 @@ type Worker struct { reconciler *TaskReconciler group string consumer string + + // Concurrency is how many events the worker reconciles at once. Events for the + // same task always go to the same slot, in order, so a task is never reconciled + // concurrently with itself. Values below 1 mean 1. + Concurrency int } // NewWorker creates a worker that joins group as consumer. An empty group uses the @@ -62,10 +73,11 @@ func NewWorker(s store.Store, reconciler *TaskReconciler, group, consumer string } // Run subscribes to task events and processes them until ctx is done. It returns -// ctx.Err() on shutdown; every event is acknowledged after processing, even when -// reconciliation fails, so a bad task cannot wedge the queue. +// ctx.Err() on shutdown, after in-flight events finish; every event is acknowledged +// after processing, even when reconciliation fails, so a bad task cannot wedge the queue. func (w *Worker) Run(ctx context.Context) error { - slog.Info("starting AX task worker", "group", w.group, "consumer", w.consumer) + concurrency := max(w.Concurrency, 1) + slog.Info("starting AX task worker", "group", w.group, "consumer", w.consumer, "concurrency", concurrency) sub, err := w.store.Subscribe(ctx, w.group, w.consumer) if err != nil { @@ -73,6 +85,27 @@ func (w *Worker) Run(ctx context.Context) error { } defer sub.Close() + // Each slot is a goroutine with its own queue. Hashing the task key to a slot + // keeps per-task ordering while unrelated tasks reconcile in parallel. + slots := make([]chan store.TaskEvent, concurrency) + var wg sync.WaitGroup + for i := range slots { + slots[i] = make(chan store.TaskEvent, slotQueueSize) + wg.Add(1) + go func(events <-chan store.TaskEvent) { + defer wg.Done() + for ev := range events { + w.handle(ctx, sub, ev) + } + }(slots[i]) + } + defer func() { + for _, ch := range slots { + close(ch) + } + wg.Wait() + }() + for { ev, err := sub.Next(ctx) if err != nil { @@ -89,21 +122,44 @@ func (w *Worker) Run(ctx context.Context) error { continue } - if err := w.processEvent(ctx, ev); err != nil { - slog.Error("error processing task event", - "id", ev.ID, - "atespace", ev.Atespace, - "name", ev.Name, - "action", ev.Action, - "error", err, - ) - } - if err := sub.Ack(ctx, ev); err != nil { - slog.Warn("failed to acknowledge task event", "id", ev.ID, "error", err) + select { + case slots[slotFor(ev, concurrency)] <- ev: + case <-ctx.Done(): + // Unacknowledged, so the event stays pending for the group. + return ctx.Err() } } } +// handle reconciles one event and acknowledges it. +func (w *Worker) handle(ctx context.Context, sub store.Subscription, ev store.TaskEvent) { + if err := w.processEvent(ctx, ev); err != nil { + slog.Error("error processing task event", + "id", ev.ID, + "atespace", ev.Atespace, + "name", ev.Name, + "action", ev.Action, + "error", err, + ) + } + if err := sub.Ack(ctx, ev); err != nil { + slog.Warn("failed to acknowledge task event", "id", ev.ID, "error", err) + } +} + +// slotFor maps an event to a worker slot by task, so every event for a given +// task is processed by the same slot. +func slotFor(ev store.TaskEvent, n int) int { + if n <= 1 { + return 0 + } + h := fnv.New32a() + h.Write([]byte(ev.Atespace)) + h.Write([]byte{0}) + h.Write([]byte(ev.Name)) + return int(h.Sum32() % uint32(n)) +} + func (w *Worker) processEvent(ctx context.Context, ev store.TaskEvent) error { if ev.Action == "delete" { slog.Info("handling task deletion event", "atespace", ev.Atespace, "name", ev.Name) diff --git a/internal/controller/worker_bench_test.go b/internal/controller/worker_bench_test.go new file mode 100644 index 00000000..3fc0483c --- /dev/null +++ b/internal/controller/worker_bench_test.go @@ -0,0 +1,260 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package controller_test + +import ( + "context" + "fmt" + "net" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/agent-substrate/substrate/pkg/proto/ateapipb" + "github.com/google/ax/internal/controller" + "github.com/google/ax/internal/store/memory" + "github.com/google/ax/internal/substrate" + "github.com/google/ax/pkg/apis/v1alpha1" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/status" +) + +// slowControlServer is a Substrate stand-in that adds a fixed latency to every RPC +// on the reconcile path and records the peak number of concurrent calls per actor. +type slowControlServer struct { + ateapipb.UnimplementedControlServer + latency time.Duration + workerIP string + + mu sync.Mutex + inFlight map[string]int + maxPerTask int + resumes int +} + +func (s *slowControlServer) enter(actor string) func() { + s.mu.Lock() + s.inFlight[actor]++ + s.maxPerTask = max(s.maxPerTask, s.inFlight[actor]) + s.mu.Unlock() + time.Sleep(s.latency) + return func() { + s.mu.Lock() + s.inFlight[actor]-- + s.mu.Unlock() + } +} + +func (s *slowControlServer) CreateAtespace(ctx context.Context, req *ateapipb.CreateAtespaceRequest) (*ateapipb.Atespace, error) { + time.Sleep(s.latency) + return &ateapipb.Atespace{}, nil +} + +func (s *slowControlServer) GetActorTemplate(ctx context.Context, req *ateapipb.GetActorTemplateRequest) (*ateapipb.ActorTemplate, error) { + time.Sleep(s.latency) + return nil, status.Error(codes.NotFound, "not found") +} + +func (s *slowControlServer) CreateActorTemplate(ctx context.Context, req *ateapipb.CreateActorTemplateRequest) (*ateapipb.ActorTemplate, error) { + time.Sleep(s.latency) + return req.GetActorTemplate(), nil +} + +func (s *slowControlServer) CreateActor(ctx context.Context, req *ateapipb.CreateActorRequest) (*ateapipb.Actor, error) { + defer s.enter(req.GetActor().GetMetadata().GetName())() + return &ateapipb.Actor{ + Metadata: req.GetActor().GetMetadata(), + Status: &ateapipb.ActorStatus{State: ateapipb.ActorState_ACTOR_STATE_SUSPENDED}, + }, nil +} + +func (s *slowControlServer) CreateActorEgressPolicy(ctx context.Context, req *ateapipb.CreateActorEgressPolicyRequest) (*ateapipb.EgressPolicy, error) { + defer s.enter(req.GetActor().GetName())() + return &ateapipb.EgressPolicy{}, nil +} + +func (s *slowControlServer) ResumeActor(ctx context.Context, req *ateapipb.ResumeActorRequest) (*ateapipb.ResumeActorResponse, error) { + defer s.enter(req.GetActor().GetName())() + s.mu.Lock() + s.resumes++ + s.mu.Unlock() + return &ateapipb.ResumeActorResponse{ + Actor: &ateapipb.Actor{ + Metadata: &ateapipb.ResourceMetadata{Name: req.GetActor().GetName()}, + Status: &ateapipb.ActorStatus{ + State: ateapipb.ActorState_ACTOR_STATE_RUNNING, + WorkerAssignment: &ateapipb.WorkerAssignment{WorkerPodIp: s.workerIP}, + }, + }, + }, nil +} + +type workerHarness struct { + store *memory.MemoryStore + server *slowControlServer + worker *controller.Worker +} + +// newWorkerHarness wires a worker to a slow Substrate and a runner /readyz that +// reports the workspace ready immediately (warm) or never (cold, so every +// reconcile polls until readyTimeout). +func newWorkerHarness(tb testing.TB, concurrency int, latency, readyTimeout time.Duration, warm bool) *workerHarness { + tb.Helper() + readyz := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !warm { + w.WriteHeader(http.StatusServiceUnavailable) + } + })) + tb.Cleanup(readyz.Close) + + lis, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + tb.Fatalf("failed to listen: %v", err) + } + srv := &slowControlServer{ + latency: latency, + workerIP: strings.TrimPrefix(readyz.URL, "http://"), + inFlight: map[string]int{}, + } + grpcServer := grpc.NewServer() + ateapipb.RegisterControlServer(grpcServer, srv) + go grpcServer.Serve(lis) + tb.Cleanup(grpcServer.Stop) + + subClient, err := substrate.NewClient(lis.Addr().String(), grpc.WithTransportCredentials(insecure.NewCredentials())) + if err != nil { + tb.Fatalf("failed to create substrate client: %v", err) + } + tb.Cleanup(func() { subClient.Close() }) + + reconciler := controller.NewTaskReconciler(subClient, "default-template", "ax-system") + reconciler.SecretResolver = noSecrets + reconciler.WorkspaceReadyTimeout = readyTimeout + + memStore := memory.NewStore() + worker := controller.NewWorker(memStore, reconciler, "bench", "bench-1") + worker.Concurrency = concurrency + return &workerHarness{store: memStore, server: srv, worker: worker} +} + +// applyAndWait saves the named tasks, then waits until every one has been +// reconciled (its status reports a worker IP). +func (h *workerHarness) applyAndWait(tb testing.TB, ctx context.Context, names []string, timeout time.Duration) { + tb.Helper() + for _, name := range names { + if err := h.store.SaveTask(ctx, &v1alpha1.Task{ + Metadata: &v1alpha1.ObjectMeta{Name: name, Atespace: "default"}, + Spec: &v1alpha1.TaskSpec{Image: "ghcr.io/test/img"}, + }); err != nil { + tb.Fatalf("SaveTask(%s): %v", name, err) + } + } + deadline := time.Now().Add(timeout) + for _, name := range names { + for { + task, err := h.store.GetTask(ctx, "default", name) + if err == nil && task.GetStatus().GetWorkerIp() != "" { + break + } + if time.Now().After(deadline) { + tb.Fatalf("task %s not reconciled within %v", name, timeout) + } + time.Sleep(5 * time.Millisecond) + } + } +} + +// BenchmarkWorkerThroughput measures how many freshly applied tasks one controller +// process reconciles per second. Substrate RPCs take 10ms each; in the cold case +// the workspace never reports ready, so every reconcile also waits out a 200ms +// readiness poll (scaled down from the 15s default). +// +// go test ./internal/controller -run '^$' -bench WorkerThroughput -benchtime 1x +func BenchmarkWorkerThroughput(b *testing.B) { + const batch = 64 + for _, warm := range []bool{true, false} { + for _, concurrency := range []int{1, 4, 16, 64} { + workspace := "cold" + if warm { + workspace = "warm" + } + b.Run(fmt.Sprintf("workspace=%s/concurrency=%d", workspace, concurrency), func(b *testing.B) { + h := newWorkerHarness(b, concurrency, 10*time.Millisecond, 200*time.Millisecond, warm) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go func() { _ = h.worker.Run(ctx) }() + + b.ResetTimer() + for i := 0; i < b.N; i++ { + names := make([]string, batch) + for j := range names { + names[j] = fmt.Sprintf("task-%d-%d", i, j) + } + h.applyAndWait(b, ctx, names, 5*time.Minute) + } + b.ReportMetric(float64(batch*b.N)/b.Elapsed().Seconds(), "tasks/s") + }) + } + } +} + +// TestWorkerConcurrencyKeepsTasksSerial applies several revisions of each task at +// once and checks that no task ever has two reconciles in flight, while distinct +// tasks still make progress in parallel. +func TestWorkerConcurrencyKeepsTasksSerial(t *testing.T) { + h := newWorkerHarness(t, 8, 20*time.Millisecond, 50*time.Millisecond, true) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + go func() { _ = h.worker.Run(ctx) }() + + const tasks, revisions = 16, 3 + start := time.Now() + // Queue every revision up front, back to back, so the same task sits in the + // queue several times in a row. + for i := 0; i < tasks; i++ { + for rev := 0; rev < revisions; rev++ { + if err := h.store.SaveTask(ctx, &v1alpha1.Task{ + Metadata: &v1alpha1.ObjectMeta{Name: fmt.Sprintf("serial-%d", i), Atespace: "default"}, + Spec: &v1alpha1.TaskSpec{Image: fmt.Sprintf("ghcr.io/test/img:%d", rev)}, + }); err != nil { + t.Fatalf("SaveTask: %v", err) + } + } + } + for { + h.server.mu.Lock() + resumes, peak := h.server.resumes, h.server.maxPerTask + h.server.mu.Unlock() + if resumes == tasks*revisions { + if peak != 1 { + t.Errorf("a task had %d concurrent Substrate calls, want 1", peak) + } + break + } + if time.Since(start) > 30*time.Second { + t.Fatalf("only %d of %d reconciles finished", resumes, tasks*revisions) + } + time.Sleep(5 * time.Millisecond) + } + // Serially, 48 reconciles of 6 RPCs at 20ms each take about 5.8s. + if elapsed := time.Since(start); elapsed > 3*time.Second { + t.Errorf("48 reconciles took %v with concurrency 8; tasks do not appear to run in parallel", elapsed) + } +}