Files
2026-08-31 00:06:24 +01:00

605 lines
22 KiB
Go

package agenthelper
import (
"bytes"
"context"
"encoding/binary"
"encoding/json"
"errors"
"net"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
)
type fakeSMARTProvider func(context.Context) (json.RawMessage, error)
func (f fakeSMARTProvider) Snapshot(ctx context.Context) (json.RawMessage, error) {
return f(ctx)
}
type fakeProxmoxProvider func(context.Context) (json.RawMessage, error)
func (f fakeProxmoxProvider) LXCFilesystems(ctx context.Context) (json.RawMessage, error) {
return f(ctx)
}
type fakeContainerProvider func(context.Context) (json.RawMessage, error)
func (f fakeContainerProvider) Inventory(ctx context.Context) (json.RawMessage, error) {
return f(ctx)
}
type fakeUpdateProvider struct {
stage func(context.Context, UpdateStageRequest) (UpdateStageResult, error)
activate func(context.Context, UpdateActivateRequest) (UpdateResult, error)
commit func(context.Context, UpdateCommitRequest) (UpdateResult, error)
rollback func(context.Context, UpdateRollbackRequest) (UpdateResult, error)
}
func (f fakeUpdateProvider) Stage(ctx context.Context, request UpdateStageRequest) (UpdateStageResult, error) {
return f.stage(ctx, request)
}
func (f fakeUpdateProvider) Activate(ctx context.Context, request UpdateActivateRequest) (UpdateResult, error) {
return f.activate(ctx, request)
}
func (f fakeUpdateProvider) Commit(ctx context.Context, request UpdateCommitRequest) (UpdateResult, error) {
return f.commit(ctx, request)
}
func (f fakeUpdateProvider) Rollback(ctx context.Context, request UpdateRollbackRequest) (UpdateResult, error) {
return f.rollback(ctx, request)
}
func authorizedResolver(uid uint32) PeerResolver {
return PeerResolverFunc(func(net.Conn) (Peer, error) {
return Peer{UID: uid, GID: 2000, PID: 3000}, nil
})
}
func newTestServer(t *testing.T, registry *Registry, resolver PeerResolver, audit AuditHook) *Server {
t.Helper()
server, err := NewServer(ServerConfig{
AllowedUID: 1000,
PeerResolver: resolver,
Registry: registry,
MaxOperationTimeout: 100 * time.Millisecond,
Audit: audit,
})
if err != nil {
t.Fatalf("NewServer: %v", err)
}
return server
}
func validRequest(operation string) Request {
return Request{
ProtocolVersion: ProtocolVersion,
RequestID: "request-1",
Operation: operation,
OperationVersion: OperationVersion1,
DeadlineMillis: 50,
Payload: json.RawMessage(`{}`),
}
}
func exchangeRequest(t *testing.T, server *Server, framed []byte) Response {
t.Helper()
serverConn, clientConn := net.Pipe()
done := make(chan struct{})
go func() {
server.HandleConnection(context.Background(), serverConn)
close(done)
}()
if _, err := writeAll(clientConn, framed); err != nil {
t.Fatalf("write request: %v", err)
}
payload, err := readFrame(clientConn, MaxResponseBytes)
if err != nil {
t.Fatalf("read response: %v", err)
}
response, err := DecodeResponse(payload)
if err != nil {
t.Fatalf("DecodeResponse: %v", err)
}
_ = clientConn.Close()
<-done
return response
}
func exchange(t *testing.T, server *Server, request Request) Response {
t.Helper()
framed, err := EncodeRequestFrame(request)
if err != nil {
t.Fatalf("EncodeRequestFrame: %v", err)
}
return exchangeRequest(t, server, framed)
}
func requireErrorCode(t *testing.T, response Response, code string) {
t.Helper()
if response.Success || response.Error == nil || response.Error.Code != code {
t.Fatalf("response = %#v, want error %q", response, code)
}
}
func TestServerHealthAndCapabilities(t *testing.T) {
smart := fakeSMARTProvider(func(context.Context) (json.RawMessage, error) {
return json.RawMessage(`{"disks":[]}`), nil
})
server := newTestServer(t, NewRegistry(smart, nil), authorizedResolver(1000), nil)
health := exchange(t, server, validRequest(OperationHealth))
if !health.Success {
t.Fatalf("health response = %#v", health)
}
var healthResult HealthResult
if err := decodeStrict(health.Result, &healthResult); err != nil || healthResult.Status != "ok" {
t.Fatalf("health result = %#v, err=%v", healthResult, err)
}
capabilities := exchange(t, server, validRequest(OperationCapabilities))
if !capabilities.Success {
t.Fatalf("capabilities response = %#v", capabilities)
}
var result CapabilitiesResult
if err := decodeStrict(capabilities.Result, &result); err != nil {
t.Fatalf("decode capabilities: %v", err)
}
availability := make(map[string]bool)
for _, capability := range result.Operations {
availability[capability.Operation] = capability.Available
}
if !availability[OperationSMARTSnapshot] || availability[OperationProxmoxLXCFilesystems] {
t.Fatalf("provider availability = %#v", availability)
}
}
func TestServerDispatchesTypedProvidersWithoutCallerArguments(t *testing.T) {
smartCalled := false
proxmoxCalled := false
registry := NewRegistry(
fakeSMARTProvider(func(context.Context) (json.RawMessage, error) {
smartCalled = true
return json.RawMessage(`{"disks":[{"device":"sda"}]}`), nil
}),
fakeProxmoxProvider(func(context.Context) (json.RawMessage, error) {
proxmoxCalled = true
return json.RawMessage(`{"containers":[]}`), nil
}),
)
server := newTestServer(t, registry, authorizedResolver(1000), nil)
if response := exchange(t, server, validRequest(OperationSMARTSnapshot)); !response.Success || !smartCalled {
t.Fatalf("SMART response = %#v called=%t", response, smartCalled)
}
if response := exchange(t, server, validRequest(OperationProxmoxLXCFilesystems)); !response.Success || !proxmoxCalled {
t.Fatalf("Proxmox response = %#v called=%t", response, proxmoxCalled)
}
request := validRequest(OperationSMARTSnapshot)
request.Payload = json.RawMessage(`{"path":"/dev/sda"}`)
requireErrorCode(t, exchange(t, server, request), ErrorInvalidRequest)
request = validRequest(OperationProxmoxLXCFilesystems)
request.Payload = json.RawMessage(`{"args":["exec","100"]}`)
requireErrorCode(t, exchange(t, server, request), ErrorInvalidRequest)
}
func TestServerDispatchesClosedContainerAndUpdateOperations(t *testing.T) {
containerCalled := false
stageCalled := false
activateCalled := false
rollbackCalled := false
commitCalled := false
updates := fakeUpdateProvider{
stage: func(_ context.Context, request UpdateStageRequest) (UpdateStageResult, error) {
stageCalled = request.ArtifactID == "release-1"
return UpdateStageResult{Action: "staged", ArtifactID: request.ArtifactID, SHA256: request.SHA256}, nil
},
activate: func(_ context.Context, request UpdateActivateRequest) (UpdateResult, error) {
activateCalled = request.ArtifactID == "release-1"
return UpdateResult{Action: "pending", ActivationID: "release-1:0123456789abcdef"}, nil
},
commit: func(_ context.Context, request UpdateCommitRequest) (UpdateResult, error) {
commitCalled = request.ActivationID == "release-1:0123456789abcdef"
return UpdateResult{Action: "committed", ActivationID: request.ActivationID}, nil
},
rollback: func(_ context.Context, request UpdateRollbackRequest) (UpdateResult, error) {
rollbackCalled = request.ActivationID == "release-1:0123456789abcdef"
return UpdateResult{Action: "rolled_back", ActivationID: request.ActivationID}, nil
},
}
registry := NewRegistryWithProviders(nil, nil, Providers{
Containers: fakeContainerProvider(func(context.Context) (json.RawMessage, error) {
containerCalled = true
return json.RawMessage(`{"runtimes":[]}`), nil
}),
Updates: updates,
})
server := newTestServer(t, registry, authorizedResolver(1000), nil)
if response := exchange(t, server, validRequest(OperationContainerInventory)); !response.Success || !containerCalled {
t.Fatalf("container response=%#v called=%t", response, containerCalled)
}
stage := validRequest(OperationAgentUpdateStage)
stage.Payload = json.RawMessage(`{"artifactId":"release-1","sha256":"0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef"}`)
if response := exchange(t, server, stage); !response.Success || !stageCalled {
t.Fatalf("stage response=%#v called=%t", response, stageCalled)
}
activate := validRequest(OperationAgentUpdateActivate)
activate.Payload = json.RawMessage(`{"artifactId":"release-1","sha256":"0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef"}`)
if response := exchange(t, server, activate); !response.Success || !activateCalled {
t.Fatalf("activate response=%#v called=%t", response, activateCalled)
}
commit := validRequest(OperationAgentUpdateCommit)
commit.Payload = json.RawMessage(`{"activationId":"release-1:0123456789abcdef","currentSha256":"0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef"}`)
if response := exchange(t, server, commit); !response.Success || !commitCalled {
t.Fatalf("commit response=%#v called=%t", response, commitCalled)
}
rollback := validRequest(OperationAgentUpdateRollback)
rollback.Payload = json.RawMessage(`{"activationId":"release-1:0123456789abcdef","currentSha256":"0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef","rollbackSha256":"abcdef0123456789abcdef0123456789abcdef0123456789abcdef0123456789"}`)
if response := exchange(t, server, rollback); !response.Success || !rollbackCalled {
t.Fatalf("rollback response=%#v called=%t", response, rollbackCalled)
}
for _, operation := range []string{OperationContainerInventory, OperationAgentUpdateStage, OperationAgentUpdateActivate, OperationAgentUpdateCommit, OperationAgentUpdateRollback} {
request := validRequest(operation)
request.Payload = json.RawMessage(`{"path":"/tmp/attacker","args":["sh"]}`)
requireErrorCode(t, exchange(t, server, request), ErrorInvalidRequest)
request = validRequest(operation)
request.OperationVersion = 2
requireErrorCode(t, exchange(t, server, request), ErrorUnsupportedOperation)
}
}
func TestServerRejectsEnvelopeAndRegistryViolations(t *testing.T) {
server := newTestServer(t, NewRegistry(nil, nil), authorizedResolver(1000), nil)
tests := []struct {
name string
mutate func(*Request)
code string
}{
{name: "protocol", mutate: func(r *Request) { r.ProtocolVersion = 2 }, code: ErrorUnsupportedProtocol},
{name: "unknown operation", mutate: func(r *Request) { r.Operation = "host.exec" }, code: ErrorUnknownOperation},
{name: "operation version", mutate: func(r *Request) { r.OperationVersion = 2 }, code: ErrorUnsupportedOperation},
{name: "missing request id", mutate: func(r *Request) { r.RequestID = "" }, code: ErrorInvalidRequest},
{name: "unsafe request id", mutate: func(r *Request) { r.RequestID = "request\nforged" }, code: ErrorInvalidRequest},
{name: "zero deadline", mutate: func(r *Request) { r.DeadlineMillis = 0 }, code: ErrorInvalidRequest},
{name: "excessive deadline", mutate: func(r *Request) { r.DeadlineMillis = 101 }, code: ErrorInvalidRequest},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
request := validRequest(OperationHealth)
test.mutate(&request)
requireErrorCode(t, exchange(t, server, request), test.code)
})
}
}
func TestServerRejectsStrictJSONViolations(t *testing.T) {
server := newTestServer(t, NewRegistry(nil, nil), authorizedResolver(1000), nil)
for _, raw := range []string{
`{"protocolVersion":1,"requestId":"r","operation":"helper.health","operationVersion":1,"deadlineMillis":50,"unexpected":true}`,
`{"protocolVersion":1,"requestId":"r","operation":"helper.health","operationVersion":1,"deadlineMillis":50}{}`,
} {
framed := make([]byte, 4+len(raw))
binary.BigEndian.PutUint32(framed[:4], uint32(len(raw)))
copy(framed[4:], raw)
requireErrorCode(t, exchangeRequest(t, server, framed), ErrorInvalidRequest)
}
}
func TestServerRejectsOversizedFrameBeforeReadingPayload(t *testing.T) {
server := newTestServer(t, NewRegistry(nil, nil), authorizedResolver(1000), nil)
framed := make([]byte, 4)
binary.BigEndian.PutUint32(framed, MaxRequestBytes+1)
requireErrorCode(t, exchangeRequest(t, server, framed), ErrorInvalidFrame)
}
func TestServerRejectsUnauthorizedPeer(t *testing.T) {
server := newTestServer(t, NewRegistry(nil, nil), authorizedResolver(1001), nil)
serverConn, clientConn := net.Pipe()
done := make(chan struct{})
go func() {
server.HandleConnection(context.Background(), serverConn)
close(done)
}()
payload, err := readFrame(clientConn, MaxResponseBytes)
if err != nil {
t.Fatalf("read unauthorized response: %v", err)
}
response, err := DecodeResponse(payload)
if err != nil {
t.Fatalf("decode unauthorized response: %v", err)
}
requireErrorCode(t, response, ErrorUnauthorizedPeer)
_ = clientConn.Close()
<-done
}
func TestServerBoundsConcurrentConnectionsAndRecovers(t *testing.T) {
const maxConnections = 2
resolved := make(chan struct{}, maxConnections+1)
audited := make(chan AuditEvent, maxConnections+1)
server, err := NewServer(ServerConfig{
AllowedUID: 1000,
PeerResolver: PeerResolverFunc(func(net.Conn) (Peer, error) {
resolved <- struct{}{}
return Peer{UID: 1000, GID: 2000, PID: 3000}, nil
}),
Registry: NewRegistry(nil, nil),
MaxConcurrentConnections: maxConnections,
MaxOperationTimeout: 5 * time.Second,
PreFrameTimeout: 5 * time.Second,
Audit: func(event AuditEvent) { audited <- event },
})
if err != nil {
t.Fatalf("NewServer: %v", err)
}
socketDir, err := os.MkdirTemp("", "pah-")
if err != nil {
t.Fatalf("create short socket directory: %v", err)
}
t.Cleanup(func() { _ = os.RemoveAll(socketDir) })
socketPath := filepath.Join(socketDir, "helper.sock")
listener, err := net.Listen("unix", socketPath)
if err != nil {
t.Fatalf("listen: %v", err)
}
ctx, cancel := context.WithCancel(context.Background())
serveDone := make(chan error, 1)
go func() { serveDone <- server.Serve(ctx, listener) }()
t.Cleanup(func() {
cancel()
_ = listener.Close()
if err := <-serveDone; err != nil {
t.Errorf("Serve: %v", err)
}
})
held := make([]net.Conn, 0, maxConnections)
for i := 0; i < maxConnections; i++ {
conn, dialErr := net.Dial("unix", socketPath)
if dialErr != nil {
t.Fatalf("dial held connection %d: %v", i, dialErr)
}
held = append(held, conn)
select {
case <-resolved:
case <-time.After(time.Second):
t.Fatalf("held connection %d was not admitted", i)
}
}
t.Cleanup(func() {
for _, conn := range held {
_ = conn.Close()
}
})
overload, err := net.Dial("unix", socketPath)
if err != nil {
t.Fatalf("dial overloaded connection: %v", err)
}
defer overload.Close()
if err := overload.SetReadDeadline(time.Now().Add(time.Second)); err != nil {
t.Fatalf("set overload deadline: %v", err)
}
buffer := make([]byte, 1)
if _, err := overload.Read(buffer); err == nil {
t.Fatal("overloaded connection remained open")
} else if netErr, ok := err.(net.Error); ok && netErr.Timeout() {
t.Fatal("overloaded connection was not rejected promptly")
}
select {
case <-resolved:
t.Fatal("overloaded connection reached peer authentication")
default:
}
if err := held[0].Close(); err != nil {
t.Fatalf("release held connection: %v", err)
}
select {
case event := <-audited:
if event.Success || event.ErrorCode != ErrorInvalidFrame {
t.Fatalf("released connection audit = %#v", event)
}
case <-time.After(time.Second):
t.Fatal("released connection was not cleaned up")
}
deadline := time.Now().Add(time.Second)
for len(server.connectionSlots) != maxConnections-1 && time.Now().Before(deadline) {
time.Sleep(time.Millisecond)
}
if got := len(server.connectionSlots); got != maxConnections-1 {
t.Fatalf("active connection slots = %d, want %d", got, maxConnections-1)
}
client, err := NewClient(ClientConfig{
SocketPath: socketPath,
MaxDeadline: time.Second,
NewRequestID: func() (string, error) { return "post-overload-health", nil },
})
if err != nil {
t.Fatalf("NewClient: %v", err)
}
var health HealthResult
if _, err := client.Call(t.Context(), OperationHealth, OperationVersion1, time.Second, nil, &health); err != nil {
t.Fatalf("health after releasing connection slot: %v", err)
}
if health.Status != "ok" {
t.Fatalf("health status = %q, want ok", health.Status)
}
}
func TestServerAppliesShortPreFrameTimeout(t *testing.T) {
audited := make(chan AuditEvent, 1)
server, err := NewServer(ServerConfig{
AllowedUID: 1000,
PeerResolver: authorizedResolver(1000),
Registry: NewRegistry(nil, nil),
MaxOperationTimeout: time.Second,
PreFrameTimeout: 25 * time.Millisecond,
Audit: func(event AuditEvent) { audited <- event },
})
if err != nil {
t.Fatalf("NewServer: %v", err)
}
serverConn, clientConn := net.Pipe()
done := make(chan struct{})
started := time.Now()
go func() {
server.HandleConnection(context.Background(), serverConn)
close(done)
}()
payload, err := readFrame(clientConn, MaxResponseBytes)
if err != nil {
t.Fatalf("read timeout response: %v", err)
}
response, err := DecodeResponse(payload)
if err != nil {
t.Fatalf("decode timeout response: %v", err)
}
requireErrorCode(t, response, ErrorInvalidFrame)
_ = clientConn.Close()
<-done
if elapsed := time.Since(started); elapsed >= 500*time.Millisecond {
t.Fatalf("pre-frame timeout took %s, want less than 500ms", elapsed)
}
event := <-audited
if event.RequestID != "" || event.Operation != "" || event.RequestBytes != 0 || event.ErrorCode != ErrorInvalidFrame {
t.Fatalf("pre-frame audit exposed unexpected metadata: %#v", event)
}
}
func TestNewServerValidatesConnectionLimitsAndTimeouts(t *testing.T) {
base := ServerConfig{
AllowedUID: 1000,
PeerResolver: authorizedResolver(1000),
Registry: NewRegistry(nil, nil),
PreFrameTimeout: time.Second,
}
tests := []struct {
name string
mutate func(*ServerConfig)
}{
{name: "negative connection limit", mutate: func(config *ServerConfig) { config.MaxConcurrentConnections = -1 }},
{name: "negative operation timeout", mutate: func(config *ServerConfig) { config.MaxOperationTimeout = -1 }},
{name: "sub-millisecond operation timeout", mutate: func(config *ServerConfig) { config.MaxOperationTimeout = time.Nanosecond }},
{name: "negative pre-frame timeout", mutate: func(config *ServerConfig) { config.PreFrameTimeout = -1 }},
{name: "sub-millisecond pre-frame timeout", mutate: func(config *ServerConfig) { config.PreFrameTimeout = time.Nanosecond }},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
config := base
test.mutate(&config)
if _, err := NewServer(config); err == nil {
t.Fatal("NewServer succeeded with invalid configuration")
}
})
}
server, err := NewServer(ServerConfig{
AllowedUID: 1000,
PeerResolver: authorizedResolver(1000),
Registry: NewRegistry(nil, nil),
})
if err != nil {
t.Fatalf("NewServer defaults: %v", err)
}
if cap(server.connectionSlots) != defaultMaxConcurrentConnections {
t.Fatalf("default connection limit = %d, want %d", cap(server.connectionSlots), defaultMaxConcurrentConnections)
}
if server.preFrameTimeout != defaultPreFrameTimeout {
t.Fatalf("default pre-frame timeout = %s, want %s", server.preFrameTimeout, defaultPreFrameTimeout)
}
}
func TestServerBoundsOperationDeadline(t *testing.T) {
provider := fakeSMARTProvider(func(ctx context.Context) (json.RawMessage, error) {
<-ctx.Done()
return nil, ctx.Err()
})
server := newTestServer(t, NewRegistry(provider, nil), authorizedResolver(1000), nil)
request := validRequest(OperationSMARTSnapshot)
request.DeadlineMillis = 5
response := exchange(t, server, request)
requireErrorCode(t, response, ErrorDeadlineExceeded)
if !response.Error.Retryable {
t.Fatal("deadline error must be retryable")
}
}
func TestServerReplacesOversizedResponseWithTypedError(t *testing.T) {
oversized := json.RawMessage(`{"data":"` + strings.Repeat("x", int(MaxResponseBytes)) + `"}`)
server := newTestServer(t, NewRegistry(fakeSMARTProvider(func(context.Context) (json.RawMessage, error) {
return oversized, nil
}), nil), authorizedResolver(1000), nil)
requireErrorCode(t, exchange(t, server, validRequest(OperationSMARTSnapshot)), ErrorResponseTooLarge)
}
func TestServerMapsProviderErrorsAndInvalidResults(t *testing.T) {
tests := []struct {
name string
provider fakeSMARTProvider
code string
}{
{name: "typed", provider: func(context.Context) (json.RawMessage, error) {
return nil, &ProviderError{Code: ErrorProviderUnavailable, Message: "not configured"}
}, code: ErrorProviderUnavailable},
{name: "untyped", provider: func(context.Context) (json.RawMessage, error) {
return nil, errors.New("secret provider detail")
}, code: ErrorInternal},
{name: "invalid JSON", provider: func(context.Context) (json.RawMessage, error) {
return json.RawMessage(`not-json`), nil
}, code: ErrorInternal},
{name: "unstable provider code", provider: func(context.Context) (json.RawMessage, error) {
return nil, &ProviderError{Code: "invented_provider_code", Message: "bounded"}
}, code: ErrorInternal},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
server := newTestServer(t, NewRegistry(test.provider, nil), authorizedResolver(1000), nil)
response := exchange(t, server, validRequest(OperationSMARTSnapshot))
requireErrorCode(t, response, test.code)
if strings.Contains(response.Error.Message, "secret provider detail") {
t.Fatal("untyped provider detail escaped helper boundary")
}
})
}
}
func TestAuditHookContainsMetadataOnly(t *testing.T) {
var (
mu sync.Mutex
events []AuditEvent
)
server := newTestServer(t, NewRegistry(fakeSMARTProvider(func(context.Context) (json.RawMessage, error) {
return json.RawMessage(`{"secret":"must-not-be-audited"}`), nil
}), nil), authorizedResolver(1000), func(event AuditEvent) {
mu.Lock()
defer mu.Unlock()
events = append(events, event)
})
response := exchange(t, server, validRequest(OperationSMARTSnapshot))
if !response.Success {
t.Fatalf("response = %#v", response)
}
mu.Lock()
defer mu.Unlock()
if len(events) != 1 || events[0].RequestID != "request-1" || !events[0].Success {
t.Fatalf("audit events = %#v", events)
}
encoded, err := json.Marshal(events[0])
if err != nil {
t.Fatal(err)
}
if bytes.Contains(encoded, []byte("must-not-be-audited")) {
t.Fatal("result data reached audit metadata")
}
}