diff --git a/src/backup.go b/src/backup.go index 10c6102..57c03f2 100644 --- a/src/backup.go +++ b/src/backup.go @@ -148,8 +148,8 @@ func xteveRestore(archive string) (newWebURL string, err error) { return } - backupVersion = newConfig["version"].(string) - if backupVersion < System.Compatibility { + backupVersion, ok := newConfig["version"].(string) + if !ok || backupVersion < System.Compatibility { err = errors.New(getErrMsg(1013)) return } @@ -167,7 +167,11 @@ func xteveRestore(archive string) (newWebURL string, err error) { return } - newPort = newConfig["port"].(string) + newPort, ok = newConfig["port"].(string) + if !ok { + err = errors.New(getErrMsg(1030)) + return + } oldPort = Settings.Port if newPort == oldPort { diff --git a/src/data.go b/src/data.go index 37e4c61..31d69cb 100644 --- a/src/data.go +++ b/src/data.go @@ -63,17 +63,31 @@ func updateServerSettings(request RequestStruct) (settings SettingsStruct, err e // Leerzeichen aus den Werten entfernen und Formatierung der Uhrzeit überprüfen (0000 - 2359) var newUpdateTimes = make([]string, 0) - for _, v := range value.([]any) { + updateTimes, ok := value.([]any) + if !ok { + err = errors.New(getErrMsg(1012)) + ShowError(err, 1012) + return Settings, err + } - v = strings.Replace(v.(string), " ", "", -1) + for _, v := range updateTimes { - _, err := time.Parse("1504", v.(string)) + t, ok := v.(string) + if !ok { + err = errors.New(getErrMsg(1012)) + ShowError(err, 1012) + return Settings, err + } + + t = strings.Replace(t, " ", "", -1) + + _, err := time.Parse("1504", t) if err != nil { ShowError(err, 1012) return Settings, err } - newUpdateTimes = append(newUpdateTimes, v.(string)) + newUpdateTimes = append(newUpdateTimes, t) } @@ -320,21 +334,30 @@ func saveFiles(request RequestStruct, fileType string) (err error) { for dataID, data := range newData { + dataMap, ok := data.(map[string]any) + if !ok { + err = errors.New(getErrMsg(1020)) + return + } + if dataID == "-" { // Neue Providerdatei dataID = indicator + randomString(19) - data.(map[string]any)["new"] = true - filesMap[dataID] = data + dataMap["new"] = true + filesMap[dataID] = dataMap } else { // Bereits vorhandene Providerdatei - for key, value := range data.(map[string]any) { + oldData, ok := filesMap[dataID].(map[string]any) + if !ok { + // Unknown file ID: nothing to update. + continue + } - var oldData = filesMap[dataID].(map[string]any) + for key, value := range dataMap { oldData[key] = value - } } @@ -353,11 +376,11 @@ func saveFiles(request RequestStruct, fileType string) (err error) { } // Neue Providerdatei - if _, ok := data.(map[string]any)["new"]; ok { + if _, ok := dataMap["new"]; ok { reloadData = true err = getProviderData(fileType, dataID) - delete(data.(map[string]any), "new") + delete(dataMap, "new") if err != nil { delete(filesMap, dataID) @@ -661,8 +684,16 @@ func saveUserData(request RequestStruct) (err error) { func saveNewUser(request RequestStruct) (err error) { var data = request.UserData - var username = data["username"].(string) - var password = data["password"].(string) + + username, ok := data["username"].(string) + if !ok { + return errors.New("Username is missing") + } + + password, ok := data["password"].(string) + if !ok { + return errors.New("Password is missing") + } delete(data, "password") delete(data, "confirm") @@ -855,7 +886,7 @@ func buildDatabaseDVR() (err error) { var playlistFile = getLocalProviderFiles(fileType) - for n, i := range playlistFile { + for _, i := range playlistFile { var channels []any var groupTitle, tvgID, uuid = 0, 0, 0 @@ -878,7 +909,8 @@ func buildDatabaseDVR() (err error) { ShowError(err, 1005) err = errors.New(playlistName + ": Local copy of the file no longer exists") ShowError(err, 0) - playlistFile = append(playlistFile[:n], playlistFile[n+1:]...) + // playlistFile is not used after this loop; removing the entry + // while ranging over it only skipped the following file. } // Streams analysieren diff --git a/src/httpclient.go b/src/httpclient.go new file mode 100644 index 0000000..3652db1 --- /dev/null +++ b/src/httpclient.go @@ -0,0 +1,14 @@ +package src + +import ( + "net/http" + "time" +) + +// providerHTTPClient : Downloads of provider files (M3U, XMLTV, HDHR lineup). +// Playlists and EPG files can be large, so the overall timeout is generous; +// it only exists so a stalled provider cannot hang an update forever. +var providerHTTPClient = &http.Client{Timeout: 5 * time.Minute} + +// apiHTTPClient : Short request/response calls to other services (Plex API). +var apiHTTPClient = &http.Client{Timeout: 30 * time.Second} diff --git a/src/internal/authentication/atomic_test.go b/src/internal/authentication/atomic_test.go new file mode 100644 index 0000000..4cb3295 --- /dev/null +++ b/src/internal/authentication/atomic_test.go @@ -0,0 +1,43 @@ +package authentication + +import ( + "os" + "path/filepath" + "testing" +) + +func TestWriteFileAtomic(t *testing.T) { + dir := t.TempDir() + file := filepath.Join(dir, "authentication.json") + + if err := writeFileAtomic(file, []byte("{\"a\":1}"), 0600); err != nil { + t.Fatal(err) + } + if err := writeFileAtomic(file, []byte("{\"a\":2}"), 0600); err != nil { + t.Fatal(err) + } + + got, err := os.ReadFile(file) + if err != nil { + t.Fatal(err) + } + if string(got) != "{\"a\":2}" { + t.Errorf("content = %q", got) + } + + fi, err := os.Stat(file) + if err != nil { + t.Fatal(err) + } + if fi.Mode().Perm() != 0600 { + t.Errorf("mode = %o, want 0600", fi.Mode().Perm()) + } + + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatal(err) + } + if len(entries) != 1 { + t.Errorf("expected only the target file, found %d entries", len(entries)) + } +} diff --git a/src/internal/authentication/authentication.go b/src/internal/authentication/authentication.go index b55cf3c..f51590a 100755 --- a/src/internal/authentication/authentication.go +++ b/src/internal/authentication/authentication.go @@ -411,11 +411,47 @@ func saveDatabase(tmpMap any) (err error) { return } - err = os.WriteFile(database, []byte(jsonString), 0600) + err = writeFileAtomic(database, []byte(jsonString), 0600) + + return +} + +// writeFileAtomic : Temp file + fsync + rename, so authentication.json is +// never observed half-written. (Private copy; this package cannot import src.) +func writeFileAtomic(name string, data []byte, perm os.FileMode) (err error) { + + tmp, err := os.CreateTemp(filepath.Dir(name), ".tmp-*") if err != nil { return } + var tmpName = tmp.Name() + + defer func() { + if err != nil { + tmp.Close() + os.Remove(tmpName) + } + }() + + if _, err = tmp.Write(data); err != nil { + return + } + + if err = tmp.Sync(); err != nil { + return + } + + if err = tmp.Close(); err != nil { + return + } + + if err = os.Chmod(tmpName, perm); err != nil { + return + } + + err = os.Rename(tmpName, name) + return } diff --git a/src/internal/imgcache/cache.go b/src/internal/imgcache/cache.go index 94a5451..80d2ef6 100644 --- a/src/internal/imgcache/cache.go +++ b/src/internal/imgcache/cache.go @@ -9,8 +9,14 @@ import ( "path/filepath" "strings" "sync" + "time" ) +// httpClient : Client for logo downloads. Providers occasionally serve +// images from hosts that never answer; without a timeout the caching +// goroutine would hang forever. +var httpClient = &http.Client{Timeout: 30 * time.Second} + // Cache : Cache strcut type Cache struct { path string @@ -41,8 +47,6 @@ func New(path, chacheURL string, caching bool) (c *Cache, err error) { c.Queue = []string{} c.Cache = []string{} - var queue []string - c.Image.GetURL = func(src string) (cacheURL string) { c.Lock() @@ -89,50 +93,30 @@ func New(path, chacheURL string, caching bool) (c *Cache, err error) { return src } + // Caching downloads every queued image. The lock is only held while the + // queue is copied and while the map / queue are updated per image; the + // HTTP download and file write happen unlocked so GetURL callers (M3U / + // XMLTV generation) are not blocked for the duration of the downloads. c.Image.Caching = func() { c.Lock() - defer c.Unlock() + var queue = make([]string, len(c.Queue)) + copy(queue, c.Queue) + c.Unlock() - var filename string + for _, src := range queue { - for _, src := range c.Queue { - - resp, err := http.Get(src) + filename, err := c.download(src) if err != nil { - continue - } - defer resp.Body.Close() - - if resp.StatusCode != http.StatusOK { + // Stays in the queue and is retried on the next run. continue } - filename = fmt.Sprintf("%s%s%s%s", c.path, string(os.PathSeparator), strToMD5(src), filepath.Ext(src)) + c.Lock() + c.images[filename] = c.cacheURL + filename + c.Queue = removeStringFromSlice(src, c.Queue) + c.Unlock() - file, err := os.Create(filename) - if err != nil { - continue - } - - defer file.Close() - - _, err = io.Copy(file, resp.Body) - if err != nil { - continue - } - - u, err := url.Parse(src) - if err == nil { - c.images[fmt.Sprintf("%s%s", strToMD5(src), filepath.Ext(u.Path))] = c.cacheURL + filename - } - - queue = append(queue, src) - - } - - for _, q := range queue { - c.Queue = removeStringFromSlice(q, c.Queue) } } @@ -175,3 +159,45 @@ func New(path, chacheURL string, caching bool) (c *Cache, err error) { return } + +// download : Fetches src into the cache folder and returns the cached file +// name (md5 of the URL plus the extension of the URL path). Must be called +// without holding the lock. +func (c *Cache) download(src string) (filename string, err error) { + + u, err := url.Parse(src) + if err != nil { + return + } + + filename = fmt.Sprintf("%s%s", strToMD5(src), filepath.Ext(u.Path)) + + resp, err := httpClient.Get(src) + if err != nil { + return + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + err = fmt.Errorf("%s: %s", src, resp.Status) + return + } + + var path = filepath.Join(c.path, filename) + + file, err := os.Create(path) + if err != nil { + return + } + + _, err = io.Copy(file, resp.Body) + if closeErr := file.Close(); err == nil { + err = closeErr + } + + if err != nil { + os.Remove(path) + } + + return +} diff --git a/src/plex_api.go b/src/plex_api.go index e1169f7..2784ec8 100644 --- a/src/plex_api.go +++ b/src/plex_api.go @@ -342,11 +342,7 @@ func doPlexRequest(method, baseURL, endpoint, token string) (status int, body [] req.Header.Set("X-Plex-Version", System.Version) } - client := &http.Client{ - Timeout: 10 * time.Second, - } - - resp, err := client.Do(req) + resp, err := apiHTTPClient.Do(req) if err != nil { return } diff --git a/src/provider.go b/src/provider.go index d7cf508..b5f3cad 100644 --- a/src/provider.go +++ b/src/provider.go @@ -142,8 +142,16 @@ func getProviderData(fileType, fileID string) (err error) { for dataID, d := range dataMap { - var data = d.(map[string]any) - var fileSource = data["file.source"].(string) + data, ok := d.(map[string]any) + if !ok { + continue + } + + fileSource, ok := data["file.source"].(string) + if !ok { + continue + } + newProvider = false if _, ok := data["new"]; ok { @@ -224,8 +232,8 @@ func getProviderData(fileType, fileID string) (err error) { if value, ok := dataMap[dataID].(map[string]any); ok { data = value - data["counter.error"] = data["counter.error"].(float64) + 1 - data["counter.download"] = data["counter.download"].(float64) + 1 + data["counter.error"] = toFloat64(data["counter.error"]) + 1 + data["counter.download"] = toFloat64(data["counter.download"]) + 1 } @@ -243,10 +251,13 @@ func getProviderData(fileType, fileID string) (err error) { var data = make(map[string]any) data = value - if data["counter.error"].(float64) == 0 { + var errCount = toFloat64(data["counter.error"]) + var dlCount = toFloat64(data["counter.download"]) + + if errCount == 0 || dlCount == 0 { data["provider.availability"] = 100 } else { - data["provider.availability"] = int(data["counter.error"].(float64)*100/data["counter.download"].(float64)*-1 + 100) + data["provider.availability"] = int(errCount*100/dlCount*-1 + 100) } } @@ -289,8 +300,7 @@ func downloadFileFromServer(providerURL string) (filename string, body []byte, e req.Header.Set("User-Agent", getUserAgent()) - client := &http.Client{} - resp, err := client.Do(req) + resp, err := providerHTTPClient.Do(req) if err != nil { return } @@ -307,15 +317,22 @@ func downloadFileFromServer(providerURL string) (filename string, body []byte, e if index > -1 { - var headerFilename = resp.Header.Get("Content-Disposition")[index:len(resp.Header.Get("Content-Disposition"))] - var value = strings.Split(headerFilename, `=`) - var f = strings.Replace(value[1], `"`, "", -1) + var headerFilename = resp.Header.Get("Content-Disposition")[index:] + var value = strings.SplitN(headerFilename, `=`, 2) - f = strings.Replace(f, `;`, "", -1) - filename = f - showInfo("Header filename:" + filename) + if len(value) == 2 { - } else { + var f = strings.Replace(value[1], `"`, "", -1) + + f = strings.Replace(f, `;`, "", -1) + filename = f + showInfo("Header filename:" + filename) + + } + + } + + if len(filename) == 0 { var cleanFilename = strings.SplitN(getFilenameFromPath(providerURL), "?", 2) filename = cleanFilename[0] diff --git a/src/screen.go b/src/screen.go index 55c6ec7..cc04ddc 100644 --- a/src/screen.go +++ b/src/screen.go @@ -10,6 +10,82 @@ import ( "time" ) +// logMu guards WebScreenLog (its slice and counters) and System.Notification. +// Log lines are appended from every goroutine (buffer, maintenance, web +// handlers) and read by the websocket handler. +var logMu sync.Mutex + +// logAppend : Appends a timestamped line to the in-memory log and trims it +// to Settings.LogEntriesRAM. +func logAppend(logMsg string) { + + logMu.Lock() + defer logMu.Unlock() + + WebScreenLog.Log = append(WebScreenLog.Log, time.Now().Format("2006-01-02 15:04:05")+" "+logMsg) + logCleanUpLocked() +} + +// logAppendCounted : Appends a warning / error line and bumps the matching +// counter. Like the original code this does not trim the log; the next +// info / debug line does. +func logAppendCounted(logMsg string, isError bool) { + + logMu.Lock() + defer logMu.Unlock() + + WebScreenLog.Log = append(WebScreenLog.Log, time.Now().Format("2006-01-02 15:04:05")+" "+logMsg) + + if isError { + WebScreenLog.Errors++ + } else { + WebScreenLog.Warnings++ + } +} + +// resetWebScreenLog : Clears the in-memory log and its counters. +func resetWebScreenLog() { + + logMu.Lock() + defer logMu.Unlock() + + WebScreenLog.Log = make([]string, 0) + WebScreenLog.Errors = 0 + WebScreenLog.Warnings = 0 +} + +// webScreenLogSnapshot : Copy of the log for the web client. The slice is +// copied so the JSON encoder never races with a concurrent append. +func webScreenLogSnapshot() (snapshot WebScreenLogStruct) { + + logMu.Lock() + defer logMu.Unlock() + + snapshot = WebScreenLog + snapshot.Log = make([]string, len(WebScreenLog.Log)) + copy(snapshot.Log, WebScreenLog.Log) + + return +} + +// notificationsSnapshot : Copy of the notifications map for the web client. +func notificationsSnapshot() (snapshot map[string]Notification) { + + logMu.Lock() + defer logMu.Unlock() + + if System.Notification == nil { + return nil + } + + snapshot = make(map[string]Notification, len(System.Notification)) + for k, v := range System.Notification { + snapshot[k] = v + } + + return +} + func showInfo(str string) { if System.Flag.Info { @@ -33,9 +109,7 @@ func showInfo(str string) { printLogOnScreen(logMsg, "info") - logMsg = strings.Replace(logMsg, " ", " ", -1) - WebScreenLog.Log = append(WebScreenLog.Log, time.Now().Format("2006-01-02 15:04:05")+" "+logMsg) - logCleanUp() + logAppend(strings.Replace(logMsg, " ", " ", -1)) } @@ -51,7 +125,6 @@ func showDebug(str string, level int) { var msg = strings.SplitN(str, ":", 2) var length = len(msg[0]) var space string - var mutex = sync.RWMutex{} if len(msg) == 2 { @@ -64,11 +137,7 @@ func showDebug(str string, level int) { printLogOnScreen(logMsg, "debug") - mutex.Lock() - logMsg = strings.Replace(logMsg, " ", " ", -1) - WebScreenLog.Log = append(WebScreenLog.Log, time.Now().Format("2006-01-02 15:04:05")+" "+logMsg) - logCleanUp() - mutex.Unlock() + logAppend(strings.Replace(logMsg, " ", " ", -1)) } @@ -99,7 +168,8 @@ func showHighlight(str string) { } notification.Type = "info" - notification.Message = msg[1] + // Messages without a "key:value" prefix are shown as they are. + notification.Message = msg[len(msg)-1] addNotification(notification) @@ -109,31 +179,22 @@ func showWarning(errCode int) { var errMsg = getErrMsg(errCode) var logMsg = fmt.Sprintf("[%s] [WARNING] %s", System.Name, errMsg) - var mutex = sync.RWMutex{} printLogOnScreen(logMsg, "warning") - mutex.Lock() - WebScreenLog.Log = append(WebScreenLog.Log, time.Now().Format("2006-01-02 15:04:05")+" "+logMsg) - WebScreenLog.Warnings++ - mutex.Unlock() + logAppendCounted(logMsg, false) } // ShowError : Zeigt die Fehlermeldungen in der Konsole func ShowError(err error, errCode int) { - var mutex = sync.RWMutex{} - var errMsg = getErrMsg(errCode) var logMsg = fmt.Sprintf("[%s] [ERROR] %s (%s) - EC: %d", System.Name, err, errMsg, errCode) printLogOnScreen(logMsg, "error") - mutex.Lock() - WebScreenLog.Log = append(WebScreenLog.Log, time.Now().Format("2006-01-02 15:04:05")+" "+logMsg) - WebScreenLog.Errors++ - mutex.Unlock() + logAppendCounted(logMsg, true) } @@ -174,24 +235,17 @@ func printLogOnScreen(logMsg string, logType string) { } -func logCleanUp() { +// logCleanUpLocked : Trims the log to the newest Settings.LogEntriesRAM +// lines and recounts warnings / errors. The caller must hold logMu. +func logCleanUpLocked() { + + WebScreenLog.Log = keepLastEntries(WebScreenLog.Log, Settings.LogEntriesRAM) - var logEntriesRAM = Settings.LogEntriesRAM var logs = WebScreenLog.Log WebScreenLog.Warnings = 0 WebScreenLog.Errors = 0 - if len(logs) > logEntriesRAM { - - var tmp = make([]string, 0) - for i := len(logs) - logEntriesRAM; i < logEntriesRAM; i++ { - tmp = append(tmp, logs[i]) - } - - logs = tmp - } - for _, log := range logs { if strings.Contains(log, "WARNING") { @@ -204,8 +258,25 @@ func logCleanUp() { } - WebScreenLog.Log = logs +} +// keepLastEntries : Returns the last max entries of logs (a copy, so the +// old backing array does not keep growing). A max <= 0 keeps nothing, which +// matches the original behaviour for an unset LogEntriesRAM. +func keepLastEntries(logs []string, max int) []string { + + if max < 0 { + max = 0 + } + + if len(logs) <= max { + return logs + } + + var tmp = make([]string, max) + copy(tmp, logs[len(logs)-max:]) + + return tmp } // Fehlercodes @@ -386,30 +457,64 @@ func getErrMsg(errCode int) (errMsg string) { return errMsg } +// maxNotifications : Number of notifications kept for the web interface. +const maxNotifications = 10 + func addNotification(notification Notification) (err error) { - var i int var t = time.Now().UnixNano() / (int64(time.Millisecond) / int64(time.Nanosecond)) notification.Time = strconv.FormatInt(t, 10) + + insertNotification(notification) + + return +} + +// insertNotification : Stores the notification under its Time key and evicts +// the oldest ones once more than maxNotifications are kept. +func insertNotification(notification Notification) { + notification.New = true if len(notification.Headline) == 0 { notification.Headline = strings.ToUpper(notification.Type) } + logMu.Lock() + defer logMu.Unlock() + if len(System.Notification) == 0 { System.Notification = make(map[string]Notification) } System.Notification[notification.Time] = notification - for key := range System.Notification { + for len(System.Notification) > maxNotifications { + delete(System.Notification, oldestNotificationKey(System.Notification)) + } - if i < len(System.Notification)-10 { - delete(System.Notification, key) +} + +// oldestNotificationKey : Key of the notification with the smallest +// timestamp. Keys are millisecond timestamps as decimal strings; anything +// that does not parse sorts before everything that does. +func oldestNotificationKey(notifications map[string]Notification) (oldest string) { + + var oldestTime int64 + var found bool + + for key, n := range notifications { + + t, err := strconv.ParseInt(n.Time, 10, 64) + if err != nil { + return key } - i++ + if !found || t < oldestTime { + oldest = key + oldestTime = t + found = true + } } diff --git a/src/screen_test.go b/src/screen_test.go new file mode 100644 index 0000000..1b85c3c --- /dev/null +++ b/src/screen_test.go @@ -0,0 +1,108 @@ +package src + +import ( + "fmt" + "reflect" + "strconv" + "strings" + "testing" +) + +func TestKeepLastEntries(t *testing.T) { + logs := []string{"1", "2", "3", "4", "5"} + + got := keepLastEntries(logs, 3) + if want := []string{"3", "4", "5"}; !reflect.DeepEqual(got, want) { + t.Errorf("keepLastEntries(5 entries, 3) = %v, want %v", got, want) + } + + if got := keepLastEntries(logs, 10); !reflect.DeepEqual(got, logs) { + t.Errorf("keepLastEntries below the limit = %v, want %v", got, logs) + } + + if got := keepLastEntries(logs, 0); len(got) != 0 { + t.Errorf("keepLastEntries(.., 0) = %v, want empty", got) + } +} + +// The ring buffer must keep the newest lines once the log is full; the old +// loop bound dropped them. +func TestLogCleanUpKeepsNewest(t *testing.T) { + oldLimit := Settings.LogEntriesRAM + resetWebScreenLog() + t.Cleanup(func() { + Settings.LogEntriesRAM = oldLimit + resetWebScreenLog() + }) + + Settings.LogEntriesRAM = 3 + + for i := 1; i <= 5; i++ { + logAppend(fmt.Sprintf("line-%d", i)) + } + + snapshot := webScreenLogSnapshot() + if len(snapshot.Log) != 3 { + t.Fatalf("log has %d entries, want 3: %v", len(snapshot.Log), snapshot.Log) + } + for i, want := range []string{"line-3", "line-4", "line-5"} { + if !strings.HasSuffix(snapshot.Log[i], want) { + t.Errorf("entry %d = %q, want suffix %q", i, snapshot.Log[i], want) + } + } + + // Counters are recomputed from the retained lines. + logAppendCounted("[ERROR] boom", true) + logAppend("[WARNING] careful") + snapshot = webScreenLogSnapshot() + if snapshot.Errors != 1 || snapshot.Warnings != 1 { + t.Errorf("errors/warnings = %d/%d, want 1/1 (%v)", snapshot.Errors, snapshot.Warnings, snapshot.Log) + } +} + +func TestInsertNotificationEvictsOldest(t *testing.T) { + old := System.Notification + System.Notification = nil + t.Cleanup(func() { System.Notification = old }) + + // Insert out of order so eviction cannot rely on insertion order. + order := []int64{5, 1, 12, 3, 9, 2, 11, 7, 4, 10, 8, 6} + for _, ts := range order { + insertNotification(Notification{Type: "info", Message: "n", Time: strconv.FormatInt(1000+ts, 10)}) + } + + got := notificationsSnapshot() + if len(got) != maxNotifications { + t.Fatalf("kept %d notifications, want %d", len(got), maxNotifications) + } + for _, ts := range []int64{1, 2} { + if _, ok := got[strconv.FormatInt(1000+ts, 10)]; ok { + t.Errorf("oldest notification %d should have been evicted", ts) + } + } + for ts := int64(3); ts <= 12; ts++ { + if _, ok := got[strconv.FormatInt(1000+ts, 10)]; !ok { + t.Errorf("notification %d missing", ts) + } + } + + // The web-client copy must be independent of the live map. + delete(got, "1012") + if _, ok := notificationsSnapshot()["1012"]; !ok { + t.Error("snapshot is not a copy") + } +} + +func TestShowHighlightWithoutColon(t *testing.T) { + old := System.Notification + System.Notification = nil + t.Cleanup(func() { System.Notification = old }) + + showHighlight("no separator here") // used to index msg[1] and panic + + for _, n := range notificationsSnapshot() { + if n.Message != "no separator here" { + t.Errorf("message = %q", n.Message) + } + } +} diff --git a/src/security.go b/src/security.go index 4c31b46..927b97e 100644 --- a/src/security.go +++ b/src/security.go @@ -124,12 +124,7 @@ func maskSettings(s SettingsStruct) SettingsStruct { // writePrivateFile : like writeByteToFile but readable only by the owner. func writePrivateFile(file string, data []byte) error { - var filename = getPlatformFile(file) - - if err := os.WriteFile(filename, data, 0600); err != nil { - return err - } - - // WriteFile keeps the mode of an existing file; tighten it. - return os.Chmod(filename, 0600) + // The temp file is chmod'ed before the rename, so an existing + // world-readable file is replaced by a 0600 one. + return writeFileAtomic(getPlatformFile(file), data, 0600) } diff --git a/src/ssdp.go b/src/ssdp.go index 85e38f2..4dd95a3 100644 --- a/src/ssdp.go +++ b/src/ssdp.go @@ -4,12 +4,38 @@ import ( "fmt" "log" "os" - "os/signal" + "sync" "time" "github.com/koron/go-ssdp" ) +// ssdpState : The running advertiser, so Shutdown can say goodbye to the +// network before the process exits. +var ssdpState struct { + sync.Mutex + adv *ssdp.Advertiser + done chan struct{} +} + +// Shutdown : Stops background services (currently the SSDP advertiser). +// Called from main on SIGINT / SIGTERM. +func Shutdown() { + + ssdpState.Lock() + defer ssdpState.Unlock() + + if ssdpState.adv == nil { + return + } + + close(ssdpState.done) + ssdpState.adv.Bye() + ssdpState.adv.Close() + ssdpState.adv = nil + +} + // SSDP : SSPD / DLNA Server func SSDP() (err error) { @@ -19,9 +45,6 @@ func SSDP() (err error) { showInfo(fmt.Sprintf("SSDP / DLNA:%t", Settings.SSDP)) - quit := make(chan os.Signal, 1) - signal.Notify(quit, os.Interrupt) - ad, err := ssdp.Advertise( "upnp:rootdevice", // send as "ST" fmt.Sprintf("uuid:%s::upnp:rootdevice", System.DeviceID), // send as "USN" @@ -38,35 +61,43 @@ func SSDP() (err error) { ssdp.Logger = log.New(os.Stderr, "[SSDP] ", log.LstdFlags) } - go func(adv *ssdp.Advertiser) { + var done = make(chan struct{}) - aliveTick := time.Tick(300 * time.Second) + ssdpState.Lock() + ssdpState.adv = ad + ssdpState.done = done + ssdpState.Unlock() + + go func(adv *ssdp.Advertiser, done chan struct{}) { + + aliveTick := time.NewTicker(300 * time.Second) + defer aliveTick.Stop() - loop: for { select { - case <-aliveTick: - err = adv.Alive() - if err != nil { - ShowError(err, 0) + case <-aliveTick.C: + if aliveErr := adv.Alive(); aliveErr != nil { + ShowError(aliveErr, 0) + ssdpState.Lock() + if ssdpState.adv == adv { + ssdpState.adv = nil + } + ssdpState.Unlock() adv.Bye() adv.Close() - break loop + return } - case <-quit: - adv.Bye() - adv.Close() - os.Exit(0) - break loop + case <-done: + return } } - }(ad) + }(ad, done) return } diff --git a/src/toolchain.go b/src/toolchain.go index baffa25..14db0e3 100644 --- a/src/toolchain.go +++ b/src/toolchain.go @@ -204,12 +204,7 @@ func saveMapToJSONFile(file string, tmpMap any) error { return err } - err = os.WriteFile(filename, []byte(jsonString), 0644) - if err != nil { - return err - } - - return nil + return writeFileAtomic(filename, []byte(jsonString), 0644) } func loadJSONFileToMap(file string) (tmpMap map[string]any, err error) { @@ -229,6 +224,37 @@ func loadJSONFileToMap(file string) (tmpMap map[string]any, err error) { return } +// toFloat64 : Numeric value of a JSON decoded field; 0 for anything that +// is not a number (missing key, wrong type). +func toFloat64(v any) float64 { + + switch n := v.(type) { + case float64: + return n + case int: + return float64(n) + case int64: + return float64(n) + } + + return 0 +} + +// removeStrings : Copy of list without the entries present in remove. +// Order of the remaining entries is kept. +func removeStrings(list []string, remove map[string]bool) []string { + + var result = make([]string, 0, len(list)) + + for _, s := range list { + if !remove[s] { + result = append(result, s) + } + } + + return result +} + // Binary func readByteFromFile(file string) (content []byte, err error) { @@ -246,7 +272,48 @@ func readByteFromFile(file string) (content []byte, err error) { func writeByteToFile(file string, data []byte) (err error) { var filename = getPlatformFile(file) - err = os.WriteFile(filename, data, 0644) + err = writeFileAtomic(filename, data, 0644) + + return +} + +// writeFileAtomic : Writes data to a temporary file in the target directory, +// fsyncs it and renames it over name, so readers (and a crash mid-write) +// never see a truncated or half-written file. The temp file is removed on +// any error. +func writeFileAtomic(name string, data []byte, perm os.FileMode) (err error) { + + tmp, err := os.CreateTemp(filepath.Dir(name), ".tmp-*") + if err != nil { + return + } + + var tmpName = tmp.Name() + + defer func() { + if err != nil { + tmp.Close() + os.Remove(tmpName) + } + }() + + if _, err = tmp.Write(data); err != nil { + return + } + + if err = tmp.Sync(); err != nil { + return + } + + if err = tmp.Close(); err != nil { + return + } + + if err = os.Chmod(tmpName, perm); err != nil { + return + } + + err = os.Rename(tmpName, name) return } diff --git a/src/toolchain_test.go b/src/toolchain_test.go new file mode 100644 index 0000000..4deddc4 --- /dev/null +++ b/src/toolchain_test.go @@ -0,0 +1,86 @@ +package src + +import ( + "os" + "path/filepath" + "reflect" + "testing" +) + +func TestWriteFileAtomic(t *testing.T) { + dir := t.TempDir() + file := filepath.Join(dir, "settings.json") + + if err := writeFileAtomic(file, []byte("first"), 0644); err != nil { + t.Fatal(err) + } + if err := writeFileAtomic(file, []byte("second"), 0600); err != nil { + t.Fatal(err) + } + + got, err := os.ReadFile(file) + if err != nil { + t.Fatal(err) + } + if string(got) != "second" { + t.Errorf("content = %q, want %q", got, "second") + } + + fi, err := os.Stat(file) + if err != nil { + t.Fatal(err) + } + if fi.Mode().Perm() != 0600 { + t.Errorf("mode = %o, want 0600", fi.Mode().Perm()) + } + + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatal(err) + } + for _, e := range entries { + if e.Name() != "settings.json" { + t.Errorf("unexpected leftover file %q", e.Name()) + } + } +} + +func TestWriteFileAtomicMissingDir(t *testing.T) { + file := filepath.Join(t.TempDir(), "missing", "settings.json") + if err := writeFileAtomic(file, []byte("x"), 0644); err == nil { + t.Fatal("expected an error for a missing directory") + } +} + +func TestRemoveStrings(t *testing.T) { + list := []string{"a.xml", "b.xml", "c.xml", "d.xml"} + got := removeStrings(list, map[string]bool{"b.xml": true, "d.xml": true}) + want := []string{"a.xml", "c.xml"} + if !reflect.DeepEqual(got, want) { + t.Errorf("removeStrings = %v, want %v", got, want) + } + + // Nothing to remove: same content, original untouched. + got = removeStrings(list, nil) + if !reflect.DeepEqual(got, list) { + t.Errorf("removeStrings(nil) = %v, want %v", got, list) + } +} + +func TestToFloat64(t *testing.T) { + cases := []struct { + in any + want float64 + }{ + {float64(3), 3}, + {int(2), 2}, + {int64(1), 1}, + {"x", 0}, + {nil, 0}, + } + for _, c := range cases { + if got := toFloat64(c.in); got != c.want { + t.Errorf("toFloat64(%v) = %v, want %v", c.in, got, c.want) + } + } +} diff --git a/src/webserver.go b/src/webserver.go index 65bab81..a302603 100644 --- a/src/webserver.go +++ b/src/webserver.go @@ -9,6 +9,7 @@ import ( "os" "strconv" "strings" + "time" "xteve/src/internal/authentication" @@ -73,7 +74,17 @@ func StartWebserver() (err error) { } - if err = http.ListenAndServe(":"+port, nil); err != nil { + // No ReadTimeout / WriteTimeout on purpose: /stream/ responses are + // long-lived. ReadHeaderTimeout still protects against slowloris + // clients and IdleTimeout reaps keep-alive connections. + var server = &http.Server{ + Addr: ":" + port, + Handler: nil, + ReadHeaderTimeout: 15 * time.Second, + IdleTimeout: 120 * time.Second, + } + + if err = server.ListenAndServe(); err != nil { ShowError(err, 1001) return } @@ -317,10 +328,6 @@ func DataImages(w http.ResponseWriter, r *http.Request) { // WS : Web Sockets /ws/ func WS(w http.ResponseWriter, r *http.Request) { - var request RequestStruct - var response ResponseStruct - response.Status = true - var newToken string /* @@ -341,6 +348,12 @@ func WS(w http.ResponseWriter, r *http.Request) { for { + // Fresh structs per command: a failed command must not leak its + // Status / Error / Base64 into the next one on the same connection. + var request RequestStruct + var response ResponseStruct + response.Status = true + err = conn.ReadJSON(&request) if err != nil { @@ -465,9 +478,7 @@ func WS(w http.ResponseWriter, r *http.Request) { } case "resetLogs": - WebScreenLog.Log = make([]string, 0) - WebScreenLog.Errors = 0 - WebScreenLog.Warnings = 0 + resetWebScreenLog() response.OpenMenu = strconv.Itoa(indexOfString("log", System.WEB.Menu)) case "xteveBackup": @@ -478,9 +489,7 @@ func WS(w http.ResponseWriter, r *http.Request) { } case "xteveRestore": - WebScreenLog.Log = make([]string, 0) - WebScreenLog.Errors = 0 - WebScreenLog.Warnings = 0 + resetWebScreenLog() if len(request.Base64) > 0 { @@ -855,8 +864,8 @@ func API(w http.ResponseWriter, r *http.Request) { default: token, err = tokenAuthentication(request.Token) - fmt.Println(err) if err != nil { + ShowError(err, 0) responseAPIError(err) return } @@ -926,6 +935,7 @@ func API(w http.ResponseWriter, r *http.Request) { if err != nil { responseAPIError(err) + return } w.Write([]byte(mapToJSON(response))) @@ -977,10 +987,10 @@ func setDefaultResponseData(response ResponseStruct, data bool) (defaults Respon defaults.ClientInfo.OS = System.OS defaults.ClientInfo.Streams = fmt.Sprintf("%d / %d", len(Data.Streams.Active), len(Data.Streams.All)) defaults.ClientInfo.UUID = Settings.UUID - defaults.ClientInfo.Errors = WebScreenLog.Errors - defaults.ClientInfo.Warnings = WebScreenLog.Warnings - defaults.Notification = System.Notification - defaults.Log = WebScreenLog + defaults.Log = webScreenLogSnapshot() + defaults.ClientInfo.Errors = defaults.Log.Errors + defaults.ClientInfo.Warnings = defaults.Log.Warnings + defaults.Notification = notificationsSnapshot() switch System.Branch { diff --git a/src/xepg.go b/src/xepg.go index ccdffeb..8d40ce6 100644 --- a/src/xepg.go +++ b/src/xepg.go @@ -172,6 +172,11 @@ func createXEPGMapping() { if len(Data.XMLTV.Files) > 0 { + // Files that could not be read are collected here and removed after + // the loop; mutating Data.XMLTV.Files while iterating over it used + // to duplicate entries instead of removing them. + var failedFiles = make(map[string]bool) + for i := len(Data.XMLTV.Files) - 1; i >= 0; i-- { var file = Data.XMLTV.Files[i] @@ -185,7 +190,7 @@ func createXEPGMapping() { err = getLocalXMLTV(file, &xmltv) if err != nil { - Data.XMLTV.Files = append(Data.XMLTV.Files, Data.XMLTV.Files[i+1:]...) + failedFiles[file] = true var errMsg = err.Error() err = errors.New(getProviderParameter(fileID, "xmltv", "name") + ": " + errMsg) ShowError(err, 000) @@ -215,6 +220,7 @@ func createXEPGMapping() { } + Data.XMLTV.Files = removeStrings(Data.XMLTV.Files, failedFiles) Data.XMLTV.Mapping = tmpMap } else { diff --git a/xteve.go b/xteve.go index 6b28c62..3877b93 100644 --- a/xteve.go +++ b/xteve.go @@ -9,9 +9,11 @@ import ( "flag" "fmt" "os" + "os/signal" "path/filepath" "runtime" "strings" + "syscall" "xteve/src" ) @@ -111,7 +113,7 @@ func main() { err := src.Init() if err != nil { src.ShowError(err, 0) - os.Exit(0) + os.Exit(1) } src.ShowSystemInfo() @@ -150,12 +152,13 @@ func main() { err := src.Init() if err != nil { src.ShowError(err, 0) - os.Exit(0) + os.Exit(1) } err = src.XteveRestoreFromCLI(*restore) if err != nil { src.ShowError(err, 0) + os.Exit(1) } os.Exit(0) @@ -164,25 +167,35 @@ func main() { err := src.Init() if err != nil { src.ShowError(err, 0) - os.Exit(0) + os.Exit(1) } err = src.StartSystem(false) if err != nil { src.ShowError(err, 0) - os.Exit(0) + os.Exit(1) } + // Graceful stop on SIGINT / SIGTERM: say goodbye on the network (SSDP) + // and exit cleanly instead of letting the runtime kill the process. + go func() { + quit := make(chan os.Signal, 1) + signal.Notify(quit, os.Interrupt, syscall.SIGTERM) + <-quit + src.Shutdown() + os.Exit(0) + }() + err = src.InitMaintenance() if err != nil { src.ShowError(err, 0) - os.Exit(0) + os.Exit(1) } err = src.StartWebserver() if err != nil { src.ShowError(err, 0) - os.Exit(0) + os.Exit(1) } }