iop/apps/agent/internal/localcontrol/server_test.go

459 lines
11 KiB
Go

package localcontrol
import (
"context"
"encoding/binary"
"errors"
"io"
"net"
"os"
"path/filepath"
"testing"
"time"
proto_socket "git.toki-labs.com/toki/proto-socket/go"
"git.toki-labs.com/toki/proto-socket/go/packets"
iop "iop/proto/gen/iop"
"google.golang.org/protobuf/proto"
)
type fixedPeerCredentials struct {
supported bool
uid uint32
err error
}
func (f fixedPeerCredentials) Supported() bool {
return f.supported
}
func (f fixedPeerCredentials) UID(net.Conn) (uint32, error) {
return f.uid, f.err
}
func TestServerSameUserProtoSocket(t *testing.T) {
t.Parallel()
host := newRecordingHost()
stateRoot := filepath.Join(t.TempDir(), "agent-state")
service, _ := newTestService(
t,
filepath.Join(stateRoot, "manager.json"),
"daemon-server",
8,
host,
)
server, err := NewServer(ServerConfig{StateRoot: stateRoot}, service)
if err != nil {
t.Fatalf("NewServer: %v", err)
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
if err := server.Start(ctx); err != nil {
t.Fatalf("Start: %v", err)
}
defer func() {
if err := server.Stop(); err != nil {
t.Errorf("Stop: %v", err)
}
}()
rootInfo, err := os.Lstat(stateRoot)
if err != nil {
t.Fatalf("Lstat state root: %v", err)
}
if rootInfo.Mode().Perm() != 0o700 {
t.Fatalf("state root mode = %o, want 700", rootInfo.Mode().Perm())
}
socketInfo, err := os.Lstat(server.Path())
if err != nil {
t.Fatalf("Lstat socket: %v", err)
}
if socketInfo.Mode()&os.ModeSocket == 0 || socketInfo.Mode().Perm() != 0o600 {
t.Fatalf("socket mode = %v, want socket 0600", socketInfo.Mode())
}
conn, err := net.DialTimeout("unix", server.Path(), time.Second)
if err != nil {
t.Fatalf("Dial: %v", err)
}
client := proto_socket.NewTcpClient(
conn,
0,
0,
localControlParserMap(),
)
defer client.Close()
response, err := proto_socket.SendRequestTyped[
*iop.AgentLocalEnvelope,
*iop.AgentLocalEnvelope,
](
&client.Communicator,
runtimeStatusRequest("message-real-socket"),
2*time.Second,
)
if err != nil {
t.Fatalf("SendRequestTyped: %v", err)
}
if response.GetResponse().GetSnapshot().GetDaemonId() != "daemon-server" {
t.Fatalf("response = %#v", response)
}
if calls := host.callCount(OperationRuntimeStatus); calls != 1 {
t.Fatalf("runtime.status calls = %d, want 1", calls)
}
}
func TestPeerUIDMismatchDeniedBeforeDispatch(t *testing.T) {
t.Parallel()
host := newRecordingHost()
stateRoot := filepath.Join(t.TempDir(), "agent-state")
service, _ := newTestService(
t,
filepath.Join(stateRoot, "manager.json"),
"daemon-denied",
8,
host,
)
server, err := newServer(
ServerConfig{StateRoot: stateRoot},
service,
fixedPeerCredentials{
supported: true,
uid: effectiveUID() + 1,
},
effectiveUID,
)
if err != nil {
t.Fatalf("newServer: %v", err)
}
if err := server.Start(context.Background()); err != nil {
t.Fatalf("Start: %v", err)
}
defer server.Stop()
conn, err := net.DialTimeout("unix", server.Path(), time.Second)
if err != nil {
t.Fatalf("Dial: %v", err)
}
defer conn.Close()
response := readEnvelopePacket(t, conn)
if response.GetError().GetCode() != ErrorPermissionDenied {
t.Fatalf("denied response = %#v", response)
}
if calls := host.totalCalls(); calls != 0 {
t.Fatalf("denied peer host calls = %d, want 0", calls)
}
}
func TestServerBroadcastsCommittedEventToConcurrentClients(t *testing.T) {
t.Parallel()
host := newRecordingHost()
stateRoot := filepath.Join(t.TempDir(), "agent-state")
service, _ := newTestService(
t,
filepath.Join(stateRoot, "manager.json"),
"daemon-broadcast",
8,
host,
)
server, err := NewServer(ServerConfig{StateRoot: stateRoot}, service)
if err != nil {
t.Fatalf("NewServer: %v", err)
}
if err := server.Start(context.Background()); err != nil {
t.Fatalf("Start: %v", err)
}
defer server.Stop()
clients := make([]*proto_socket.TcpClient, 0, 2)
eventChannels := make([]chan *iop.AgentLocalEnvelope, 0, 2)
for index := 0; index < 2; index++ {
conn, err := net.DialTimeout("unix", server.Path(), time.Second)
if err != nil {
t.Fatalf("Dial client %d: %v", index, err)
}
client := proto_socket.NewTcpClient(conn, 0, 0, localControlParserMap())
events := make(chan *iop.AgentLocalEnvelope, 1)
proto_socket.AddListenerTyped(
&client.Communicator,
func(envelope *iop.AgentLocalEnvelope) {
if envelope.GetKind() == iop.AgentLocalKind_AGENT_LOCAL_KIND_EVENT {
events <- envelope
}
},
)
if _, err := proto_socket.SendRequestTyped[
*iop.AgentLocalEnvelope,
*iop.AgentLocalEnvelope,
](
&client.Communicator,
runtimeStatusRequest("message-client-ready-"+string(rune('0'+index))),
2*time.Second,
); err != nil {
t.Fatalf("client %d readiness request: %v", index, err)
}
clients = append(clients, client)
eventChannels = append(eventChannels, events)
}
defer func() {
for _, client := range clients {
_ = client.Close()
}
}()
response, err := proto_socket.SendRequestTyped[
*iop.AgentLocalEnvelope,
*iop.AgentLocalEnvelope,
](
&clients[0].Communicator,
projectRequest(
"message-broadcast",
"command-broadcast",
OperationProjectStart,
"project-a",
),
2*time.Second,
)
if err != nil {
t.Fatalf("SendRequestTyped: %v", err)
}
if response.GetResponse().GetMutation() == nil {
t.Fatalf("response = %#v", response)
}
for index, events := range eventChannels {
select {
case event := <-events:
if event.GetEvent().GetEventSequence() != 1 {
t.Fatalf("client %d event = %#v", index, event)
}
case <-time.After(2 * time.Second):
t.Fatalf("client %d did not receive live event", index)
}
}
}
func TestServerRejectsUnsafePaths(t *testing.T) {
t.Parallel()
host := newRecordingHost()
tests := []struct {
name string
prepare func(*testing.T, string)
want error
}{
{
name: "broad state root",
prepare: func(t *testing.T, root string) {
t.Helper()
if err := os.MkdirAll(root, 0o755); err != nil {
t.Fatalf("MkdirAll: %v", err)
}
if err := os.Chmod(root, 0o755); err != nil {
t.Fatalf("Chmod: %v", err)
}
},
want: ErrUnsafeStateRoot,
},
{
name: "state root symlink",
prepare: func(t *testing.T, root string) {
t.Helper()
target := root + "-target"
if err := os.MkdirAll(target, 0o700); err != nil {
t.Fatalf("MkdirAll target: %v", err)
}
if err := os.Symlink(target, root); err != nil {
t.Fatalf("Symlink: %v", err)
}
},
want: ErrUnsafeStateRoot,
},
{
name: "existing socket path replacement",
prepare: func(t *testing.T, root string) {
t.Helper()
if err := os.MkdirAll(root, 0o700); err != nil {
t.Fatalf("MkdirAll: %v", err)
}
if err := os.WriteFile(
filepath.Join(root, defaultSocketName),
[]byte("replacement"),
0o600,
); err != nil {
t.Fatalf("WriteFile: %v", err)
}
},
want: ErrUnsafeSocketPath,
},
}
for _, test := range tests {
test := test
t.Run(test.name, func(t *testing.T) {
t.Parallel()
root := filepath.Join(t.TempDir(), "state")
test.prepare(t, root)
service, _ := newTestService(
t,
filepath.Join(t.TempDir(), "manager.json"),
"daemon-path-"+test.name,
8,
host,
)
server, err := NewServer(ServerConfig{StateRoot: root}, service)
if err != nil {
t.Fatalf("NewServer: %v", err)
}
if err := server.Start(context.Background()); !errors.Is(err, test.want) {
t.Fatalf("Start error = %v, want %v", err, test.want)
}
})
}
}
func TestServerStopPreservesReplacedSocketPath(t *testing.T) {
t.Parallel()
host := newRecordingHost()
stateRoot := filepath.Join(t.TempDir(), "state")
service, _ := newTestService(
t,
filepath.Join(t.TempDir(), "manager.json"),
"daemon-replacement",
8,
host,
)
server, err := NewServer(ServerConfig{StateRoot: stateRoot}, service)
if err != nil {
t.Fatalf("NewServer: %v", err)
}
if err := server.Start(context.Background()); err != nil {
t.Fatalf("Start: %v", err)
}
if err := os.Remove(server.Path()); err != nil {
t.Fatalf("Remove socket: %v", err)
}
replacement := []byte("do-not-remove")
if err := os.WriteFile(server.Path(), replacement, 0o600); err != nil {
t.Fatalf("WriteFile replacement: %v", err)
}
if err := server.Stop(); !errors.Is(err, ErrUnsafeSocketPath) {
t.Fatalf("Stop error = %v, want ErrUnsafeSocketPath", err)
}
got, err := os.ReadFile(server.Path())
if err != nil {
t.Fatalf("ReadFile replacement: %v", err)
}
if string(got) != string(replacement) {
t.Fatalf("replacement content = %q", got)
}
}
func TestServerFailsBeforeListeningWithoutPeerCredentials(t *testing.T) {
t.Parallel()
host := newRecordingHost()
stateRoot := filepath.Join(t.TempDir(), "state")
service, _ := newTestService(
t,
filepath.Join(t.TempDir(), "manager.json"),
"daemon-unsupported",
8,
host,
)
server, err := newServer(
ServerConfig{StateRoot: stateRoot},
service,
fixedPeerCredentials{supported: false},
effectiveUID,
)
if err != nil {
t.Fatalf("newServer: %v", err)
}
if err := server.Start(context.Background()); !errors.Is(
err,
ErrPeerCredentialsUnsupported,
) {
t.Fatalf("Start error = %v", err)
}
if _, err := os.Lstat(stateRoot); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("state root exists after unsupported preflight: %v", err)
}
}
func TestServerOversizeFrameHasZeroMutation(t *testing.T) {
t.Parallel()
host := newRecordingHost()
stateRoot := filepath.Join(t.TempDir(), "state")
service, _ := newTestService(
t,
filepath.Join(t.TempDir(), "manager.json"),
"daemon-frame-bound",
8,
host,
)
server, err := NewServer(ServerConfig{StateRoot: stateRoot}, service)
if err != nil {
t.Fatalf("NewServer: %v", err)
}
if err := server.Start(context.Background()); err != nil {
t.Fatalf("Start: %v", err)
}
defer server.Stop()
conn, err := net.DialTimeout("unix", server.Path(), time.Second)
if err != nil {
t.Fatalf("Dial: %v", err)
}
defer conn.Close()
header := make([]byte, 4)
binary.BigEndian.PutUint32(header, MaxFrameBytes+1)
if _, err := conn.Write(header); err != nil {
t.Fatalf("Write header: %v", err)
}
if err := conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil {
t.Fatalf("SetReadDeadline: %v", err)
}
var one [1]byte
if _, err := conn.Read(one[:]); err == nil {
t.Fatal("oversize frame connection remained open")
}
if calls := host.totalCalls(); calls != 0 {
t.Fatalf("oversize frame host calls = %d, want 0", calls)
}
}
func localControlParserMap() proto_socket.ParserMap {
return proto_socket.ParserMap{
proto_socket.TypeNameOf(&iop.AgentLocalEnvelope{}): func(
payload []byte,
) (proto.Message, error) {
return ParseEnvelope(payload)
},
}
}
func readEnvelopePacket(t *testing.T, conn net.Conn) *iop.AgentLocalEnvelope {
t.Helper()
if err := conn.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil {
t.Fatalf("SetReadDeadline: %v", err)
}
header := make([]byte, 4)
if _, err := io.ReadFull(conn, header); err != nil {
t.Fatalf("ReadFull header: %v", err)
}
length := binary.BigEndian.Uint32(header)
if length == 0 || length > MaxFrameBytes {
t.Fatalf("packet length = %d", length)
}
payload := make([]byte, int(length))
if _, err := io.ReadFull(conn, payload); err != nil {
t.Fatalf("ReadFull payload: %v", err)
}
packet := &packets.PacketBase{}
if err := proto.Unmarshal(payload, packet); err != nil {
t.Fatalf("Unmarshal packet: %v", err)
}
envelope, err := ParseEnvelope(packet.GetData())
if err != nil {
t.Fatalf("ParseEnvelope: %v", err)
}
return envelope
}