// Package load installs and starts a kernel driver as a Windows service. package load import ( "fmt" "path/filepath" "strings" "time" "golang.org/x/sys/windows" "golang.org/x/sys/windows/svc" "golang.org/x/sys/windows/svc/mgr" ) // LoadDriver installs the driver as a demand-start kernel service and starts it. // If the service already exists and is running, it is left untouched. It returns // the service name and whether the service was newly created (as opposed to // already existing). func LoadDriver(path string) (name string, created bool, err error) { name = ServiceName(path) m, err := mgr.Connect() if err != nil { return name, false, fmt.Errorf("connect to service control manager: %w", err) } defer m.Disconnect() s, err := m.OpenService(name) if err != nil { s, err = m.CreateService(name, path, mgr.Config{ ServiceType: windows.SERVICE_KERNEL_DRIVER, StartType: windows.SERVICE_DEMAND_START, ErrorControl: windows.SERVICE_ERROR_NORMAL, }) if err != nil { return name, false, fmt.Errorf("create service %q: %w", name, err) } created = true } defer s.Close() status, err := s.Query() if err != nil { return name, created, fmt.Errorf("query service %q: %w", name, err) } if status.State == svc.Running { return name, created, nil } if err := s.Start(); err != nil { return name, created, fmt.Errorf("start service %q: %w", name, err) } // Wait for the driver to reach RUNNING. A driver whose DriverEntry fails // (unsigned, or a load-time error) transitions back to STOPPED. for i := 0; i < 100; i++ { status, err = s.Query() if err != nil { return name, created, fmt.Errorf("query service %q: %w", name, err) } switch status.State { case svc.Running: return name, created, nil case svc.Stopped: return name, created, fmt.Errorf("service %q stopped immediately (unsigned driver? test mode off?)", name) } time.Sleep(100 * time.Millisecond) } return name, created, fmt.Errorf("service %q did not reach RUNNING state", name) } // UnloadDriver stops and deletes a driver service. It is a no-op if the service // does not exist. func UnloadDriver(name string) error { m, err := mgr.Connect() if err != nil { return fmt.Errorf("connect to service control manager: %w", err) } defer m.Disconnect() s, err := m.OpenService(name) if err != nil { if err == windows.ERROR_SERVICE_DOES_NOT_EXIST { return nil } return fmt.Errorf("open service %q: %w", name, err) } defer s.Close() status, err := s.Query() if err != nil { return fmt.Errorf("query service %q: %w", name, err) } if status.State != svc.Stopped { if _, err := s.Control(svc.Stop); err != nil { return fmt.Errorf("stop service %q: %w", name, err) } // Wait for the service to reach STOPPED. for i := 0; i < 100; i++ { status, err = s.Query() if err != nil { return fmt.Errorf("query service %q: %w", name, err) } if status.State == svc.Stopped { break } time.Sleep(100 * time.Millisecond) } } if err := s.Delete(); err != nil { return fmt.Errorf("delete service %q: %w", name, err) } return nil } // ServiceName derives a service name from the driver file name. func ServiceName(path string) string { base := filepath.Base(path) return strings.TrimSuffix(base, filepath.Ext(base)) }