Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions .github/workflows/Release.yml
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,18 @@ on:
permissions: write-all # Necessary for the generate-build-provenance action with containers

jobs:
stress-tests:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- uses: actions/setup-go@v5
with:
go-version: stable
- name: Run release stress tests
run: go test -tags=stress -count=1 ./cmd/cloud-init-server ./pkg/wgtunnel ./internal/memstore ./internal/smdclient

release:
needs: stress-tests
uses: OpenCHAMI/github-actions/.github/workflows/go-build-release.yml@v3.2
with:
cgo-enabled: "1"
Expand Down
15 changes: 6 additions & 9 deletions cmd/cloud-init-server/handlers.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,12 +5,11 @@ import (
"io"
"net/http"

// Import to run swag.Register() to generated docs
"github.com/go-chi/chi/v5"
// Import to run swag.Register() to generated docs
_ "github.com/openchami/cloud-init/docs"
"github.com/openchami/cloud-init/internal/smdclient"
"github.com/openchami/cloud-init/pkg/cistore"
"github.com/openchami/cloud-init/pkg/wgtunnel"
"github.com/rs/zerolog/log"
"github.com/swaggo/swag"
)
Expand Down Expand Up @@ -207,7 +206,7 @@ func InstanceInfoHandler(sm smdclient.SMDClientInterface, store cistore.Store) h
// @Param hostname formData string true "Node's given hostname"
// @Param fqdn formData string true "Node's given fully-qualified domain name"
// @Router /phone-home/{id} [post]
func PhoneHomeHandler(wg *wgtunnel.InterfaceManager, sm smdclient.SMDClientInterface) http.HandlerFunc {
func PhoneHomeHandler(peerRemovalQueue *PeerRemovalQueue, sm smdclient.SMDClientInterface) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
w.WriteHeader(http.StatusMethodNotAllowed)
Expand Down Expand Up @@ -249,12 +248,10 @@ func PhoneHomeHandler(wg *wgtunnel.InterfaceManager, sm smdclient.SMDClientInter
Msgf("Received phone home data: pub_key_rsa=%s, pub_key_ecdsa=%s, pub_key_ed25519=%s, instance_id=%s, hostname=%s, fqdn=%s",
pubKeyRsa, pubKeyEcdsa, pubKeyEd25519, instanceId, hostname, fqdn)

if wg != nil {
go func() {
_ = wg.RemovePeer(peerName) // Explicitly ignoring the error here. There's nothing to do with it within the goroutine.
}()

w.WriteHeader(http.StatusOK)
if peerRemovalQueue != nil && !peerRemovalQueue.TryEnqueue(peerName) {
http.Error(w, "WireGuard peer removal queue is full", http.StatusServiceUnavailable)
return
}
w.WriteHeader(http.StatusOK)
}
}
10 changes: 7 additions & 3 deletions cmd/cloud-init-server/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -269,6 +269,10 @@ func startServer() error {

// Create router
router := chi.NewRouter()
var peerRemovalQueue *PeerRemovalQueue
if wgInterfaceManager != nil {
peerRemovalQueue = NewPeerRemovalQueue(wgInterfaceManager)
}

// Add middleware
router.Use(
Expand All @@ -282,7 +286,7 @@ func startServer() error {
)

// Setup routes
initCiClientRouter(router, handler, wgInterfaceManager)
initCiClientRouter(router, handler, wgInterfaceManager, peerRemovalQueue)
initCiAdminRouter(router, handler)

// Add secure routes if JWKS is configured
Expand Down Expand Up @@ -315,7 +319,7 @@ func parseBool(str string) bool {
return strings.EqualFold(str, "true") || str == "1"
}

func initCiClientRouter(router chi.Router, handler *CiHandler, wgInterfaceManager *wgtunnel.InterfaceManager) {
func initCiClientRouter(router chi.Router, handler *CiHandler, wgInterfaceManager *wgtunnel.InterfaceManager, peerRemovalQueue *PeerRemovalQueue) {
// Add cloud-init endpoints to router
router.Get("/openapi.json", DocsHandler)
router.Get("/version", VersionHandler)
Expand All @@ -330,7 +334,7 @@ func initCiClientRouter(router chi.Router, handler *CiHandler, wgInterfaceManage
router.Get("/vendor-data", VendorDataHandler(handler.sm, handler.store, baseUrl))
router.Get("/{group}.yaml", GroupUserDataHandler(handler.sm, handler.store))
}
router.Post("/phone-home/{id}", PhoneHomeHandler(wgInterfaceManager, handler.sm))
router.Post("/phone-home/{id}", PhoneHomeHandler(peerRemovalQueue, handler.sm))
router.Post("/wg-init", wgtunnel.AddClientHandler(wgInterfaceManager, handler.sm))
}

Expand Down
52 changes: 52 additions & 0 deletions cmd/cloud-init-server/peer_removal_queue.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
package main

import "github.com/rs/zerolog/log"

const (
defaultPeerRemovalWorkers = 2
Comment thread
travisbcotton marked this conversation as resolved.
defaultPeerRemovalBuffer = 64
)

type peerRemover interface {
RemovePeer(peerName string) error
}

type PeerRemovalQueue struct {
remover peerRemover
jobs chan string
}

func NewPeerRemovalQueue(remover peerRemover) *PeerRemovalQueue {
return newPeerRemovalQueue(remover, defaultPeerRemovalWorkers, defaultPeerRemovalBuffer)
}

func newPeerRemovalQueue(remover peerRemover, workers int, buffer int) *PeerRemovalQueue {
queue := &PeerRemovalQueue{
remover: remover,
jobs: make(chan string, buffer),
}
for range workers {
go queue.work()
}
return queue
}

func (q *PeerRemovalQueue) TryEnqueue(peerName string) bool {
if q == nil || q.remover == nil {
return true
}
select {
case q.jobs <- peerName:
return true
default:
return false
}
}

func (q *PeerRemovalQueue) work() {
for peerName := range q.jobs {
if err := q.remover.RemovePeer(peerName); err != nil {
log.Error().Err(err).Str("peer", peerName).Msg("failed to remove WireGuard peer")
}
}
}
74 changes: 74 additions & 0 deletions cmd/cloud-init-server/peer_removal_queue_stress_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
//go:build stress

package main

import (
"net/http"
"net/http/httptest"
"sync"
"sync/atomic"
"testing"
"time"
)

func TestStressPhoneHomeQueueBackpressure10K(t *testing.T) {
remover := newBlockingPeerRemover()
queue := newPeerRemovalQueue(remover, defaultPeerRemovalWorkers, defaultPeerRemovalBuffer)
handler := PhoneHomeHandler(queue, &phoneHomeSMDClient{})

for range defaultPeerRemovalWorkers {
recorder := httptest.NewRecorder()
handler(recorder, phoneHomeRequest(t))
if recorder.Code != http.StatusOK {
t.Fatalf("worker-fill response status = %d, want %d", recorder.Code, http.StatusOK)
}
}
for range defaultPeerRemovalWorkers {
select {
case <-remover.started:
case <-time.After(time.Second):
t.Fatal("worker did not start removal")
}
}

const requestCount = 10_000
var okCount atomic.Int64
var unavailableCount atomic.Int64
var ready sync.WaitGroup
var start sync.WaitGroup
var done sync.WaitGroup
ready.Add(requestCount)
start.Add(1)
done.Add(requestCount)

for range requestCount {
go func() {
defer done.Done()
ready.Done()
start.Wait()

recorder := httptest.NewRecorder()
handler(recorder, phoneHomeRequest(t))
switch recorder.Code {
case http.StatusOK:
okCount.Add(1)
case http.StatusServiceUnavailable:
unavailableCount.Add(1)
default:
t.Errorf("response status = %d, want %d or %d", recorder.Code, http.StatusOK, http.StatusServiceUnavailable)
}
}()
}

ready.Wait()
start.Done()
done.Wait()

if got := okCount.Load(); got != defaultPeerRemovalBuffer {
t.Fatalf("accepted removals = %d, want %d", got, defaultPeerRemovalBuffer)
}
if got := unavailableCount.Load(); got != requestCount-defaultPeerRemovalBuffer {
t.Fatalf("backpressured removals = %d, want %d", got, requestCount-defaultPeerRemovalBuffer)
}
close(remover.release)
}
126 changes: 126 additions & 0 deletions cmd/cloud-init-server/peer_removal_queue_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,126 @@
package main

import (
"context"
"net/http"
"net/http/httptest"
"testing"
"time"

"github.com/go-chi/chi/v5"
"github.com/openchami/cloud-init/internal/smdclient"
)

type blockingPeerRemover struct {
started chan struct{}
release chan struct{}
removed chan string
}

func newBlockingPeerRemover() *blockingPeerRemover {
return &blockingPeerRemover{
started: make(chan struct{}, 16),
release: make(chan struct{}),
removed: make(chan string, 16),
}
}

func (r *blockingPeerRemover) RemovePeer(peerName string) error {
r.started <- struct{}{}
<-r.release
r.removed <- peerName
return nil
}

type phoneHomeSMDClient struct {
smdclient.FakeSMDClient
}

func (phoneHomeSMDClient) IDfromIP(string) (string, error) {
return "x0c0s0b0n0", nil
}

func (phoneHomeSMDClient) IPfromID(string) (string, error) {
return "10.1.0.1", nil
}

func TestPeerRemovalQueueBoundsWork(t *testing.T) {
remover := newBlockingPeerRemover()
queue := newPeerRemovalQueue(remover, 1, 1)

if !queue.TryEnqueue("peer-1") {
t.Fatal("first enqueue unexpectedly failed")
}
select {
case <-remover.started:
case <-time.After(time.Second):
t.Fatal("worker did not start first removal")
}
if !queue.TryEnqueue("peer-2") {
t.Fatal("buffered enqueue unexpectedly failed")
}
if queue.TryEnqueue("peer-3") {
t.Fatal("enqueue succeeded when worker and buffer were saturated")
}

close(remover.release)
for range 2 {
select {
case <-remover.removed:
case <-time.After(time.Second):
t.Fatal("queued removal did not finish")
}
}
}

func TestPhoneHomeHandlerReturnsUnavailableWhenRemovalQueueFull(t *testing.T) {
remover := newBlockingPeerRemover()
queue := newPeerRemovalQueue(remover, 1, 1)
handler := PhoneHomeHandler(queue, &phoneHomeSMDClient{})

first := httptest.NewRecorder()
handler(first, phoneHomeRequest(t))
if first.Code != http.StatusOK {
t.Fatalf("first response status = %d, want %d", first.Code, http.StatusOK)
}
select {
case <-remover.started:
case <-time.After(time.Second):
t.Fatal("worker did not start first removal")
}

second := httptest.NewRecorder()
handler(second, phoneHomeRequest(t))
if second.Code != http.StatusOK {
t.Fatalf("second response status = %d, want %d", second.Code, http.StatusOK)
}

third := httptest.NewRecorder()
handler(third, phoneHomeRequest(t))
if third.Code != http.StatusServiceUnavailable {
t.Fatalf("third response status = %d, want %d", third.Code, http.StatusServiceUnavailable)
}
close(remover.release)
}

func TestPhoneHomeHandlerWithoutWireGuardStillReturnsOK(t *testing.T) {
handler := PhoneHomeHandler(nil, &phoneHomeSMDClient{})
recorder := httptest.NewRecorder()
handler(recorder, phoneHomeRequest(t))
if recorder.Code != http.StatusOK {
t.Fatalf("response status = %d, want %d", recorder.Code, http.StatusOK)
}
}

func phoneHomeRequest(t *testing.T) *http.Request {
t.Helper()
r := httptest.NewRequest(http.MethodPost, "/phone-home/x0c0s0b0n0", nil)
r.RemoteAddr = "10.1.0.1:12345"
rctx := chi.NewRouteContext()
rctx.URLParams.Add("id", "x0c0s0b0n0")
return r.WithContext(contextWithRoute(r.Context(), rctx))
}

func contextWithRoute(ctx context.Context, rctx *chi.Context) context.Context {
return context.WithValue(ctx, chi.RouteCtxKey, rctx)
}
Loading
Loading