546 lines
12 KiB
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
|
|
}
|