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 }