feat: statefile versioning

This commit is contained in:
Marco98
2025-11-17 17:24:18 +01:00
parent 5e39f26252
commit 898567679e
3 changed files with 72 additions and 21 deletions

View File

@@ -36,7 +36,7 @@ func (d *Daemon) getUserPass(ctx context.Context, username string) (string, erro
slog.InfoContext(ctx, "created new app password", "user", username)
d.userTokensLock.Lock()
d.userTokens[username] = pass
d.userTokensUnsaved = true
d.stateUnsaved = true
d.userTokensLock.Unlock()
return pass, nil
}

View File

@@ -24,13 +24,13 @@ var (
)
type Daemon struct {
userTokens map[string]string
userTokensLock *sync.RWMutex
userTokensUnsaved bool
httpClient *http.Client
baseURL string
mailcowClient mailcow.Client
statefile string
httpClient *http.Client
baseURL string
mailcowClient mailcow.Client
userTokens map[string]string
userTokensLock *sync.RWMutex
stateFilepath string
stateUnsaved bool
}
func main() {
@@ -46,22 +46,22 @@ func run() error {
userTokens: make(map[string]string),
userTokensLock: &sync.RWMutex{},
baseURL: os.Getenv("MAILCOW_BASE"),
statefile: os.Getenv("STATEFILE"),
stateFilepath: os.Getenv("STATEFILE"),
httpClient: &http.Client{
Transport: &http.Transport{
Proxy: http.ProxyFromEnvironment,
},
},
}
if len(d.statefile) == 0 {
d.statefile = "state.json"
if len(d.stateFilepath) == 0 {
d.stateFilepath = "state.json"
}
d.mailcowClient = mailcow.New(
d.httpClient,
d.baseURL,
os.Getenv("MAILCOW_APIKEY"),
)
if err := d.LoadFromDisk(); err != nil {
if err := d.loadState(); err != nil {
return err
}
d.daemonLoop()
@@ -92,12 +92,12 @@ func (d *Daemon) daemonRun() error {
})
}
eg.Wait()
if d.userTokensUnsaved {
if d.stateUnsaved {
slog.Info("saving tokens to disk", "count", len(d.userTokens))
if err := d.SaveToDisk(); err != nil {
if err := d.saveState(); err != nil {
return err
}
d.userTokensUnsaved = false
d.stateUnsaved = false
}
return nil
}

View File

@@ -1,13 +1,64 @@
package main
import (
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"log/slog"
"os"
)
func (d *Daemon) LoadFromDisk() error {
f, err := os.OpenFile(d.statefile, os.O_RDONLY, 0o660)
func (d *Daemon) loadState() error {
stateVer := struct {
Version int `json:"version"`
}{}
if err := d.loadFromDisk(&stateVer); err != nil {
return fmt.Errorf("cant detect state version: %w", err)
}
switch stateVer.Version {
case 0:
slog.Warn("loading old state version", "stateVer", stateVer.Version)
if err := d.loadFromDisk(&d.userTokens); err != nil {
return fmt.Errorf("cant load state v%d: %w", stateVer.Version, err)
}
d.stateUnsaved = true
case 1:
state := struct {
Version int `json:"version"`
UserTokens map[string]string `json:"userTokens"`
}{}
if err := d.loadFromDisk(&state); err != nil {
return fmt.Errorf("cant load state v%d: %w", stateVer.Version, err)
}
for k, v := range state.UserTokens {
dec, err := base64.StdEncoding.DecodeString(v)
if err != nil {
return fmt.Errorf("cant decode pass from %s: %w", k, err)
}
d.userTokens[k] = string(dec)
}
}
return nil
}
func (d *Daemon) saveState() error {
encTokens := make(map[string]string, len(d.userTokens))
for k, v := range d.userTokens {
encTokens[k] = base64.StdEncoding.EncodeToString([]byte(v))
}
state := struct {
Version int `json:"version"`
UserTokens map[string]string `json:"userTokens"`
}{
Version: 1,
UserTokens: encTokens,
}
return d.saveToDisk(state)
}
func (d *Daemon) loadFromDisk(state any) error {
f, err := os.OpenFile(d.stateFilepath, os.O_RDONLY, 0o660)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil
@@ -15,14 +66,14 @@ func (d *Daemon) LoadFromDisk() error {
return err
}
defer f.Close()
return json.NewDecoder(f).Decode(&d.userTokens)
return json.NewDecoder(f).Decode(state)
}
func (d *Daemon) SaveToDisk() error {
f, err := os.OpenFile(d.statefile, os.O_CREATE|os.O_WRONLY, 0o660)
func (d *Daemon) saveToDisk(state any) error {
f, err := os.OpenFile(d.stateFilepath, os.O_CREATE|os.O_WRONLY, 0o660)
if err != nil {
return err
}
defer f.Close()
return json.NewEncoder(f).Encode(d.userTokens)
return json.NewEncoder(f).Encode(state)
}