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

546 lines
12 KiB
Go

package localcontrol
import (
"context"
"encoding/binary"
"errors"
"fmt"
"io"
"net"
"os"
"path/filepath"
"sync"
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"
)
const defaultSocketName = "iop-agent.sock"
var (
ErrServerStarted = errors.New("localcontrol: server is already started")
ErrUnsafeStateRoot = errors.New("localcontrol: state root is not owner-only")
ErrUnsafeSocketPath = errors.New("localcontrol: socket path is unsafe")
)
type ServerConfig struct {
StateRoot string
SocketName string
}
// Server owns one Unix local proto-socket and authorizes every accepted peer
// from kernel credentials before constructing a protocol session.
type Server struct {
config ServerConfig
service *Service
credentials peerCredentialSource
effectiveUID func() uint32
mu sync.Mutex
listener *net.UnixListener
socketInfo os.FileInfo
sessions map[*framedSession]struct{}
started bool
stopping bool
acceptDone chan struct{}
eventCancel func()
eventDone chan struct{}
sessionWait sync.WaitGroup
}
func NewServer(config ServerConfig, service *Service) (*Server, error) {
return newServer(
config,
service,
kernelPeerCredentialSource{},
effectiveUID,
)
}
func newServer(
config ServerConfig,
service *Service,
credentials peerCredentialSource,
euid func() uint32,
) (*Server, error) {
if service == nil {
return nil, fmt.Errorf("localcontrol: service is required")
}
if credentials == nil || euid == nil {
return nil, fmt.Errorf("localcontrol: peer credential boundary is required")
}
if !filepath.IsAbs(config.StateRoot) {
return nil, fmt.Errorf("localcontrol: state root must be absolute")
}
config.StateRoot = filepath.Clean(config.StateRoot)
if config.SocketName == "" {
config.SocketName = defaultSocketName
}
if filepath.Base(config.SocketName) != config.SocketName ||
config.SocketName == "." ||
config.SocketName == string(filepath.Separator) {
return nil, fmt.Errorf("localcontrol: socket name must be one file name")
}
return &Server{
config: config,
service: service,
credentials: credentials,
effectiveUID: euid,
sessions: make(map[*framedSession]struct{}),
}, nil
}
func (s *Server) Path() string {
return filepath.Join(s.config.StateRoot, s.config.SocketName)
}
func (s *Server) Start(ctx context.Context) error {
if ctx == nil {
return fmt.Errorf("localcontrol: context is required")
}
if !s.credentials.Supported() {
return ErrPeerCredentialsUnsupported
}
s.mu.Lock()
defer s.mu.Unlock()
if s.started || s.stopping {
return ErrServerStarted
}
if err := s.prepareStateRoot(); err != nil {
return err
}
path := s.Path()
if _, err := os.Lstat(path); err == nil {
return fmt.Errorf("%w: socket path already exists", ErrUnsafeSocketPath)
} else if !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("%w: inspect socket path", ErrUnsafeSocketPath)
}
address, err := net.ResolveUnixAddr("unix", path)
if err != nil {
return fmt.Errorf("localcontrol: resolve socket address: %w", err)
}
listener, err := net.ListenUnix("unix", address)
if err != nil {
return fmt.Errorf("localcontrol: listen on Unix socket: %w", err)
}
listener.SetUnlinkOnClose(false)
createdInfo, err := os.Lstat(path)
if err != nil {
_ = listener.Close()
return fmt.Errorf("localcontrol: inspect created socket: %w", err)
}
cleanup := true
defer func() {
if cleanup {
_ = listener.Close()
current, err := os.Lstat(path)
if err == nil &&
current.Mode()&os.ModeSymlink == 0 &&
current.Mode()&os.ModeSocket != 0 &&
os.SameFile(createdInfo, current) {
_ = os.Remove(path)
}
}
}()
if err := os.Chmod(path, 0o600); err != nil {
return fmt.Errorf("localcontrol: set socket permissions: %w", err)
}
info, err := os.Lstat(path)
if err != nil {
return fmt.Errorf("localcontrol: inspect created socket: %w", err)
}
if info.Mode()&os.ModeSymlink != 0 ||
info.Mode()&os.ModeSocket == 0 ||
info.Mode().Perm() != 0o600 {
return fmt.Errorf("%w: created socket metadata is invalid", ErrUnsafeSocketPath)
}
owner, err := fileOwnerUID(info)
if err != nil || owner != s.effectiveUID() {
return fmt.Errorf("%w: created socket owner is invalid", ErrUnsafeSocketPath)
}
s.listener = listener
s.socketInfo = info
s.acceptDone = make(chan struct{})
events, cancelEvents := s.service.ledger.Subscribe()
s.eventCancel = cancelEvents
s.eventDone = make(chan struct{})
s.started = true
s.stopping = false
cleanup = false
go s.acceptLoop(listener, s.acceptDone)
go s.broadcastLoop(events, s.eventDone)
if ctx.Done() != nil {
acceptDone := s.acceptDone
go func() {
select {
case <-ctx.Done():
_ = s.Stop()
case <-acceptDone:
}
}()
}
return nil
}
func (s *Server) Stop() error {
s.mu.Lock()
if !s.started && !s.stopping {
s.mu.Unlock()
return nil
}
if s.stopping {
done := s.acceptDone
eventDone := s.eventDone
s.mu.Unlock()
if done != nil {
<-done
}
if eventDone != nil {
<-eventDone
}
s.sessionWait.Wait()
return nil
}
s.stopping = true
s.started = false
listener := s.listener
done := s.acceptDone
cancelEvents := s.eventCancel
eventDone := s.eventDone
info := s.socketInfo
sessions := make([]*framedSession, 0, len(s.sessions))
for session := range s.sessions {
sessions = append(sessions, session)
}
s.mu.Unlock()
var closeErr error
if listener != nil {
closeErr = listener.Close()
}
if cancelEvents != nil {
cancelEvents()
}
for _, session := range sessions {
_ = session.Close()
}
if done != nil {
<-done
}
if eventDone != nil {
<-eventDone
}
s.sessionWait.Wait()
removeErr := s.removeOwnedSocket(info)
s.mu.Lock()
s.listener = nil
s.socketInfo = nil
s.acceptDone = nil
s.eventCancel = nil
s.eventDone = nil
s.stopping = false
s.sessions = make(map[*framedSession]struct{})
s.mu.Unlock()
return errors.Join(closeErr, removeErr)
}
func (s *Server) broadcastLoop(
events <-chan *iop.AgentLocalEnvelope,
done chan<- struct{},
) {
defer close(done)
for event := range events {
s.mu.Lock()
sessions := make([]*framedSession, 0, len(s.sessions))
for session := range s.sessions {
sessions = append(sessions, session)
}
s.mu.Unlock()
for _, session := range sessions {
if err := session.communicator.Send(event); err != nil {
_ = session.Close()
}
}
}
}
func (s *Server) prepareStateRoot() error {
info, err := os.Lstat(s.config.StateRoot)
if errors.Is(err, os.ErrNotExist) {
if err := os.MkdirAll(s.config.StateRoot, 0o700); err != nil {
return fmt.Errorf("localcontrol: create state root: %w", err)
}
info, err = os.Lstat(s.config.StateRoot)
}
if err != nil {
return fmt.Errorf("localcontrol: inspect state root: %w", err)
}
if info.Mode()&os.ModeSymlink != 0 ||
!info.IsDir() ||
info.Mode().Perm() != 0o700 {
return ErrUnsafeStateRoot
}
owner, err := fileOwnerUID(info)
if err != nil || owner != s.effectiveUID() {
return ErrUnsafeStateRoot
}
return nil
}
func (s *Server) acceptLoop(
listener *net.UnixListener,
done chan struct{},
) {
defer close(done)
for {
conn, err := listener.AcceptUnix()
if err != nil {
if errors.Is(err, net.ErrClosed) {
return
}
s.mu.Lock()
active := s.started && s.listener == listener
s.mu.Unlock()
if !active {
return
}
continue
}
peerUID, credentialErr := s.credentials.UID(conn)
if credentialErr != nil || peerUID != s.effectiveUID() {
_ = writeEnvelopePacket(
conn,
errorEnvelope(
nil,
protocolError(
ErrorPermissionDenied,
"local-control peer is not authorized",
),
),
0,
)
_ = conn.Close()
continue
}
session := newFramedSession(conn, s.service)
s.mu.Lock()
if !s.started {
s.mu.Unlock()
_ = session.Close()
continue
}
s.sessions[session] = struct{}{}
s.sessionWait.Add(1)
s.mu.Unlock()
go func() {
defer s.sessionWait.Done()
session.Run()
s.mu.Lock()
delete(s.sessions, session)
s.mu.Unlock()
}()
}
}
func (s *Server) removeOwnedSocket(expected os.FileInfo) error {
path := s.Path()
current, err := os.Lstat(path)
if errors.Is(err, os.ErrNotExist) {
return nil
}
if err != nil {
return err
}
if expected == nil {
s.mu.Lock()
expected = s.socketInfo
s.mu.Unlock()
}
if expected == nil ||
current.Mode()&os.ModeSymlink != 0 ||
current.Mode()&os.ModeSocket == 0 ||
!os.SameFile(expected, current) {
return fmt.Errorf("%w: socket path was replaced", ErrUnsafeSocketPath)
}
return os.Remove(path)
}
type framedSession struct {
conn net.Conn
service *Service
ctx context.Context
cancel context.CancelFunc
communicator *proto_socket.Communicator
writeMu sync.Mutex
closeOnce sync.Once
}
func newFramedSession(conn net.Conn, service *Service) *framedSession {
ctx, cancel := context.WithCancel(context.Background())
session := &framedSession{
conn: conn,
service: service,
ctx: ctx,
cancel: cancel,
}
parserMap := proto_socket.ParserMap{
proto_socket.TypeNameOf(&iop.AgentLocalEnvelope{}): func(
payload []byte,
) (proto.Message, error) {
return ParseEnvelope(payload)
},
}
session.communicator = proto_socket.NewCommunicator(session, parserMap)
proto_socket.AddRequestListenerTyped[
*iop.AgentLocalEnvelope,
*iop.AgentLocalEnvelope,
](session.communicator, func(
request *iop.AgentLocalEnvelope,
) (*iop.AgentLocalEnvelope, error) {
return service.Handle(session.ctx, true, request), nil
})
return session
}
func (s *framedSession) Run() {
defer s.Close()
header := make([]byte, 4)
for {
if _, err := io.ReadFull(s.conn, header); err != nil {
return
}
length := binary.BigEndian.Uint32(header)
if length == 0 {
continue
}
if length > MaxFrameBytes {
return
}
payload := make([]byte, int(length))
if _, err := io.ReadFull(s.conn, payload); err != nil {
return
}
base := &packets.PacketBase{}
if err := proto.Unmarshal(payload, base); err != nil {
return
}
if base.GetResponseNonce() != 0 ||
(base.GetTypeName() != proto_socket.TypeNameOf(&iop.AgentLocalEnvelope{}) &&
base.GetTypeName() != "AgentLocalEnvelope") {
_ = s.writeEnvelope(
errorEnvelope(nil, malformed("unexpected local-control frame type")),
base.GetNonce(),
)
continue
}
if _, err := ParseEnvelope(base.GetData()); err != nil {
_ = s.writeEnvelope(
errorEnvelope(nil, malformed("local-control payload is not valid protobuf")),
base.GetNonce(),
)
continue
}
s.communicator.OnReceivedData(
base.GetTypeName(),
base.GetData(),
base.GetNonce(),
0,
)
}
}
func (s *framedSession) WritePacket(base *packets.PacketBase) error {
payload, err := proto.Marshal(base)
if err != nil {
return err
}
if len(payload) > MaxFrameBytes {
return fmt.Errorf("localcontrol: outgoing frame exceeds limit")
}
header := make([]byte, 4)
binary.BigEndian.PutUint32(header, uint32(len(payload)))
s.writeMu.Lock()
defer s.writeMu.Unlock()
if err := writeAll(s.conn, header); err != nil {
return err
}
return writeAll(s.conn, payload)
}
func (s *framedSession) Close() error {
var err error
s.closeOnce.Do(func() {
s.cancel()
err = s.conn.Close()
_ = s.communicator.Close()
})
return err
}
func (s *framedSession) writeEnvelope(
envelope *iop.AgentLocalEnvelope,
responseNonce int32,
) error {
envelopePayload, err := proto.Marshal(envelope)
if err != nil {
return err
}
return s.WritePacket(&packets.PacketBase{
TypeName: proto_socket.TypeNameOf(envelope),
Nonce: 1,
ResponseNonce: responseNonce,
Data: envelopePayload,
})
}
func writeEnvelopePacket(
conn net.Conn,
envelope *iop.AgentLocalEnvelope,
responseNonce int32,
) error {
envelopePayload, err := proto.Marshal(envelope)
if err != nil {
return err
}
packetPayload, err := proto.Marshal(&packets.PacketBase{
TypeName: proto_socket.TypeNameOf(envelope),
Nonce: 1,
ResponseNonce: responseNonce,
Data: envelopePayload,
})
if err != nil {
return err
}
if len(packetPayload) > MaxFrameBytes {
return fmt.Errorf("localcontrol: outgoing frame exceeds limit")
}
header := make([]byte, 4)
binary.BigEndian.PutUint32(header, uint32(len(packetPayload)))
if err := writeAll(conn, header); err != nil {
return err
}
return writeAll(conn, packetPayload)
}
func writeAll(writer io.Writer, payload []byte) error {
for len(payload) != 0 {
written, err := writer.Write(payload)
if err != nil {
return err
}
if written == 0 {
return io.ErrShortWrite
}
payload = payload[written:]
}
return nil
}