mirror of
https://github.com/rcourtman/Pulse.git
synced 2026-09-10 18:45:53 +00:00
605 lines
22 KiB
Go
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")
|
|
}
|
|
}
|