Files

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})
}