200 lines
6.8 KiB
Go
200 lines
6.8 KiB
Go
|
|
// Package daemonserver accepts local CLI requests over the frozen Unix Socket.
|
||
|
|
package daemonserver
|
||
|
|
|
||
|
|
import (
|
||
|
|
"context"
|
||
|
|
"encoding/json"
|
||
|
|
"errors"
|
||
|
|
"fmt"
|
||
|
|
"io"
|
||
|
|
"log/slog"
|
||
|
|
"net"
|
||
|
|
"os"
|
||
|
|
"path/filepath"
|
||
|
|
"strings"
|
||
|
|
"sync"
|
||
|
|
"time"
|
||
|
|
|
||
|
|
"yms-daemon/internal/backendupdate"
|
||
|
|
"yms-daemon/internal/daemonapi"
|
||
|
|
"yms-daemon/internal/transaction"
|
||
|
|
)
|
||
|
|
|
||
|
|
const maximumRequestBytes = 1 << 20
|
||
|
|
|
||
|
|
type backendUpdater interface {
|
||
|
|
UpdateRepack(context.Context, string, backendupdate.ProgressReporter) (transaction.Transaction, error)
|
||
|
|
UpdateNativeJAR(context.Context, string, backendupdate.ProgressReporter) (transaction.Transaction, error)
|
||
|
|
Restart(context.Context, backendupdate.ProgressReporter) (transaction.Transaction, error)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Server owns the local Unix Socket and dispatches requests to the transaction orchestrator.
|
||
|
|
type Server struct {
|
||
|
|
socketPath string
|
||
|
|
updater backendUpdater
|
||
|
|
logger *slog.Logger
|
||
|
|
}
|
||
|
|
|
||
|
|
func New(socketPath string, updater backendUpdater, logger *slog.Logger) (*Server, error) {
|
||
|
|
if !filepath.IsAbs(socketPath) || filepath.Clean(socketPath) != socketPath {
|
||
|
|
return nil, errors.New("daemon Unix Socket path must be a clean absolute path")
|
||
|
|
}
|
||
|
|
if updater == nil {
|
||
|
|
return nil, errors.New("backend updater is required")
|
||
|
|
}
|
||
|
|
if logger == nil {
|
||
|
|
logger = slog.Default()
|
||
|
|
}
|
||
|
|
return &Server{socketPath: socketPath, updater: updater, logger: logger}, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// Serve listens until ctx is canceled. Each accepted request owns its response connection.
|
||
|
|
func (s *Server) Serve(ctx context.Context) error {
|
||
|
|
if err := prepareSocketPath(s.socketPath); err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
listener, err := net.Listen("unix", s.socketPath)
|
||
|
|
if err != nil {
|
||
|
|
return fmt.Errorf("listen on daemon Unix Socket %s: %w", s.socketPath, err)
|
||
|
|
}
|
||
|
|
if err := os.Chmod(s.socketPath, 0o600); err != nil {
|
||
|
|
_ = listener.Close()
|
||
|
|
return fmt.Errorf("set daemon Unix Socket permissions: %w", err)
|
||
|
|
}
|
||
|
|
defer func() {
|
||
|
|
_ = listener.Close()
|
||
|
|
_ = os.Remove(s.socketPath)
|
||
|
|
}()
|
||
|
|
|
||
|
|
s.logger.InfoContext(ctx, "daemon server listening", "socket", s.socketPath)
|
||
|
|
var connections sync.WaitGroup
|
||
|
|
defer connections.Wait()
|
||
|
|
go func() {
|
||
|
|
<-ctx.Done()
|
||
|
|
_ = listener.Close()
|
||
|
|
}()
|
||
|
|
for {
|
||
|
|
connection, err := listener.Accept()
|
||
|
|
if err != nil {
|
||
|
|
if ctx.Err() != nil {
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
return fmt.Errorf("accept daemon Unix Socket connection: %w", err)
|
||
|
|
}
|
||
|
|
connections.Add(1)
|
||
|
|
go func() {
|
||
|
|
defer connections.Done()
|
||
|
|
s.handle(ctx, connection)
|
||
|
|
}()
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *Server) handle(ctx context.Context, connection net.Conn) {
|
||
|
|
defer connection.Close()
|
||
|
|
_ = connection.SetReadDeadline(time.Now().Add(10 * time.Second))
|
||
|
|
request, err := decodeRequest(connection)
|
||
|
|
_ = connection.SetReadDeadline(time.Time{})
|
||
|
|
if err != nil {
|
||
|
|
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: err.Error()})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
if request.Service != "backend" {
|
||
|
|
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "service must be backend"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
progressWritable := true
|
||
|
|
report := func(progress backendupdate.Progress) {
|
||
|
|
if !progressWritable {
|
||
|
|
return
|
||
|
|
}
|
||
|
|
err := s.writeResponse(connection, daemonapi.Response{
|
||
|
|
Kind: daemonapi.ResponseProgress,
|
||
|
|
TransactionID: progress.TransactionID,
|
||
|
|
State: string(progress.State),
|
||
|
|
Message: progress.Message,
|
||
|
|
})
|
||
|
|
if err != nil {
|
||
|
|
progressWritable = false
|
||
|
|
s.logger.WarnContext(ctx, "write daemon update progress", "error", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
var record transaction.Transaction
|
||
|
|
var updateErr error
|
||
|
|
switch request.Operation {
|
||
|
|
case daemonapi.OperationUpdate:
|
||
|
|
switch request.InputType {
|
||
|
|
case daemonapi.InputTypeRepackZIP:
|
||
|
|
record, updateErr = s.updater.UpdateRepack(ctx, request.File, report)
|
||
|
|
case daemonapi.InputTypeNativeJAR:
|
||
|
|
record, updateErr = s.updater.UpdateNativeJAR(ctx, request.File, report)
|
||
|
|
default:
|
||
|
|
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "inputType must be repack-zip or native-jar"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
case daemonapi.OperationRestart:
|
||
|
|
if request.InputType != "" || request.File != "" {
|
||
|
|
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "restart does not accept inputType or file"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
record, updateErr = s.updater.Restart(ctx, report)
|
||
|
|
default:
|
||
|
|
_ = s.writeResponse(connection, daemonapi.Response{Kind: daemonapi.ResponseResult, Error: "operation must be update or restart"})
|
||
|
|
return
|
||
|
|
}
|
||
|
|
response := daemonapi.Response{Kind: daemonapi.ResponseResult, TransactionID: record.ID, State: string(record.State)}
|
||
|
|
if updateErr != nil {
|
||
|
|
response.Error = updateErr.Error()
|
||
|
|
s.logger.ErrorContext(ctx, "backend operation failed", "operation", request.Operation, "transaction_id", record.ID, "state", record.State, "error", updateErr)
|
||
|
|
} else {
|
||
|
|
s.logger.InfoContext(ctx, "backend operation completed", "operation", request.Operation, "transaction_id", record.ID, "state", record.State)
|
||
|
|
}
|
||
|
|
if err := s.writeResponse(connection, response); err != nil {
|
||
|
|
s.logger.ErrorContext(ctx, "write daemon operation result", "operation", request.Operation, "transaction_id", record.ID, "error", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *Server) writeResponse(connection net.Conn, response daemonapi.Response) error {
|
||
|
|
return json.NewEncoder(connection).Encode(response)
|
||
|
|
}
|
||
|
|
|
||
|
|
func decodeRequest(reader io.Reader) (daemonapi.Request, error) {
|
||
|
|
limited := io.LimitReader(reader, maximumRequestBytes+1)
|
||
|
|
decoder := json.NewDecoder(limited)
|
||
|
|
decoder.DisallowUnknownFields()
|
||
|
|
var request daemonapi.Request
|
||
|
|
if err := decoder.Decode(&request); err != nil {
|
||
|
|
return daemonapi.Request{}, fmt.Errorf("decode daemon request: %w", err)
|
||
|
|
}
|
||
|
|
var trailing any
|
||
|
|
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
|
||
|
|
if err == nil {
|
||
|
|
return daemonapi.Request{}, errors.New("daemon request contains multiple JSON values")
|
||
|
|
}
|
||
|
|
return daemonapi.Request{}, fmt.Errorf("decode daemon request trailing content: %w", err)
|
||
|
|
}
|
||
|
|
if strings.TrimSpace(request.Operation) != request.Operation || strings.TrimSpace(request.Service) != request.Service || strings.TrimSpace(request.InputType) != request.InputType || strings.TrimSpace(request.File) != request.File {
|
||
|
|
return daemonapi.Request{}, errors.New("daemon request fields must not contain surrounding whitespace")
|
||
|
|
}
|
||
|
|
return request, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func prepareSocketPath(socketPath string) error {
|
||
|
|
if err := os.MkdirAll(filepath.Dir(socketPath), 0o755); err != nil {
|
||
|
|
return fmt.Errorf("create daemon Unix Socket directory: %w", err)
|
||
|
|
}
|
||
|
|
info, err := os.Lstat(socketPath)
|
||
|
|
if errors.Is(err, os.ErrNotExist) {
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
if err != nil {
|
||
|
|
return fmt.Errorf("inspect daemon Unix Socket path: %w", err)
|
||
|
|
}
|
||
|
|
if info.Mode()&os.ModeSocket == 0 {
|
||
|
|
return fmt.Errorf("daemon Unix Socket path is occupied by a non-socket file: %s", socketPath)
|
||
|
|
}
|
||
|
|
if err := os.Remove(socketPath); err != nil {
|
||
|
|
return fmt.Errorf("remove stale daemon Unix Socket: %w", err)
|
||
|
|
}
|
||
|
|
return nil
|
||
|
|
}
|