diff --git a/src/ectl/main.go b/src/ectl/main.go index e13e8c8..9c851c2 100644 --- a/src/ectl/main.go +++ b/src/ectl/main.go @@ -9,6 +9,7 @@ import ( "net" "os" "path" + "strconv" "syscall" "time" ) @@ -58,7 +59,7 @@ func main() { if len(flag.Args()) <= 1 { fmt.Println("Usage: ectl service [service]") return - } else if flag.Arg(1) == "start" || flag.Arg(1) == "stop" || flag.Arg(1) == "restart" || flag.Arg(1) == "enable" || flag.Arg(1) == "disable" { + } else if flag.Arg(1) == "start" || flag.Arg(1) == "stop" || flag.Arg(1) == "restart" { // Ensure service name argument has been set if len(flag.Args()) <= 2 { fmt.Printf("Usage: ectl service %s \n", flag.Args()[1]) @@ -118,6 +119,82 @@ func main() { log.Fatal("Connection returned empty string!") } + return + } else if flag.Arg(1) == "enable" || flag.Arg(1) == "disable" { + // Ensure service name argument has been set + if len(flag.Args()) <= 2 { + fmt.Printf("Usage: ectl service %s [stage]\n", flag.Args()[1]) + return + } + + // Get service stage + stage := 2 + if len(flag.Args()) > 3 { + flagStr := flag.Arg(3) + _stage, err := strconv.ParseInt(flagStr, 10, 32) + if err != nil { + log.Fatalf("Error: could not parse stage number: %s", err) + } + stage = int(_stage) + } else if flag.Arg(1) == "disable" { + stage = 0 + } + + type ServiceCommandJsonStruct struct { + Command string `json:"command"` + Service string `json:"service"` + Stage int `json:"stage"` + } + serviceCommandJson := ServiceCommandJsonStruct{ + Command: "set_enabled", + Service: flag.Arg(2), + Stage: stage, + } + + // Encode struct to json string + jsonData, err := json.Marshal(serviceCommandJson) + if err != nil { + log.Fatalf("Could not encode JSON data! Error: %s\n", err) + } + + _, err = conn.Write(jsonData) + if err != nil { + log.Fatalf("Could not write JSON data to socket! Error: %s\n", err) + } + + // Create a buffer for incoming data. + buf := make([]byte, 4096) + + // Read data from the connection. + n, err := conn.Read(buf) + if err == io.EOF { + return + } + if err != nil { + return + } + + // Print json data if flag is set + if *printJson { + fmt.Println(string(buf[:n])) + return + } + + // Decoode JSON data + var returnedJsonData map[string]any + err = json.Unmarshal(buf[:n], &returnedJsonData) + if err != nil { + log.Fatalf("Could not decode JSON data from connection!") + } + + if err, ok := returnedJsonData["error"]; ok { + log.Fatal(err) + } else if msg, ok := returnedJsonData["success"]; ok { + fmt.Println(msg) + } else { + log.Fatal("Connection returned empty string!") + } + return } else if flag.Args()[1] == "status" { // Ensure service name argument has been set @@ -177,10 +254,15 @@ func main() { serviceState := returnedJsonData["state"].(string) serviceEnabled := returnedJsonData["is_enabled"].(bool) + serviceStage := int(returnedJsonData["stage"].(float64)) fmt.Printf("Name: %s\n", flag.Arg(2)) fmt.Printf("State: %s\n", serviceState) - fmt.Printf("Enabled: %t\n", serviceEnabled) + if serviceEnabled { + fmt.Printf("Enabled: %t (Stage %d)\n", serviceEnabled, serviceStage) + } else { + fmt.Printf("Enabled: %t\n", serviceEnabled) + } return } else if flag.Arg(1) == "list" { @@ -213,7 +295,7 @@ func main() { if err != nil { return } - + // Print json data if flag is set if *printJson { fmt.Println(string(buf[:n])) @@ -235,10 +317,15 @@ func main() { serviceName := serviceMap.(map[string]any)["name"].(string) serviceState := serviceMap.(map[string]any)["state"].(string) serviceEnabled := serviceMap.(map[string]any)["is_enabled"].(bool) + serviceStage := serviceMap.(map[string]any)["stage"].(int) fmt.Printf("Name: %s\n", serviceName) fmt.Printf("State: %s\n", serviceState) - fmt.Printf("Enabled: %t\n", serviceEnabled) + if serviceEnabled { + fmt.Printf("Enabled: %t (Stage %d)\n", serviceEnabled, serviceStage) + } else { + fmt.Printf("Enabled: %t\n", serviceEnabled) + } fmt.Println() } diff --git a/src/enit/main.go b/src/enit/main.go index 4dc5181..6b037ce 100644 --- a/src/enit/main.go +++ b/src/enit/main.go @@ -149,40 +149,40 @@ func startServiceManager() { func stopServiceManager() { fmt.Println("Stopping service manager... ") - err := syscall.Kill(serviceManagerPid, syscall.SIGTERM) - if err != nil { + process, _ := os.FindProcess(serviceManagerPid) + + // Send SIGTERM signal to service manager + if err := process.Signal(syscall.SIGTERM); err != nil { log.Println("Could not stop service manager!") + syscall.Kill(serviceManagerPid, syscall.SIGKILL) + return } // Check if service manager has stopped gracefully, otherwise send sigkill on timeout - exit := false - for timeout := time.After(60 * time.Second); ; { - if exit { - break - } - select { - case <-timeout: - log.Println("Could not stop service manager!") - err := syscall.Kill(serviceManagerPid, syscall.SIGKILL) - if err != nil { - log.Println("Could not stop service manager!") - } - exit = true - default: - waitZombieProcesses() - p, err := os.FindProcess(serviceManagerPid) - if err != nil { - exit = true + exited := make(chan bool) + go func() { + for { + if err := process.Signal(syscall.Signal(0)); err != nil { break } - err = p.Signal(syscall.Signal(0)) - if err != nil { - exit = true - } + } + exited <- true + }() + + for { + select { + case <-exited: + fmt.Println("Done.") + return + case <-time.After(60 * time.Second): + log.Println("Could not stop service manager!") + syscall.Kill(serviceManagerPid, syscall.SIGKILL) + return + default: + waitZombieProcesses() } } - fmt.Println("Done.") } func setHostname() { diff --git a/src/esvm/main.go b/src/esvm/main.go index 62a1941..2a3c1e1 100644 --- a/src/esvm/main.go +++ b/src/esvm/main.go @@ -3,11 +3,14 @@ package main import ( "flag" "fmt" + "io" "log" + "maps" "net" "os" "os/signal" "path" + "slices" "strings" "syscall" "time" @@ -21,8 +24,6 @@ var version = "dev" var runtimeServiceDir string var serviceConfigDir string -var Services = make([]EnitService, 0) - var logger *log.Logger var socket net.Listener @@ -44,7 +45,7 @@ func main() { // Setup main logger err := setupESVMLogger() if err != nil { - log.Printf("Could not setup main ESVM logger! Error: %s\n", err) + log.Printf("Error: could not setup main ESVM logger: %s\n", err) logger = log.Default() } @@ -94,8 +95,11 @@ func setupESVMLogger() error { return err } + // Setup multiwriter + w := io.MultiWriter(loggerFile, os.Stderr) + // Initialize logger and print a header line - logger = log.New(loggerFile, "[ESVM] ", log.Lshortfile|log.LstdFlags) + logger = log.New(w, "[ESVM] ", log.Lshortfile|log.LstdFlags) _, err = loggerFile.WriteString("------ " + time.Now().Format(time.UnixDate) + " ------\n") return nil @@ -105,17 +109,17 @@ func Init() { logger.Println("Initializing ESVM...") if _, err := os.Stat(runtimeServiceDir); err == nil { - logger.Fatalf("Could not initialize ESVM! Error: %s", fmt.Errorf("runtime service directory %s already exists", runtimeServiceDir)) + logger.Fatalf("Error: could not initialize ESVM: %s", fmt.Errorf("runtime service directory %s already exists", runtimeServiceDir)) } err := os.MkdirAll(runtimeServiceDir, 0755) if err != nil { - logger.Fatalf("Could not initialize ESVM! Error: %s", err) + logger.Fatalf("Error: could not initialize ESVM: %s", err) } socket, err = initSocket() if err != nil { - logger.Fatalf("Could not initialize ESVM! Error: %s", err) + logger.Fatalf("Error: could not initialize ESVM: %s", err) } if stat, err := os.Stat(serviceConfigDir); err != nil || !stat.IsDir() { @@ -125,7 +129,7 @@ func Init() { dirEntries, err := os.ReadDir(path.Join(serviceConfigDir, "services")) if err != nil { - logger.Fatalf("Could not initialize ESVM! Error: %s", err) + logger.Fatalf("Error: Could not initialize ESVM: %s", err) } // Read and initialize service files @@ -134,7 +138,7 @@ func Init() { logger.Printf("Initializing service (%s)...\n", entry.Name()) bytes, err := os.ReadFile(path.Join(serviceConfigDir, "services", entry.Name())) if err != nil { - logger.Printf("Could not read service file at %s!\n", path.Join(serviceConfigDir, "services", entry.Name())) + logger.Printf("Error: Could not read service file (%s)", path.Join(serviceConfigDir, "services", entry.Name())) continue } @@ -154,27 +158,27 @@ func Init() { LogOutput: true, } if err := yaml.Unmarshal(bytes, &service); err != nil { - logger.Printf("Could not read service file at %s!\n", path.Join(serviceConfigDir, "services", entry.Name())) + logger.Printf("Error: could not read service file %s", path.Join(serviceConfigDir, "services", entry.Name())) continue } for _, sv := range Services { if sv.Name == service.Name { - logger.Printf("Service with name (%s) has already been initialized!", service.Name) + logger.Printf("Error: service with name (%s) has already been initialized", service.Name) } } switch service.Type { case "simple", "background": default: - logger.Printf("Unknown service type: %s\n", service.Type) + logger.Printf("Error: unknown service type (%s)", service.Type) continue } switch service.ExitMethod { case "stop_command", "kill": default: - logger.Printf("Unknown exit method: %s\n", service.ExitMethod) + logger.Printf("Error: unknown exit method (%s)\n", service.ExitMethod) continue } @@ -187,12 +191,12 @@ func Init() { service.ServiceRunPath = path.Join(runtimeServiceDir, service.Name) err = os.MkdirAll(path.Join(service.ServiceRunPath), 0755) if err != nil { - logger.Fatalf("Could not initialize ESVM! Error: %s", err) + logger.Fatalf("Error: could not initialize ESVM: %s", err) } err = service.setCurrentState(EnitServiceUnloaded) if err != nil { - logger.Fatalf("Could not initialize ESVM! Error: %s", err) + logger.Fatalf("Error: could not initialize ESVM: %s", err) } Services = append(Services, service) @@ -201,51 +205,33 @@ func Init() { } } - // Get enabled services that meet their dependencies - servicesWithMetDepends := make([]EnitService, 0) - for _, service := range Services { - if service.isEnabled() && len(service.GetUnmetDependencies()) == 0 { - servicesWithMetDepends = append(servicesWithMetDepends, service) - } - } + // Read enabled services + ReadEnabledServices() - // Loop until all enabled services have started or timed out - for start := time.Now(); time.Since(start) < 60*time.Second; { - if len(servicesWithMetDepends) == 0 { - break - } + // Start enabled services + stages := slices.Collect(maps.Keys(EnabledServices)) + slices.Sort(stages) + for stage := 0; stage <= stages[len(stages)-1]; stage++ { + logger.Printf("Starting stage %d services...", stage) - for i := len(servicesWithMetDepends) - 1; i >= 0; i-- { - service := servicesWithMetDepends[i] - canStart := true - for _, dependency := range service.Dependencies { - if strings.HasPrefix(dependency, "/") { - // File dependency - if _, err := os.Stat(dependency); err != nil { - canStart = false - break - } - } else { - // Service dependency - if GetServiceByName(dependency).GetCurrentState() != EnitServiceRunning && GetServiceByName(dependency).GetCurrentState() != EnitServiceCompleted { - canStart = false - break + services := EnabledServices[stage] + remainingServices := len(services) + for remainingServices != 0 { + for _, serviceName := range services { + service := GetServiceByName(serviceName) + if service == nil { + remainingServices-- + continue + } + + if len(service.GetUnmetDependencies()) == 0 { + err := service.StartService() + if err != nil { + logger.Printf("Error: could not start service (%s): %s", service.Name, err) } + remainingServices-- } } - if canStart { - err := service.StartService() - if err != nil { - logger.Printf("Could not start service (%s)! Error: %s", service.Name, err) - } - servicesWithMetDepends = append(servicesWithMetDepends[:i], servicesWithMetDepends[i+1:]...) - } - } - } - - if len(servicesWithMetDepends) > 0 { - for _, service := range servicesWithMetDepends { - logger.Printf("Could not start service (%s)! Error: dependencies not met", service.Name) } } @@ -254,11 +240,21 @@ func Init() { func Destroy() { logger.Println("Stopping all ESVM services...") - for _, service := range Services { + + // Loop through all started services in reverse + for i := len(startedServicesOrder) - 1; i >= 0; i-- { + // Get service by name + service := GetServiceByName(startedServicesOrder[i]) + if service == nil { + continue + } + + // Stop service if err := service.StopService(); err != nil { - logger.Printf("Error stopping service %s! Error: %s\n", service.Name, err) + logger.Printf("Error: could not stop service (%s): %s", service.Name, err) } } + logger.Println("All ESVM services have stopped!") } diff --git a/src/esvm/service.go b/src/esvm/service.go index ed565ce..82b3b27 100644 --- a/src/esvm/service.go +++ b/src/esvm/service.go @@ -1,14 +1,17 @@ package main import ( - "io" + "fmt" "os" "os/exec" "path" + "slices" "strconv" "strings" "syscall" "time" + + "gopkg.in/yaml.v3" ) type EnitServiceState uint8 @@ -47,6 +50,10 @@ type EnitService struct { stopChannel chan bool } +var Services = make([]EnitService, 0) +var EnabledServices = make(map[int][]string) +var startedServicesOrder = make([]string, 0) + func (service *EnitService) GetUnmetDependencies() (missingDependencies []string) { for _, dependency := range service.Dependencies { if strings.HasPrefix(dependency, "/") { @@ -211,7 +218,6 @@ func (service *EnitService) StartService() error { select { case <-service.stopChannel: service.restartCount = 0 - _ = service.setCurrentState(EnitServiceStopped) default: if service.Type == "simple" && err == nil { service.restartCount = 0 @@ -239,6 +245,11 @@ func (service *EnitService) StartService() error { } }() + // Add to started services order slice + if !slices.Contains(startedServicesOrder, service.Name) { + startedServicesOrder = append(startedServicesOrder, service.Name) + } + logger.Printf("Service (%s) has started!\n", service.Name) return nil @@ -249,45 +260,44 @@ func (service *EnitService) StopService() error { return nil } - logger.Printf("Stopping service (%s)...\n", service.Name) + logger.Printf("Stopping service (%s)...", service.Name) + + newServiceStatus := EnitServiceCrashed + defer service.setCurrentState(newServiceStatus) + defer service.setProcessID(0) if service.ExitMethod == "kill" { process := service.GetProcess() if err := process.Signal(syscall.Signal(0)); err != nil { - logger.Printf("Service (%s) has stopped. (Process already dead)\n", service.Name) + newServiceStatus = EnitServiceStopped + logger.Printf("Service (%s) has stopped (Process already dead)", service.Name) return nil } go func() { service.stopChannel <- true }() - err := service.GetProcess().Signal(syscall.SIGTERM) - if err != nil { - return err + // Send SIGTERM signal to process + if err := process.Signal(syscall.SIGTERM); err != nil { + process.Signal(syscall.SIGKILL) + return fmt.Errorf("could not stop process gracefully") } - exit := false - for timeout := time.After(5 * time.Second); ; { - if exit { - break - } - select { - case <-timeout: - logger.Println("Process took too long to finish. Forcefully killing process...") - err := service.GetProcess().Kill() - if err != nil { - return err - } - exit = true - default: - if process == nil { - exit = true + // Check if the process has stopped gracefully, otherwise send sigkill on timeout + exited := make(chan bool) + go func() { + for { + if err := process.Signal(syscall.Signal(0)); err != nil { break } - err = process.Signal(syscall.Signal(0)) - if err != nil { - exit = true - } } + exited <- true + }() + + select { + case <-exited: + case <-time.After(5 * time.Second): + process.Signal(syscall.SIGKILL) + return fmt.Errorf("could not stop process gracefully") } } else { cmd := exec.Command("/bin/sh", "-c", service.StopCmd) @@ -296,16 +306,7 @@ func (service *EnitService) StopService() error { } } - err := service.setCurrentState(EnitServiceStopped) - if err != nil { - return err - } - - err = service.setProcessID(0) - if err != nil { - return err - } - + newServiceStatus = EnitServiceStopped logger.Printf("Service (%s) has stopped!\n", service.Name) return nil @@ -325,61 +326,73 @@ func (service *EnitService) RestartService() error { // Functions will be rewritten at some point to allow enabling unloaded services -func (service *EnitService) isEnabled() bool { - contents, err := os.ReadFile(path.Join(serviceConfigDir, "enabled_services")) - if err != nil { - return false - } - - for _, line := range strings.Split(string(contents), "\n") { - line = strings.TrimSpace(line) - if line == "" { - continue - } - - if line == service.Name { - return true +func (service *EnitService) isEnabled() (bool, int) { + for stage, services := range EnabledServices { + if slices.Contains(services, service.Name) { + return true, stage } } - return false + return false, 0 } -func (service *EnitService) SetEnabled(isEnabled bool) error { +func (service *EnitService) SetEnabled(stage int) error { + // Get current service enabled status + _, s := service.isEnabled() + // Return if service is already in correct state - if service.isEnabled() == isEnabled { + if s == stage { return nil } - // Create or open enabled_services file - file, err := os.OpenFile(path.Join(serviceConfigDir, "enabled_services"), os.O_CREATE|os.O_RDWR, 0644) + // Remove service from current stage + if s != 0 { + EnabledServices[s] = slices.DeleteFunc(EnabledServices[s], func(name string) bool { + return name == service.Name + }) + } + + // Add service to stage + EnabledServices[stage] = append(EnabledServices[stage], service.Name) + + // Save enabled services to file + data, err := yaml.Marshal(EnabledServices) if err != nil { return err } - defer file.Close() - - // Get enabled_services file contents - contents, err := io.ReadAll(file) + err = os.WriteFile(path.Join(serviceConfigDir, "enabled_services"), data, 0644) if err != nil { return err } - // Modify contents - strContents := string(contents) - if isEnabled { - strContents += service.Name + "\n" - } else { - strContents = strings.ReplaceAll(strContents, service.Name+"\n", "") + return nil +} + +func ReadEnabledServices() error { + data, err := os.ReadFile(path.Join(serviceConfigDir, "enabled_services")) + if err != nil { + return err } - // Write new contents to file - file.Truncate(0) - file.Seek(0, 0) - _, err = file.WriteString(strContents) + err = yaml.Unmarshal(data, &EnabledServices) if err != nil { - return err + // Assume old plain text format + for _, service := range strings.Split(strings.TrimSpace(string(data)), "\n") { + EnabledServices[2] = append(EnabledServices[2], service) + } + + // Update enabled_services file + data, err := yaml.Marshal(EnabledServices) + if err != nil { + return err + } + err = os.WriteFile(path.Join(serviceConfigDir, "enabled_services"), data, 0644) + if err != nil { + return err + } + + return nil } - file.Sync() return nil } diff --git a/src/esvm/socket.go b/src/esvm/socket.go index edde9cf..588fe36 100644 --- a/src/esvm/socket.go +++ b/src/esvm/socket.go @@ -20,8 +20,7 @@ func initSocket() (socket net.Listener, err error) { commandHandlers["start"] = handleStartServiceCommand commandHandlers["stop"] = handleStopServiceCommand commandHandlers["restart"] = handleRestartServiceCommand - commandHandlers["enable"] = handleEnableServiceCommand - commandHandlers["disable"] = handleDisableServiceCommand + commandHandlers["set_enabled"] = handleSetEnabledServiceCommand commandHandlers["status"] = handleStatusServiceCommand commandHandlers["list"] = handleListServicesCommand @@ -147,7 +146,7 @@ func handleRestartServiceCommand(conn net.Conn, jsonData map[string]any) { conn.Write(wrapSuccessMsgInJson(fmt.Sprintf("Service (%s) has restarted sucessfully", serviceName.(string)))) } -func handleEnableServiceCommand(conn net.Conn, jsonData map[string]any) { +func handleSetEnabledServiceCommand(conn net.Conn, jsonData map[string]any) { // Get service name from json data serviceName, ok := jsonData["service"] if !ok { @@ -155,6 +154,18 @@ func handleEnableServiceCommand(conn net.Conn, jsonData map[string]any) { return } + // Get service stage from json json data + _serviceStage, ok := jsonData["stage"] + if !ok { + conn.Write(wrapErrorInJson(fmt.Errorf("'stage' field missing"))) + return + } + serviceStage, ok := _serviceStage.(float64) + if !ok { + conn.Write(wrapErrorInJson(fmt.Errorf("'stage' field is not a number"))) + return + } + // Ensure service exists service := GetServiceByName(serviceName.(string)) if service == nil { @@ -162,53 +173,33 @@ func handleEnableServiceCommand(conn net.Conn, jsonData map[string]any) { return } - // Check if service is already enabled - if service.isEnabled() { - conn.Write(wrapSuccessMsgInJson(fmt.Sprintf("Service (%s) is already enabled", serviceName.(string)))) + // Get current service enabled status + _, s := service.isEnabled() + + // Return if service is already in correct state + if s == int(serviceStage) { + if serviceStage == 0 { + conn.Write(wrapSuccessMsgInJson(fmt.Sprintf("Service (%s) is already disabled", serviceName.(string)))) + } else { + conn.Write(wrapSuccessMsgInJson(fmt.Sprintf("Service (%s) is already enabled", serviceName.(string)))) + } return } // Enable service - err := service.SetEnabled(true) + err := service.SetEnabled(int(serviceStage)) if err != nil { - conn.Write(wrapErrorInJson(fmt.Errorf("Could not enable service! Error: %s", err))) + if serviceStage == 0 { + conn.Write(wrapErrorInJson(fmt.Errorf("Could not disable service! Error: %s", err))) + } else { + conn.Write(wrapErrorInJson(fmt.Errorf("Could not enable service! Error: %s", err))) + } return } conn.Write(wrapSuccessMsgInJson(fmt.Sprintf("Service (%s) was enabled sucessfully", serviceName.(string)))) } -func handleDisableServiceCommand(conn net.Conn, jsonData map[string]any) { - // Get service name from json data - serviceName, ok := jsonData["service"] - if !ok { - conn.Write(wrapErrorInJson(fmt.Errorf("'service' field missing"))) - return - } - - // Ensure service exists - service := GetServiceByName(serviceName.(string)) - if service == nil { - conn.Write(wrapErrorInJson(fmt.Errorf("Service (%s) not found", serviceName.(string)))) - return - } - - // Check if service is already disabled - if !service.isEnabled() { - conn.Write(wrapSuccessMsgInJson(fmt.Sprintf("Service (%s) is already disabled", serviceName.(string)))) - return - } - - // Disable service - err := service.SetEnabled(false) - if err != nil { - conn.Write(wrapErrorInJson(fmt.Errorf("Could not disable service! Error: %s", err))) - return - } - - conn.Write(wrapSuccessMsgInJson(fmt.Sprintf("Service (%s) was disabled sucessfully", serviceName.(string)))) -} - func handleStatusServiceCommand(conn net.Conn, jsonData map[string]any) { // Get service name from json data serviceName, ok := jsonData["service"] @@ -227,7 +218,7 @@ func handleStatusServiceCommand(conn net.Conn, jsonData map[string]any) { statusMap := make(map[string]any) statusMap["name"] = service.Name statusMap["state"] = EnitServiceStateNames[service.GetCurrentState()] - statusMap["is_enabled"] = service.isEnabled() + statusMap["is_enabled"], statusMap["stage"] = service.isEnabled() // Encode map to json string newJsonData, err := json.Marshal(statusMap) @@ -248,7 +239,7 @@ func handleListServicesCommand(conn net.Conn, _ map[string]any) { statusMap := make(map[string]any) statusMap["name"] = service.Name statusMap["state"] = EnitServiceStateNames[service.GetCurrentState()] - statusMap["is_enabled"] = service.isEnabled() + statusMap["is_enabled"], statusMap["stage"] = service.isEnabled() servicesMap["services"] = append(servicesMap["services"].([]map[string]any), statusMap) }