282 lines
7.7 KiB
Go
282 lines
7.7 KiB
Go
//go:build windows
|
|
|
|
package windowsservice
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"syscall"
|
|
"time"
|
|
|
|
"golang.org/x/sys/windows"
|
|
"golang.org/x/sys/windows/registry"
|
|
"golang.org/x/sys/windows/svc"
|
|
"golang.org/x/sys/windows/svc/mgr"
|
|
)
|
|
|
|
func connect() (*mgr.Mgr, error) {
|
|
return mgr.Connect()
|
|
}
|
|
|
|
func installNative(spec InstallSpec) error {
|
|
if err := spec.Validate(); err != nil {
|
|
return err
|
|
}
|
|
manager, err := connect()
|
|
if err != nil {
|
|
return fmt.Errorf("connect to service control manager: %w", err)
|
|
}
|
|
defer manager.Disconnect()
|
|
startup := uint32(mgr.StartAutomatic)
|
|
if spec.Startup == StartupManual {
|
|
startup = mgr.StartManual
|
|
}
|
|
configuration := mgr.Config{
|
|
ServiceType: windows.SERVICE_WIN32_OWN_PROCESS,
|
|
StartType: startup,
|
|
ErrorControl: mgr.ErrorNormal,
|
|
DisplayName: DisplayName,
|
|
Description: Description,
|
|
ServiceStartName: "LocalSystem",
|
|
}
|
|
service, openErr := manager.OpenService(Name)
|
|
if errors.Is(openErr, windows.ERROR_SERVICE_DOES_NOT_EXIST) {
|
|
service, err = manager.CreateService(Name, spec.ExecutablePath, configuration, "--service", "--config", spec.ConfigPath)
|
|
if err != nil {
|
|
return fmt.Errorf("create %s service: %w", Name, err)
|
|
}
|
|
} else if openErr != nil {
|
|
return fmt.Errorf("open %s service: %w", Name, openErr)
|
|
} else {
|
|
configuration, err = service.Config()
|
|
if err != nil {
|
|
return fmt.Errorf("query %s service configuration: %w", Name, err)
|
|
}
|
|
configuration.ServiceType = windows.SERVICE_WIN32_OWN_PROCESS
|
|
configuration.StartType = startup
|
|
configuration.ErrorControl = mgr.ErrorNormal
|
|
configuration.BinaryPathName = serviceImage(spec)
|
|
configuration.DisplayName = DisplayName
|
|
configuration.Description = Description
|
|
configuration.ServiceStartName = "LocalSystem"
|
|
if err := service.UpdateConfig(configuration); err != nil {
|
|
return fmt.Errorf("update %s service configuration: %w", Name, err)
|
|
}
|
|
}
|
|
defer service.Close()
|
|
if err := service.Start(); err != nil && !errors.Is(err, windows.ERROR_SERVICE_ALREADY_RUNNING) {
|
|
return fmt.Errorf("start %s service: %w", Name, err)
|
|
}
|
|
if err := registerTray(spec); err != nil {
|
|
return fmt.Errorf("register per-user tray: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func serviceImage(spec InstallSpec) string {
|
|
return syscall.EscapeArg(spec.ExecutablePath) + " --service --config " + syscall.EscapeArg(spec.ConfigPath)
|
|
}
|
|
|
|
func uninstallNative() error {
|
|
var serviceErr error
|
|
manager, err := connect()
|
|
if err != nil {
|
|
return fmt.Errorf("connect to service control manager: %w", err)
|
|
}
|
|
defer manager.Disconnect()
|
|
service, err := manager.OpenService(Name)
|
|
if errors.Is(err, windows.ERROR_SERVICE_DOES_NOT_EXIST) {
|
|
serviceErr = nil
|
|
} else if err != nil {
|
|
return fmt.Errorf("open %s service: %w", Name, err)
|
|
} else {
|
|
defer service.Close()
|
|
if err := stopService(service, 30*time.Second); err != nil {
|
|
return err
|
|
}
|
|
if err := service.Delete(); err != nil && !errors.Is(err, windows.ERROR_SERVICE_MARKED_FOR_DELETE) {
|
|
return fmt.Errorf("delete %s service: %w", Name, err)
|
|
}
|
|
}
|
|
if err := removeTrayRegistration(); err != nil {
|
|
return fmt.Errorf("remove per-user tray: %w", err)
|
|
}
|
|
return serviceErr
|
|
}
|
|
|
|
func configureNative(startup StartupMode) error {
|
|
if startup != StartupAutomatic && startup != StartupManual {
|
|
return ErrInvalidInstallSpec
|
|
}
|
|
manager, err := connect()
|
|
if err != nil {
|
|
return fmt.Errorf("connect to service control manager: %w", err)
|
|
}
|
|
defer manager.Disconnect()
|
|
service, err := manager.OpenService(Name)
|
|
if errors.Is(err, windows.ERROR_SERVICE_DOES_NOT_EXIST) {
|
|
return fmt.Errorf("%s service is not installed", Name)
|
|
}
|
|
if err != nil {
|
|
return fmt.Errorf("open %s service: %w", Name, err)
|
|
}
|
|
defer service.Close()
|
|
configuration, err := service.Config()
|
|
if err != nil {
|
|
return fmt.Errorf("query %s service configuration: %w", Name, err)
|
|
}
|
|
if startup == StartupAutomatic {
|
|
configuration.StartType = mgr.StartAutomatic
|
|
} else {
|
|
configuration.StartType = mgr.StartManual
|
|
}
|
|
if err := service.UpdateConfig(configuration); err != nil {
|
|
return fmt.Errorf("update %s startup type: %w", Name, err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
const trayRunValue = "RVBoxTray"
|
|
|
|
func registerTray(spec InstallSpec) error {
|
|
key, _, err := registry.CreateKey(registry.LOCAL_MACHINE, `SOFTWARE\Microsoft\Windows\CurrentVersion\Run`, registry.SET_VALUE)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer key.Close()
|
|
image := syscall.EscapeArg(spec.ExecutablePath) + " --tray --config " + syscall.EscapeArg(spec.ConfigPath)
|
|
return key.SetStringValue(trayRunValue, image)
|
|
}
|
|
|
|
func removeTrayRegistration() error {
|
|
key, err := registry.OpenKey(registry.LOCAL_MACHINE, `SOFTWARE\Microsoft\Windows\CurrentVersion\Run`, registry.SET_VALUE)
|
|
if errors.Is(err, registry.ErrNotExist) {
|
|
return nil
|
|
}
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer key.Close()
|
|
if err := key.DeleteValue(trayRunValue); err != nil && !errors.Is(err, registry.ErrNotExist) {
|
|
return err
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func startNative() error {
|
|
manager, err := connect()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer manager.Disconnect()
|
|
service, err := manager.OpenService(Name)
|
|
if errors.Is(err, windows.ERROR_SERVICE_DOES_NOT_EXIST) {
|
|
return fmt.Errorf("%s service is not installed", Name)
|
|
}
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer service.Close()
|
|
if err := service.Start(); err != nil && !errors.Is(err, windows.ERROR_SERVICE_ALREADY_RUNNING) {
|
|
return err
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func stopNative(timeoutSeconds uint32) error {
|
|
manager, err := connect()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer manager.Disconnect()
|
|
service, err := manager.OpenService(Name)
|
|
if errors.Is(err, windows.ERROR_SERVICE_DOES_NOT_EXIST) {
|
|
return nil
|
|
}
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer service.Close()
|
|
timeout := 30 * time.Second
|
|
if timeoutSeconds > 0 {
|
|
timeout = time.Duration(timeoutSeconds) * time.Second
|
|
}
|
|
return stopService(service, timeout)
|
|
}
|
|
|
|
func restartNative(timeoutSeconds uint32) error {
|
|
if err := stopNative(timeoutSeconds); err != nil {
|
|
return err
|
|
}
|
|
return startNative()
|
|
}
|
|
|
|
func stopService(service *mgr.Service, timeout time.Duration) error {
|
|
status, err := service.Query()
|
|
if err != nil {
|
|
return fmt.Errorf("query %s service: %w", Name, err)
|
|
}
|
|
if status.State == svc.Stopped {
|
|
return nil
|
|
}
|
|
if _, err := service.Control(svc.Stop); err != nil && !errors.Is(err, windows.ERROR_SERVICE_NOT_ACTIVE) {
|
|
return fmt.Errorf("stop %s service: %w", Name, err)
|
|
}
|
|
deadline := time.Now().Add(timeout)
|
|
for {
|
|
status, err = service.Query()
|
|
if err != nil {
|
|
return fmt.Errorf("query %s service while stopping: %w", Name, err)
|
|
}
|
|
if status.State == svc.Stopped {
|
|
return nil
|
|
}
|
|
if time.Now().After(deadline) {
|
|
return fmt.Errorf("timed out stopping %s service", Name)
|
|
}
|
|
time.Sleep(100 * time.Millisecond)
|
|
}
|
|
}
|
|
|
|
type handler struct {
|
|
run func(context.Context) error
|
|
}
|
|
|
|
func (serviceHandler handler) Execute(_ []string, changes <-chan svc.ChangeRequest, status chan<- svc.Status) (bool, uint32) {
|
|
if serviceHandler.run == nil {
|
|
return false, 1
|
|
}
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
defer cancel()
|
|
status <- svc.Status{State: svc.StartPending, WaitHint: 10_000}
|
|
done := make(chan error, 1)
|
|
go func() { done <- serviceHandler.run(ctx) }()
|
|
status <- svc.Status{State: svc.Running, Accepts: svc.AcceptStop | svc.AcceptShutdown}
|
|
for {
|
|
select {
|
|
case request := <-changes:
|
|
if request.Cmd == svc.Stop || request.Cmd == svc.Shutdown {
|
|
status <- svc.Status{State: svc.StopPending, WaitHint: 30_000}
|
|
cancel()
|
|
err := <-done
|
|
if err != nil {
|
|
return true, 1
|
|
}
|
|
status <- svc.Status{State: svc.Stopped}
|
|
return false, 0
|
|
}
|
|
case err := <-done:
|
|
if err != nil {
|
|
return true, 1
|
|
}
|
|
status <- svc.Status{State: svc.Stopped}
|
|
return false, 0
|
|
}
|
|
}
|
|
}
|
|
|
|
func runNative(run func(context.Context) error) error {
|
|
return svc.Run(Name, handler{run: run})
|
|
}
|