|
25 | 25 | package devicemgmtstate |
26 | 26 |
|
27 | 27 | import ( |
| 28 | + "encoding/json" |
28 | 29 | "errors" |
29 | 30 | "fmt" |
30 | 31 | "sort" |
|
58 | 59 | maxSequences = 256 |
59 | 60 | maxBlockedMessagesPerSequence = 8 |
60 | 61 |
|
| 62 | + awaitSubsystemRetryInterval = 30 * time.Second |
| 63 | + |
61 | 64 | deviceMgmtExchangeChangeKind = swfeats.RegisterChangeKind("device-management-exchange") |
62 | 65 | ) |
63 | 66 |
|
@@ -180,6 +183,23 @@ func (ms *deviceMgmtState) getRequestMessage(id string) (*RequestMessage, error) |
180 | 183 | return nil, fmt.Errorf("cannot find message %q", id) |
181 | 184 | } |
182 | 185 |
|
| 186 | +// removeRequestMessage removes a processed request message from its sequence, |
| 187 | +// leaving the sequence entry in place so its Applied progress is preserved for |
| 188 | +// later messages in the same sequence. |
| 189 | +func (ms *deviceMgmtState) removeRequestMessage(msg *RequestMessage) { |
| 190 | + seq := ms.Sequences[msg.BaseID] |
| 191 | + if seq == nil { |
| 192 | + return |
| 193 | + } |
| 194 | + |
| 195 | + for i, m := range seq.Messages { |
| 196 | + if m.SeqNum == msg.SeqNum { |
| 197 | + seq.Messages = append(seq.Messages[:i], seq.Messages[i+1:]...) |
| 198 | + return |
| 199 | + } |
| 200 | + } |
| 201 | +} |
| 202 | + |
183 | 203 | // enqueueRequestMessages queues incoming request messages for processing |
184 | 204 | // and updates polling state accordingly. |
185 | 205 | func (ms *deviceMgmtState) enqueueRequestMessages(pollResp *store.MessageExchangeResponse) { |
@@ -479,8 +499,6 @@ func (m *DeviceMgmtManager) dispatchSequence(dispatchTask *state.Task, seq *sequ |
479 | 499 | // the final task so callers can chain subsequent messages after it. |
480 | 500 | func (m *DeviceMgmtManager) dispatchMessage(prevTask *state.Task, msg *RequestMessage) *state.Task { |
481 | 501 | chg := prevTask.Change() |
482 | | - // TODO: add tests verifying that a failure in one message's task chain does not |
483 | | - // affect other messages (lanes provide this isolation, but it needs test coverage). |
484 | 502 | lane := m.state.NewLane() |
485 | 503 |
|
486 | 504 | addTask := func(kind, summary string) { |
@@ -586,9 +604,99 @@ func (m *DeviceMgmtManager) doApplyMessage(t *state.Task, _ *tomb.Tomb) error { |
586 | 604 | } |
587 | 605 |
|
588 | 606 | // doQueueResponse builds a response, signs it, and queues it for transmission on the next exchange. |
589 | | -// Retries until subsystem change completes. |
| 607 | +// Retries until the subsystem change (if any) completes. |
590 | 608 | func (m *DeviceMgmtManager) doQueueResponse(t *state.Task, _ *tomb.Tomb) error { |
591 | | - // TODO: implement this task, no-op for now. |
| 609 | + m.state.Lock() |
| 610 | + defer m.state.Unlock() |
| 611 | + |
| 612 | + ms, err := m.getState() |
| 613 | + if err != nil { |
| 614 | + return err |
| 615 | + } |
| 616 | + |
| 617 | + var msgID string |
| 618 | + err = t.Get("message-id", &msgID) |
| 619 | + if err != nil { |
| 620 | + return err |
| 621 | + } |
| 622 | + |
| 623 | + msg, err := ms.getRequestMessage(msgID) |
| 624 | + if err != nil { |
| 625 | + // Message already processed on a prior run. |
| 626 | + return nil |
| 627 | + } |
| 628 | + |
| 629 | + err = m.setMessageResponseFromChange(msg) |
| 630 | + if err != nil { |
| 631 | + return err |
| 632 | + } |
| 633 | + |
| 634 | + bodyBytes, err := json.Marshal(msg.ResponseBody) |
| 635 | + if err != nil { |
| 636 | + return fmt.Errorf("cannot marshal response body: %w", err) |
| 637 | + } |
| 638 | + |
| 639 | + // TODO: determine reasonable behavior for internal errors (e.g., signing or marshal failures). |
| 640 | + // Since tasks are idempotent, a failed message will not be re-dispatched or re-applied on the |
| 641 | + // next change, but the request message remains in state until doQueueResponse completes, |
| 642 | + // so such failures leave it hanging indefinitely. |
| 643 | + |
| 644 | + resAs, err := m.signer.SignResponseMessage(msg.AccountID, msg.ID(), msg.ResponseStatus, bodyBytes) |
| 645 | + if err != nil { |
| 646 | + return fmt.Errorf("cannot sign response message: %w", err) |
| 647 | + } |
| 648 | + |
| 649 | + ms.ReadyResponses[msg.ID()] = store.Message{ |
| 650 | + Format: "assertion", |
| 651 | + Data: string(asserts.Encode(resAs)), |
| 652 | + } |
| 653 | + |
| 654 | + // TODO: rejecting sequences currently happens in 2 ways: |
| 655 | + // 1. doDispatchMessage can evict the sequence immediately if it's rejected early. |
| 656 | + // 2. If it errors elsewhere (in validate, apply, or queue-response), we end |
| 657 | + // up not advancing Applied, which means we accumulate messages until we |
| 658 | + // hit the sequence cap. |
| 659 | + // Refactor sequence rejection to always evict immediately. |
| 660 | + if msg.SeqNum > 0 && msg.ResponseStatus == asserts.MessageStatusSuccess { |
| 661 | + ms.Sequences[msg.BaseID].Applied = msg.SeqNum |
| 662 | + } |
| 663 | + ms.removeRequestMessage(msg) |
| 664 | + |
| 665 | + m.setState(ms) |
| 666 | + |
| 667 | + return nil |
| 668 | +} |
| 669 | + |
| 670 | +// setMessageResponseFromChange populates msg's response fields from the completed apply change. |
| 671 | +func (m *DeviceMgmtManager) setMessageResponseFromChange(msg *RequestMessage) error { |
| 672 | + if msg.ResponseStatus != "" { |
| 673 | + return nil |
| 674 | + } |
| 675 | + |
| 676 | + handler, ok := m.handlers[msg.Kind] |
| 677 | + if !ok { |
| 678 | + msg.ResponseStatus = asserts.MessageStatusError |
| 679 | + msg.ResponseBody = map[string]any{"message": fmt.Sprintf("cannot find handler for message kind %q", msg.Kind)} |
| 680 | + return nil |
| 681 | + } |
| 682 | + |
| 683 | + change := m.state.Change(msg.ApplyChangeID) |
| 684 | + if change == nil { |
| 685 | + return fmt.Errorf("internal error: cannot find subsystem change %q", msg.ApplyChangeID) |
| 686 | + } |
| 687 | + if !change.Status().Ready() { |
| 688 | + return &state.Retry{After: awaitSubsystemRetryInterval} |
| 689 | + } |
| 690 | + |
| 691 | + body, err := handler.ResultFromChange(change) |
| 692 | + if err != nil { |
| 693 | + msg.ResponseStatus = asserts.MessageStatusError |
| 694 | + msg.ResponseBody = map[string]any{"message": fmt.Sprintf("cannot process message: %v", err)} |
| 695 | + } else { |
| 696 | + msg.ResponseStatus = asserts.MessageStatusSuccess |
| 697 | + msg.ResponseBody = body |
| 698 | + } |
| 699 | + |
592 | 700 | return nil |
593 | 701 | } |
594 | 702 |
|
|
0 commit comments