Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions cmd/ember/app.go
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@ type App struct {
viewWaiters viewWaiters
viewRecheck time.Duration
viewWaitHook func()
otaReadHook func()
knobStats *knobStatsStore
wifiDrops checkinDropLog
diagDrops checkinDropLog
Expand Down
1 change: 1 addition & 0 deletions cmd/ember/devices.go
Original file line number Diff line number Diff line change
Expand Up @@ -197,6 +197,7 @@ type deviceRegistry struct {
loadErr error
persistedAt time.Time
dirty bool
fwGen uint64

// Runs under mu: it must not block or call back into the registry.
onChange func()
Expand Down
88 changes: 66 additions & 22 deletions cmd/ember/devices_ota.go
Original file line number Diff line number Diff line change
Expand Up @@ -34,13 +34,15 @@ const (
otaAutoNotStarted = 24 * time.Hour
otaBlockedMax = 16
otaWriteDeadline = 10 * time.Minute
otaPutTries = 3
)

var (
errNoRollback = errors.New("no_rollback_bootloader")
errOTAInProgress = errors.New("ota_in_progress")
errOTANotOffered = errors.New("this firmware version is not offered to this device")
errOTAUnknownImage = fmt.Errorf("%w: unknown firmware version", errDeviceBody)
errFirmwareChanged = errors.New("firmware_changed")

otaErrorPattern = regexp.MustCompile(`^[a-z0-9_]{1,24}$`)
)
Expand Down Expand Up @@ -328,50 +330,92 @@ func (r *deviceRegistry) updateOTA(id string, fn func(o *knobOTA, last *deviceCh
return o, last, err
}

func (o knobOTA) parked(version string) bool {
return o.Version == version && (o.Phase == otaPhaseFailed || o.Phase == otaPhaseRolledBack)
}

func (o knobOTA) holds(version string) bool {
return (o.Target == version && !o.parked(version)) || (o.Version == version && otaActive(o.Phase))
}

func (r *deviceRegistry) firmwareGen() uint64 {
gen, _ := r.otaGens()
return gen
}

func (r *deviceRegistry) otaGens() (uint64, uint64) {
r.mu.Lock()
defer r.mu.Unlock()
return r.fwGen, r.state.Epoch
}

func (r *deviceRegistry) otaTargets(version string) bool {
r.mu.Lock()
defer r.mu.Unlock()
for _, d := range r.state.Devices {
if o := d.OTA; o != nil && (o.Target == version || (o.Version == version && otaActive(o.Phase))) {
if d.OTA != nil && d.OTA.holds(version) {
return true
}
}
r.fwGen++
return false
}

func (r *deviceRegistry) otaKeeps(version string) bool {
r.mu.Lock()
defer r.mu.Unlock()
for _, d := range r.state.Devices {
if (d.OTA != nil && (d.OTA.holds(version) || d.OTA.Target == version)) ||
(d.LastCheckin != nil && d.LastCheckin.FW == version) {
return true
}
}
r.fwGen++
return false
}

func (r *deviceRegistry) otaUnblock(versions []string) error {
func (r *deviceRegistry) otaRetire(versions []string, stored string) error {
r.mu.Lock()
defer r.mu.Unlock()
gone := func(v string) bool { return slices.Contains(versions, v) }
if !slices.ContainsFunc(r.state.Devices, func(d deviceRecord) bool {
return d.OTA != nil && slices.ContainsFunc(d.OTA.Blocked, gone)
}) {
r.fwGen++
unblock := func(v string) bool { return slices.Contains(versions, v) }
gone := func(v string) bool { return v != "" && v != stored && unblock(v) }
stale := func(o *knobOTA) bool {
if o == nil {
return false
}
idle := !otaActive(o.Phase)
return slices.ContainsFunc(o.Blocked, unblock) ||
(idle && gone(o.Target)) || (idle && o.Phase != otaPhaseDone && gone(o.Version))
}
if !slices.ContainsFunc(r.state.Devices, func(d deviceRecord) bool { return stale(d.OTA) }) {
return nil
}
return r.mutateLocked(func(st *deviceState) error {
bump := false
for i := range st.Devices {
if o := st.Devices[i].OTA; o != nil {
o.Blocked = slices.DeleteFunc(o.Blocked, gone)
o := st.Devices[i].OTA
if !stale(o) {
continue
}
o.Blocked = slices.DeleteFunc(o.Blocked, unblock)
if otaActive(o.Phase) {
continue
}
if gone(o.Target) {
o.Target, o.Retry, bump = "", false, true
}
if o.Phase != otaPhaseDone && gone(o.Version) {
*o = knobOTA{Mode: o.Mode, Target: o.Target, Retry: o.Retry, Blocked: o.Blocked, Attempt: o.Attempt}
}
}
if bump {
st.Epoch++
}
return nil
})
}

func (r *deviceRegistry) otaKeeps(version string) bool {
if r.otaTargets(version) {
return true
}
r.mu.Lock()
defer r.mu.Unlock()
for _, d := range r.state.Devices {
if d.LastCheckin != nil && d.LastCheckin.FW == version {
return true
}
}
return false
}

type otaProgress struct {
version string
bytes int64
Expand Down
77 changes: 64 additions & 13 deletions cmd/ember/devices_ota_http.go
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ func (a *App) otaCheckin(id string, report deviceCheckin, res *checkinResult) er
if a.knobFW == nil {
return nil
}
gen, epoch := a.devices.otaGens()
snap, _, err := a.devices.otaSnapshot(id)
if err != nil {
return err
Expand All @@ -50,14 +51,21 @@ func (a *App) otaCheckin(id string, report deviceCheckin, res *checkinResult) er
in.auto = &m
}
}
if a.otaReadHook != nil {
a.otaReadHook()
}
var offer *otaOffer
var waiting string
_, _, err = a.devices.updateOTA(id, func(o *knobOTA, _ *deviceCheckin) (bool, error) {
offer, waiting = nil, ""
if o.Target != snap.Target || o.mode() != snap.mode() {
if o.Target != snap.Target || o.mode() != snap.mode() || o.Attempt != snap.Attempt || a.devices.fwGen != gen {
applyOTAResult(o, report, in.now)
return false, nil
}
if o.Target != "" && in.target == nil && !otaActive(o.Phase) && a.devices.state.Epoch == epoch {
o.Target, o.Retry = "", false
return true, nil
}
offer, waiting = stepOTA(o, in)
if offer == nil || o.Phase != otaPhaseOffered {
a.ota.notOffered(id)
Expand Down Expand Up @@ -187,23 +195,25 @@ func (a *App) handleDeviceOTAPut(w http.ResponseWriter, r *http.Request) {
a.writeDeviceError(w, r, fmt.Errorf("%w: target must be a version string or null", errDeviceBody))
return
}
if a.knobFW == nil {
writeError(w, http.StatusServiceUnavailable, errFirmwareOff)
return
}
if _, ok := a.knobFW.get(v); !ok {
a.writeDeviceError(w, r, errOTAUnknownImage)
return
}
target = &v
}
retry := req.Retry != nil && *req.Retry
if (target != nil || retry) && a.knobFW == nil {
writeError(w, http.StatusServiceUnavailable, errFirmwareOff)
return
}
id := r.PathValue("id")
o, last, err := a.devices.updateOTA(id, func(o *knobOTA, last *deviceCheckin) (bool, error) {
return applyOTAPut(o, last, req.Mode, target, clearTarget, retry)
})
var o knobOTA
var last *deviceCheckin
var err error
for range otaPutTries {
o, last, err = a.putOTA(id, req.Mode, target, clearTarget, retry)
if !errors.Is(err, errFirmwareChanged) {
break
}
}
switch {
case errors.Is(err, errNoRollback), errors.Is(err, errOTAInProgress):
case errors.Is(err, errNoRollback), errors.Is(err, errOTAInProgress), errors.Is(err, errFirmwareChanged):
writeError(w, http.StatusConflict, err)
return
case err != nil:
Expand All @@ -214,6 +224,41 @@ func (a *App) handleDeviceOTAPut(w http.ResponseWriter, r *http.Request) {
writeJSON(w, http.StatusOK, a.otaStatus(id, o, last))
}

func (a *App) putOTA(id string, mode, target *string, clearTarget, retry bool) (knobOTA, *deviceCheckin, error) {
gen := a.devices.firmwareGen()
snap, _, err := a.devices.otaSnapshot(id)
if err != nil {
return knobOTA{}, nil, err
}
want := otaWants(snap, target, retry)
if want != "" {
if _, ok := a.knobFW.get(want); !ok {
return knobOTA{}, nil, errOTAUnknownImage
}
}
if a.otaReadHook != nil {
a.otaReadHook()
}
return a.devices.updateOTA(id, func(o *knobOTA, last *deviceCheckin) (bool, error) {
if otaWants(*o, target, retry) != want || a.devices.fwGen != gen {
return false, errFirmwareChanged
}
return applyOTAPut(o, last, mode, target, clearTarget, retry)
})
}

func otaWants(o knobOTA, target *string, retry bool) string {
switch {
case target != nil:
return *target
case !retry:
return ""
case o.Target != "":
return o.Target
}
return o.Version
}

func applyOTAPut(o *knobOTA, last *deviceCheckin, mode, target *string, clearTarget, retry bool) (bool, error) {
if (target != nil || retry) && (last == nil || last.OTA == nil || !last.OTA.Rollback) {
return false, errNoRollback
Expand All @@ -225,6 +270,9 @@ func applyOTAPut(o *knobOTA, last *deviceCheckin, mode, target *string, clearTar
if committed && (retry || (clearTarget && o.Target != "") || (target != nil && *target != o.Version)) {
return false, errOTAInProgress
}
if clearTarget && o.Target != "" && o.Phase == otaPhaseDownloading {
return false, errOTAInProgress
}
if target != nil && ((*target == o.Version && (o.Phase == otaPhaseFailed || o.Phase == otaPhaseRolledBack)) || slices.Contains(o.Blocked, *target)) {
retry = true
}
Expand All @@ -238,6 +286,9 @@ func applyOTAPut(o *knobOTA, last *deviceCheckin, mode, target *string, clearTar
o.Phase = otaPhaseIdle
}
}
if clearTarget && (o.Phase == otaPhaseFailed || o.Phase == otaPhaseRolledBack) {
o.Phase, o.Error = otaPhaseIdle, ""
}
if target != nil && (*target != o.Target || !otaActive(o.phase())) {
o.Target, bump = *target, true
if o.Version != *target || !otaActive(o.phase()) {
Expand Down
Loading
Loading