459 lines
11 KiB
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
|
|
}
|