diff --git a/overlord/assertstate/assertstate.go b/overlord/assertstate/assertstate.go index b6eea88b10..ad0d82f1d4 100644 --- a/overlord/assertstate/assertstate.go +++ b/overlord/assertstate/assertstate.go @@ -1485,3 +1485,33 @@ func ValidatedIntegrityData(st *state.State, snapID string, rev snap.Revision) ( return integrity.NewIntegrityDataParamsFromRevision(revAssertion) } + +// AccountKey returns the account-key assertion for the given signing key ID, +// if it's present in the system assertion database. +func AccountKey(st *state.State, signKeyID string) (*asserts.AccountKey, error) { + db := DB(st) + as, err := db.Find(asserts.AccountKeyType, map[string]string{ + "public-key-sha3-384": signKeyID, + }) + if err != nil { + return nil, err + } + + return as.(*asserts.AccountKey), nil +} + +// FetchAccountKey fetches the account-key assertion for the given signing key ID. +func FetchAccountKey(st *state.State, userID int, signKeyID string) error { + deviceCtx, err := snapstate.DevicePastSeeding(st, nil) + if err != nil { + return err + } + + return doFetch(st, userID, deviceCtx, nil, func(f asserts.Fetcher) error { + ref := &asserts.Ref{ + Type: asserts.AccountKeyType, + PrimaryKey: []string{signKeyID}, + } + return f.Fetch(ref) + }) +} diff --git a/overlord/assertstate/assertstate_test.go b/overlord/assertstate/assertstate_test.go index 56c5cfcfee..9a7092e890 100644 --- a/overlord/assertstate/assertstate_test.go +++ b/overlord/assertstate/assertstate_test.go @@ -6315,3 +6315,53 @@ func (s *assertMgrSuite) TestValidatedIntegrityDataErrorNoRevisionsFound(c *C) { expErr: regexp.QuoteMeta("no snap-revision assertion found that matches (snap-id=snap-id-1, snap-revision=10)."), }) } + +func (s *assertMgrSuite) TestFetchAccountKeyOK(c *C) { + s.state.Lock() + defer s.state.Unlock() + + storeAs := s.setupModelAndStore(c) + err := s.storeSigning.Add(storeAs) + c.Assert(err, IsNil) + + db, err := asserts.OpenDatabase(&asserts.DatabaseConfig{ + Backstore: asserts.NewMemoryBackstore(), + Trusted: s.storeSigning.Trusted, + }) + c.Assert(err, IsNil) + + assertstate.ReplaceDB(s.state, db) + + keyID := s.dev1AcctKey.PublicKeyID() + + // not found locally + _, err = assertstate.AccountKey(s.state, keyID) + c.Assert(err, testutil.ErrorIs, &asserts.NotFoundError{}) + + err = assertstate.FetchAccountKey(s.state, 0, keyID) + c.Assert(err, IsNil) + + // found after fetch + _, err = assertstate.AccountKey(s.state, keyID) + c.Assert(err, IsNil) +} + +func (s *assertMgrSuite) TestFetchAccountKeyError(c *C) { + s.state.Lock() + defer s.state.Unlock() + + storeAs := s.setupModelAndStore(c) + err := s.storeSigning.Add(storeAs) + c.Assert(err, IsNil) + + db, err := asserts.OpenDatabase(&asserts.DatabaseConfig{ + Backstore: asserts.NewMemoryBackstore(), + Trusted: s.storeSigning.Trusted, + }) + c.Assert(err, IsNil) + + assertstate.ReplaceDB(s.state, db) + + err = assertstate.FetchAccountKey(s.state, 0, "no-such-key-id") + c.Assert(err, testutil.ErrorIs, &asserts.NotFoundError{}) +} diff --git a/overlord/devicemgmtstate/devicemgmtmgr.go b/overlord/devicemgmtstate/devicemgmtmgr.go index b75198b54a..f3d26d0edd 100644 --- a/overlord/devicemgmtstate/devicemgmtmgr.go +++ b/overlord/devicemgmtstate/devicemgmtmgr.go @@ -36,6 +36,7 @@ import ( "github.com/snapcore/snapd/asserts" "github.com/snapcore/snapd/features" "github.com/snapcore/snapd/logger" + "github.com/snapcore/snapd/overlord/assertstate" "github.com/snapcore/snapd/overlord/configstate/config" "github.com/snapcore/snapd/overlord/snapstate" "github.com/snapcore/snapd/overlord/state" @@ -54,7 +55,8 @@ const ( ) var ( - timeNow = time.Now + timeNow = time.Now + assertstateFetchAccountKey = assertstate.FetchAccountKey maxSequences = 256 maxBlockedMessagesPerSequence = 8 @@ -64,6 +66,12 @@ var ( deviceMgmtExchangeChangeKind = swfeats.RegisterChangeKind("device-management-exchange") ) +// deviceBackend provides device identity and response message signing. +type deviceBackend interface { + Serial() (*asserts.Serial, error) + SignResponseMessage(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) +} + // MessageHandler processes request messages of a specific kind. // Caller must hold state lock when using this interface. type MessageHandler interface { @@ -79,6 +87,16 @@ type MessageHandler interface { ResultFromChange(chg *state.Change) (body map[string]any, err error) } +// UnauthorizedError is returned by MessageHandler.Validate when the operator +// does not have permission to perform the requested action. +type UnauthorizedError struct { + Operator string +} + +func (e *UnauthorizedError) Error() string { + return fmt.Sprintf("cannot perform action: operator %q is not authorized", e.Operator) +} + // MarkChangeForMessage records the message ID on the change created by an Apply // implementation. It must be called after change creation and before releasing // the state lock, so that doApplyMessage can recover the change ID on retry @@ -87,11 +105,6 @@ func MarkChangeForMessage(chg *state.Change, msg *RequestMessage) { chg.Set(mgmtMessageIDKey, msg.ID()) } -// responseMessageSigner can sign response-message assertions. -type responseMessageSigner interface { - SignResponseMessage(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) -} - // RequestMessage represents a request-message being processed. // Messages remain pending until their associated change completes, // at which point a response is queued and the message is removed. @@ -104,6 +117,7 @@ type RequestMessage struct { Devices []string `json:"devices"` ValidSince time.Time `json:"valid-since"` ValidUntil time.Time `json:"valid-until"` + Assumes []string `json:"assumes,omitempty"` Body string `json:"body"` ReceiveTime time.Time `json:"receive-time"` @@ -116,6 +130,9 @@ type RequestMessage struct { // A non-empty ResponseStatus means the message has been fully processed. ResponseStatus asserts.MessageStatus `json:"response-status,omitempty"` ResponseBody map[string]any `json:"response-body,omitempty"` + + // RawAssertion holds the original encoded assertion bytes. + RawAssertion []byte `json:"raw-assertion"` } // ID returns the full message identifier `BaseID[-SeqNum]`. @@ -127,6 +144,23 @@ func (msg *RequestMessage) ID() string { return msg.BaseID } +// ValidAt returns whether the request-message is valid at 'when' time. +func (msg *RequestMessage) ValidAt(when time.Time) bool { + return (when.Equal(msg.ValidSince) || when.After(msg.ValidSince)) && when.Before(msg.ValidUntil) +} + +// Targets returns whether the given device is listed in the message's devices header. +func (msg *RequestMessage) Targets(devID asserts.DeviceID) bool { + target := devID.String() + for _, d := range msg.Devices { + if d == target { + return true + } + } + + return false +} + // sequenceState holds the messages and progress for a single base ID, // covering both sequenced & unsequenced messages. type sequenceState struct { @@ -238,9 +272,14 @@ func (ms *deviceMgmtState) enqueueRequestMessages(pollResp *store.MessageExchang } if len(pollResp.Messages) > 0 { + // Record the token of the last message as the ack position. It will be + // sent on the next exchange to advance the store's queue pointer. token := pollResp.Messages[len(pollResp.Messages)-1].Token ms.LastReceivedToken = token } else { + // No messages were returned: the store has already advanced the queue + // pointer from the previous exchange. Clear the token so the next + // exchange sends no ack, since there is nothing new to acknowledge. ms.LastReceivedToken = "" } @@ -260,15 +299,15 @@ func (ms *deviceMgmtState) removeSequenceFromLRU(baseID string) { // DeviceMgmtManager handles device management operations. type DeviceMgmtManager struct { state *state.State - signer responseMessageSigner + device deviceBackend handlers map[string]MessageHandler } // Manager creates a new DeviceMgmtManager. -func Manager(state *state.State, runner *state.TaskRunner, signer responseMessageSigner) *DeviceMgmtManager { +func Manager(state *state.State, runner *state.TaskRunner, backend deviceBackend) *DeviceMgmtManager { m := &DeviceMgmtManager{ state: state, - signer: signer, + device: backend, handlers: make(map[string]MessageHandler), } @@ -312,6 +351,21 @@ func (m *DeviceMgmtManager) getState() (*deviceMgmtState, error) { return &ms, nil } +// getMessageAndState retrieves the current state along with the given message. +func (m *DeviceMgmtManager) getMessageAndState(msgID string) (*deviceMgmtState, *RequestMessage, error) { + ms, err := m.getState() + if err != nil { + return nil, nil, err + } + + msg, err := ms.getRequestMessage(msgID) + if err != nil { + return nil, nil, err + } + + return ms, msg, nil +} + // setState persists the device management state. func (m *DeviceMgmtManager) setState(ms *deviceMgmtState) { m.state.Set(deviceMgmtStateKey, ms) @@ -546,27 +600,143 @@ func (m *DeviceMgmtManager) rejectSequence(ms *deviceMgmtState, chg *state.Chang // doValidateMessage performs snapd-level and subsystem-level validation on a message. func (m *DeviceMgmtManager) doValidateMessage(t *state.Task, _ *tomb.Tomb) error { - // TODO: implement this task, no-op for now. + m.state.Lock() + defer m.state.Unlock() + + var msgID string + err := t.Get("message-id", &msgID) + if err != nil { + return err + } + + ms, msg, err := m.getMessageAndState(msgID) + if err != nil { + return err + } + + if msg.ResponseStatus != "" { + return nil + } + + setMsgResponse := func(status asserts.MessageStatus, message string) { + msg.ResponseStatus = status + msg.ResponseBody = map[string]any{"message": message} + m.setState(ms) + } + rejectMsg := func(reason string) { + setMsgResponse(asserts.MessageStatusRejected, reason) + } + + a, err := asserts.Decode(msg.RawAssertion) + if err != nil { + rejectMsg(fmt.Sprintf("cannot decode message: %v", err)) + return nil + } + + fetched, err := m.ensureAccountKey(a.SignKeyID()) + if err != nil { + // TODO: need to distinguish between: + // - transient errors (like store unreachable) - retry + // - permanent errors - determine whether to fail the task or reject the message. + return err + } + if fetched { + // The state lock was dropped during the store fetch. Concurrent tasks in + // other lanes may have mutated state in that window, so re-read before mutating. + ms, msg, err = m.getMessageAndState(msgID) + if err != nil { + return err + } + } + + err = assertstate.DB(m.state).Check(a) + if err != nil { + rejectMsg(fmt.Sprintf("cannot verify message signature: %v", err)) + return nil + } + + serial, err := m.device.Serial() + if err != nil { + return err + } + devID := serial.DeviceID() + if !msg.Targets(devID) { + rejectMsg(fmt.Sprintf("cannot process message: not intended for device %s", devID)) + return nil + } + + now := timeNow() + if !msg.ValidAt(now) { + rejectMsg(fmt.Sprintf("cannot process message: not valid at %s", now.UTC().Format(time.RFC3339))) + return nil + } + + // TODO: implement assumes checks (SD187, SD251). The design is somewhat in + // flux: SD251 "extends" messages to non-run-mode contexts (e.g. first boot, + // before the device has an identity), which will require assumes entries + // like "seeding". For now, only the confdb subsystem is supported and + // no specific features need to be declared. + + handler, ok := m.handlers[msg.Kind] + if !ok { + rejectMsg(fmt.Sprintf("cannot find handler for message kind %q", msg.Kind)) + return nil + } + + err = handler.Validate(m.state, msg) + if err != nil { + var unauthorizedErr *UnauthorizedError + status := asserts.MessageStatusRejected + if errors.As(err, &unauthorizedErr) { + status = asserts.MessageStatusUnauthorized + } + reason := err.Error() + + // handler.Validate may drop the state lock internally. Concurrent tasks + // in other lanes may have mutated the state in that window, so re-read before mutating. + ms, msg, err = m.getMessageAndState(msgID) + if err != nil { + return err + } + + setMsgResponse(status, reason) + return nil + } + return nil } +// ensureAccountKey fetches the account-key assertion for signKeyID from the +// store if it is not already in the local database. +func (m *DeviceMgmtManager) ensureAccountKey(signKeyID string) (fetched bool, err error) { + _, err = assertstate.AccountKey(m.state, signKeyID) + if err == nil { + return false, nil + } + if !errors.Is(err, &asserts.NotFoundError{}) { + return false, err + } + + err = assertstateFetchAccountKey(m.state, 0, signKeyID) + if err != nil && !errors.Is(err, &asserts.NotFoundError{}) { + return true, err + } + + return true, nil +} + // doApplyMessage dispatches the message to its subsystem handler for processing. func (m *DeviceMgmtManager) doApplyMessage(t *state.Task, _ *tomb.Tomb) error { m.state.Lock() defer m.state.Unlock() - ms, err := m.getState() - if err != nil { - return err - } - var msgID string - err = t.Get("message-id", &msgID) + err := t.Get("message-id", &msgID) if err != nil { return err } - msg, err := ms.getRequestMessage(msgID) + ms, msg, err := m.getMessageAndState(msgID) if err != nil { return err } @@ -576,12 +746,11 @@ func (m *DeviceMgmtManager) doApplyMessage(t *state.Task, _ *tomb.Tomb) error { return nil } - defer m.setState(ms) - // Check if a change was already created for this message before persisting its ApplyChangeID. chg := findChangeByMgmtMessageID(m.state, msgID) if chg != nil { msg.ApplyChangeID = chg.ID() + m.setState(ms) return nil } @@ -589,16 +758,26 @@ func (m *DeviceMgmtManager) doApplyMessage(t *state.Task, _ *tomb.Tomb) error { if !ok { msg.ResponseStatus = asserts.MessageStatusError msg.ResponseBody = map[string]any{"message": fmt.Sprintf("cannot find handler for message kind %q", msg.Kind)} + m.setState(ms) return nil } - chgID, err := handler.Apply(m.state, msg) + chgID, applyErr := handler.Apply(m.state, msg) + + // handler.Apply may drop the state lock internally. Concurrent tasks in + // other lanes may have mutated the state in that window, so re-read before mutating. + ms, msg, err = m.getMessageAndState(msgID) if err != nil { + return err + } + + if applyErr != nil { msg.ResponseStatus = asserts.MessageStatusError - msg.ResponseBody = map[string]any{"message": err.Error()} + msg.ResponseBody = map[string]any{"message": applyErr.Error()} } else { msg.ApplyChangeID = chgID } + m.setState(ms) return nil } @@ -631,6 +810,19 @@ func (m *DeviceMgmtManager) doQueueResponse(t *state.Task, _ *tomb.Tomb) error { return err } + // handler.ResultFromChange may drop the state lock internally. Concurrent tasks + // in other lanes may have mutated the state in that window, so re-read before mutating. + responseStatus := msg.ResponseStatus + responseBody := msg.ResponseBody + + ms, msg, err = m.getMessageAndState(msgID) + if err != nil { + return err + } + + msg.ResponseStatus = responseStatus + msg.ResponseBody = responseBody + bodyBytes, err := json.Marshal(msg.ResponseBody) if err != nil { return fmt.Errorf("cannot marshal response body: %w", err) @@ -641,7 +833,7 @@ func (m *DeviceMgmtManager) doQueueResponse(t *state.Task, _ *tomb.Tomb) error { // next change, but the request message remains in state until doQueueResponse completes, // so such failures leave it hanging indefinitely. - resAs, err := m.signer.SignResponseMessage(msg.AccountID, msg.ID(), msg.ResponseStatus, bodyBytes) + resAs, err := m.device.SignResponseMessage(msg.AccountID, msg.ID(), msg.ResponseStatus, bodyBytes) if err != nil { return fmt.Errorf("cannot sign response message: %w", err) } @@ -691,7 +883,7 @@ func (m *DeviceMgmtManager) setMessageResponseFromChange(msg *RequestMessage) er body, err := handler.ResultFromChange(change) if err != nil { msg.ResponseStatus = asserts.MessageStatusError - msg.ResponseBody = map[string]any{"message": fmt.Sprintf("cannot process message: %v", err)} + msg.ResponseBody = map[string]any{"message": err.Error()} } else { msg.ResponseStatus = asserts.MessageStatusSuccess msg.ResponseBody = body @@ -723,16 +915,18 @@ func parseRequestMessage(msg store.Message) (*RequestMessage, error) { } return &RequestMessage{ - AccountID: reqAs.AccountID(), - AuthorityID: reqAs.AuthorityID(), - BaseID: reqAs.ID(), - SeqNum: reqAs.SeqNum(), - Kind: reqAs.Kind(), - Devices: deviceIDs, - ValidSince: reqAs.ValidSince(), - ValidUntil: reqAs.ValidUntil(), - Body: string(reqAs.Body()), - ReceiveTime: timeNow(), + AccountID: reqAs.AccountID(), + AuthorityID: reqAs.AuthorityID(), + BaseID: reqAs.ID(), + SeqNum: reqAs.SeqNum(), + Kind: reqAs.Kind(), + Devices: deviceIDs, + ValidSince: reqAs.ValidSince(), + ValidUntil: reqAs.ValidUntil(), + Assumes: reqAs.Assumes(), + Body: string(reqAs.Body()), + ReceiveTime: timeNow(), + RawAssertion: []byte(msg.Data), }, nil } diff --git a/overlord/devicemgmtstate/devicemgmtmgr_test.go b/overlord/devicemgmtstate/devicemgmtmgr_test.go index 2512d35c28..70ee0fe1ca 100644 --- a/overlord/devicemgmtstate/devicemgmtmgr_test.go +++ b/overlord/devicemgmtstate/devicemgmtmgr_test.go @@ -38,6 +38,8 @@ import ( "github.com/snapcore/snapd/features" "github.com/snapcore/snapd/logger" "github.com/snapcore/snapd/overlord" + "github.com/snapcore/snapd/overlord/assertstate" + "github.com/snapcore/snapd/overlord/auth" "github.com/snapcore/snapd/overlord/configstate/config" "github.com/snapcore/snapd/overlord/devicemgmtstate" "github.com/snapcore/snapd/overlord/snapstate" @@ -55,13 +57,32 @@ var noopTask = func(*state.Task, *tomb.Tomb) error { return nil } type mockStore struct { storetest.Store + assertionDB *asserts.Database exchangeMessages func(ctx context.Context, req *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) } +func (s *mockStore) Assertion(assertType *asserts.AssertionType, key []string, _ *auth.UserState) (asserts.Assertion, error) { + ref := &asserts.Ref{Type: assertType, PrimaryKey: key} + return ref.Resolve(s.assertionDB.Find) +} + func (s *mockStore) ExchangeMessages(ctx context.Context, req *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { return s.exchangeMessages(ctx, req) } +type mockDeviceBackend struct { + serial *asserts.Serial + sign func(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) +} + +func (m *mockDeviceBackend) Serial() (*asserts.Serial, error) { + return m.serial, nil +} + +func (m *mockDeviceBackend) SignResponseMessage(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) { + return m.sign(accountID, messageID, status, body) +} + type mockMessageHandler struct { validate func(st *state.State, msg *devicemgmtstate.RequestMessage) error apply func(st *state.State, msg *devicemgmtstate.RequestMessage) (string, error) @@ -92,14 +113,6 @@ func (h *mockMessageHandler) ResultFromChange(chg *state.Change) (map[string]any return nil, nil } -type mockSigner struct { - sign func(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) -} - -func (s *mockSigner) SignResponseMessage(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) { - return s.sign(accountID, messageID, status, body) -} - func setRemoteMgmtFeatureFlag(c *C, st *state.State, value any) { tr := config.NewTransaction(st) _, confOption := features.RemoteDeviceManagement.ConfigOption() @@ -140,13 +153,23 @@ func (s *deviceMgmtMgrSuite) SetUpTest(c *C) { s.mockModel() s.storeStack = assertstest.NewStoreStack("my-brand", nil) + db, err := asserts.OpenDatabase(&asserts.DatabaseConfig{ + Backstore: asserts.NewMemoryBackstore(), + Trusted: s.storeStack.Trusted, + }) + c.Assert(err, IsNil) + c.Assert(db.Add(s.storeStack.StoreAccountKey("")), IsNil) + assertstate.ReplaceDB(s.st, db) + s.runner = s.o.TaskRunner() s.o.AddManager(s.runner) s.mgr = devicemgmtstate.Manager(s.st, s.runner, nil) s.o.AddManager(s.mgr) - err := s.o.StartUp() + s.mgr.MockBackend(&mockDeviceBackend{serial: s.makeSerial(c, "serial-1")}) + + err = s.o.StartUp() c.Assert(err, IsNil) var restoreLogger func() @@ -165,7 +188,7 @@ func (s *deviceMgmtMgrSuite) SetUpTest(c *C) { return chg.ID(), nil }, resultFromChange: func(*state.Change) (map[string]any, error) { - return map[string]any{"key": "value"}, nil + return map[string]any{"values": "ok"}, nil }, }) } @@ -188,8 +211,30 @@ func (s *deviceMgmtMgrSuite) mockModel() { s.st.Set("seeded", true) } +func (s *deviceMgmtMgrSuite) makeSerial(c *C, serial string) *asserts.Serial { + devKey, _ := assertstest.GenerateKey(752) + encDevKey, err := asserts.EncodePublicKey(devKey.PublicKey()) + c.Assert(err, IsNil) + + as, err := s.storeStack.Sign(asserts.SerialType, map[string]any{ + "authority-id": "my-brand", + "brand-id": "my-brand", + "model": "my-model", + "serial": serial, + "device-key": string(encDevKey), + "device-key-sha3-384": devKey.PublicKey().ID(), + "timestamp": fixedTestTime.UTC().Format(time.RFC3339), + }, nil, "") + c.Assert(err, IsNil) + + return as.(*asserts.Serial) +} + func (s *deviceMgmtMgrSuite) mockStore(exchangeMessages func(context.Context, *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error)) { - snapstate.ReplaceStore(s.st, &mockStore{exchangeMessages: exchangeMessages}) + snapstate.ReplaceStore(s.st, &mockStore{ + assertionDB: s.storeStack.Database, + exchangeMessages: exchangeMessages, + }) } func (s *deviceMgmtMgrSuite) makeStoreRequestMessage(c *C, messageID, kind, token string) store.MessageWithToken { @@ -481,7 +526,7 @@ func (s *deviceMgmtMgrSuite) TestDoExchangeMessagesReplyOK(c *C) { c.Check(ms.Sequences, HasLen, 0) } -func (s *deviceMgmtMgrSuite) TestDoExchangeMessagesSequenceLRU(c *C) { +func (s *deviceMgmtMgrSuite) TestDoExchangeMessagesSequenceLRUOrdering(c *C) { s.st.Lock() defer s.st.Unlock() @@ -1040,7 +1085,7 @@ func (s *deviceMgmtMgrSuite) TestDoDispatchMessagesIdempotent(c *C) { c.Check(ti.queue["msg2"], NotNil) } -func (s *deviceMgmtMgrSuite) TestDispatchMessagesLaneIsolation(c *C) { +func (s *deviceMgmtMgrSuite) TestDoDispatchMessagesLaneIsolation(c *C) { s.st.Lock() defer s.st.Unlock() @@ -1081,7 +1126,8 @@ func (s *deviceMgmtMgrSuite) TestDispatchMessagesLaneIsolation(c *C) { }, }) - s.mgr.MockSigner(&mockSigner{ + s.mgr.MockBackend(&mockDeviceBackend{ + serial: s.makeSerial(c, "serial-1"), sign: func(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) { return assertstest.FakeAssertionWithBody(body, map[string]any{ "type": "response-message", @@ -1116,6 +1162,473 @@ func (s *deviceMgmtMgrSuite) TestDispatchMessagesLaneIsolation(c *C) { c.Check(ms.ReadyResponses["msg2"].Format, Equals, "assertion") } +func (s *deviceMgmtMgrSuite) TestDoValidateMessageOK(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1"), + }, + }, nil + }) + + s.runner.AddHandler("apply-mgmt-message", noopTask, nil) + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatus("")) // message wasn't rejected +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageAccountKeyFetchedFromStoreOK(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1"), + }, + }, nil + }) + + // Remove the account key from the local DB so ensureAccountKey must fetch. + db, err := asserts.OpenDatabase(&asserts.DatabaseConfig{ + Backstore: asserts.NewMemoryBackstore(), + Trusted: s.storeStack.Trusted, + }) + c.Assert(err, IsNil) + assertstate.ReplaceDB(s.st, db) + + s.AddCleanup(devicemgmtstate.MockFetchAccountKey(func(_ *state.State, _ int, _ string) error { + err := db.Add(s.storeStack.StoreAccountKey("")) + c.Assert(err, IsNil) + return nil + })) + + s.runner.AddHandler("apply-mgmt-message", noopTask, nil) + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatus("")) // message wasn't rejected +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageBadRawAssertion(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{}, nil + }) + + ms := &devicemgmtstate.DeviceMgmtState{ + Sequences: map[string]*devicemgmtstate.SequenceState{ + "msg1": { + Messages: []*devicemgmtstate.RequestMessage{ + { + AccountID: "my-brand", + AuthorityID: "my-brand", + BaseID: "msg1", + Kind: "test-kind", + Devices: []string{"serial-1.my-model.my-brand"}, + ValidSince: fixedTestTime.Add(-time.Hour), + ValidUntil: fixedTestTime.Add(24 * time.Hour), + Body: `{"action": "get"}`, + RawAssertion: []byte("not a valid assertion"), + }, + }, + }, + }, + } + s.mgr.SetState(ms) + + chg := s.st.NewChange("test", "test change") + t := s.st.NewTask("validate-mgmt-message", "validate msg1") + t.Set("message-id", "msg1") + chg.AddTask(t) + + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatusRejected) + c.Check(msg.ResponseBody["message"], Equals, "cannot decode message: assertion content/signature separator not found") +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageFetchAccountKeyError(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1"), + }, + }, nil + }) + + // Remove the account key from the local DB so ensureAccountKey reaches the fetch. + db, err := asserts.OpenDatabase(&asserts.DatabaseConfig{ + Backstore: asserts.NewMemoryBackstore(), + Trusted: s.storeStack.Trusted, + }) + c.Assert(err, IsNil) + assertstate.ReplaceDB(s.st, db) + + fetchErr := fmt.Errorf("store unavailable") + s.AddCleanup(devicemgmtstate.MockFetchAccountKey(func(_ *state.State, _ int, _ string) error { + return fetchErr + })) + + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatus("")) +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageBadSignature(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{}, nil + }) + + storeMsg := s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1") + reqMsg, err := devicemgmtstate.ParseRequestMessage(storeMsg.Message) + c.Assert(err, IsNil) + + // tamper the raw assertion body + reqMsg.RawAssertion = bytes.Replace(reqMsg.RawAssertion, []byte("get"), []byte("set"), 1) + + ms := &devicemgmtstate.DeviceMgmtState{ + Sequences: map[string]*devicemgmtstate.SequenceState{ + "msg1": {Messages: []*devicemgmtstate.RequestMessage{reqMsg}}, + }, + } + s.mgr.SetState(ms) + + chg := s.st.NewChange("test", "test change") + t := s.st.NewTask("validate-mgmt-message", "validate msg1") + t.Set("message-id", "msg1") + chg.AddTask(t) + + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err = s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatusRejected) + c.Check(msg.ResponseBody["message"], Equals, + "cannot verify message signature: failed signature verification: openpgp: invalid signature: hash tag doesn't match") +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageDeviceNotTargeted(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1"), + }, + }, nil + }) + + s.mgr.MockBackend(&mockDeviceBackend{serial: s.makeSerial(c, "other-serial")}) + + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatusRejected) + c.Check(msg.ResponseBody["message"], Equals, "cannot process message: not intended for device other-serial.my-model.my-brand") +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageExpired(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1"), + }, + }, nil + }) + + s.AddCleanup(devicemgmtstate.MockTimeNow(fixedTestTime.Add(48 * time.Hour))) + + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatusRejected) + c.Check(msg.ResponseBody["message"], Equals, "cannot process message: not valid at 2025-06-16T12:00:00Z") +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageUnknownKind(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "unknown-kind", "token-1"), + }, + }, nil + }) + + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatusRejected) + c.Check(msg.ResponseBody["message"], Equals, `cannot find handler for message kind "unknown-kind"`) +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageUnauthorized(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1"), + }, + }, nil + }) + + s.mgr.RegisterHandler("test-kind", &mockMessageHandler{ + validate: func(_ *state.State, _ *devicemgmtstate.RequestMessage) error { + return &devicemgmtstate.UnauthorizedError{Operator: "alice"} + }, + }) + + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatusUnauthorized) + c.Check(msg.ResponseBody["message"], Equals, `cannot perform action: operator "alice" is not authorized`) +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageHandlerError(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1"), + }, + }, nil + }) + + s.mgr.RegisterHandler("test-kind", &mockMessageHandler{ + validate: func(_ *state.State, _ *devicemgmtstate.RequestMessage) error { + return fmt.Errorf("cannot validate message: invalid payload") + }, + }) + + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatusRejected) + c.Check(msg.ResponseBody["message"], Equals, "cannot validate message: invalid payload") +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageIdempotent(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{}, nil + }) + + // Prevent FetchAccountKey from dropping the state lock, which would let the + // concurrent retry tasks race past the idempotency guard. + s.AddCleanup(devicemgmtstate.MockFetchAccountKey(func(_ *state.State, _ int, _ string) error { + return nil + })) + + validateCalls := 0 + s.mgr.RegisterHandler("test-kind", &mockMessageHandler{ + validate: func(_ *state.State, _ *devicemgmtstate.RequestMessage) error { + validateCalls++ + return fmt.Errorf("cannot validate message: invalid payload") + }, + }) + + storeMsg := s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1") + reqMsg, err := devicemgmtstate.ParseRequestMessage(storeMsg.Message) + c.Assert(err, IsNil) + + ms := &devicemgmtstate.DeviceMgmtState{ + Sequences: map[string]*devicemgmtstate.SequenceState{ + "msg1": { + Messages: []*devicemgmtstate.RequestMessage{reqMsg}, + }, + }, + } + s.mgr.SetState(ms) + + chg := s.st.NewChange("test", "test change") + for i := 1; i <= 3; i++ { + t := s.st.NewTask("validate-mgmt-message", fmt.Sprintf("validate msg1 attempt %d", i)) + t.Set("message-id", "msg1") + chg.AddTask(t) + } + + s.settle(c) + + c.Check(chg.Status(), Equals, state.DoneStatus) + c.Check(validateCalls, Equals, 1) + + ms, err = s.mgr.GetState() + c.Assert(err, IsNil) + + msg := ms.Sequences["msg1"].Messages[0] + c.Check(msg.ResponseStatus, Equals, asserts.MessageStatusRejected) +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageConcurrentWriteAfterFetch(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "unknown-kind", "token-1"), + s.makeStoreRequestMessage(c, "msg2", "unknown-kind", "token-2"), + }, + }, nil + }) + + firstCall := true + s.AddCleanup(devicemgmtstate.MockFetchAccountKey(func(st *state.State, _ int, _ string) error { + if firstCall { + firstCall = false + // Wait for the other lane's full write before resuming. + lane2Done := make(chan struct{}) + st.AddTaskStatusChangedHandler(func(t *state.Task, _, new state.Status) (remove bool) { + if t.Kind() == "validate-mgmt-message" && new == state.DoneStatus { + close(lane2Done) + return true + } + return false + }) + + st.Unlock() + <-lane2Done + st.Lock() + } + return nil + })) + + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + c.Check(ms.Sequences["msg1"].Messages[0].ResponseStatus, Equals, asserts.MessageStatusRejected) + c.Check(ms.Sequences["msg2"].Messages[0].ResponseStatus, Equals, asserts.MessageStatusRejected) +} + +func (s *deviceMgmtMgrSuite) TestDoValidateMessageConcurrentWriteAfterValidate(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1"), + s.makeStoreRequestMessage(c, "msg2", "test-kind", "token-2"), + }, + }, nil + }) + + firstCall := true + s.mgr.RegisterHandler("test-kind", &mockMessageHandler{ + validate: func(st *state.State, _ *devicemgmtstate.RequestMessage) error { + if firstCall { + firstCall = false + // Wait for the other lane's full write before resuming. + lane2Done := make(chan struct{}) + st.AddTaskStatusChangedHandler(func(t *state.Task, _, new state.Status) (remove bool) { + if t.Kind() == "validate-mgmt-message" && new == state.DoneStatus { + close(lane2Done) + return true + } + return false + }) + + st.Unlock() + <-lane2Done + st.Lock() + } + + return fmt.Errorf("cannot validate message: rejected") + }, + }) + + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + c.Check(ms.Sequences["msg1"].Messages[0].ResponseStatus, Equals, asserts.MessageStatusRejected) + c.Check(ms.Sequences["msg2"].Messages[0].ResponseStatus, Equals, asserts.MessageStatusRejected) +} + func (s *deviceMgmtMgrSuite) TestDoApplyMessageOK(c *C) { s.st.Lock() defer s.st.Unlock() @@ -1164,7 +1677,7 @@ func (s *deviceMgmtMgrSuite) TestDoApplyMessageSkipIfAlreadyFailed(c *C) { ms.Sequences[messageID].Messages[0].ResponseStatus = asserts.MessageStatusRejected ms.Sequences[messageID].Messages[0].ResponseBody = map[string]any{ - "message": "cannot process message: device not in target list", + "message": "cannot process message: not intended for device serial-1.my-model.my-brand", } s.mgr.SetState(ms) @@ -1197,14 +1710,33 @@ func (s *deviceMgmtMgrSuite) TestDoApplyMessageNoHandlerForMessageKind(c *C) { defer s.st.Unlock() s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { - return &store.MessageExchangeResponse{ - Messages: []store.MessageWithToken{ - s.makeStoreRequestMessage(c, "msg1", "unknown-kind", "token-1"), - }, - }, nil + return &store.MessageExchangeResponse{}, nil }) - s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + ms := &devicemgmtstate.DeviceMgmtState{ + Sequences: map[string]*devicemgmtstate.SequenceState{ + "msg1": { + Messages: []*devicemgmtstate.RequestMessage{ + { + AccountID: "my-brand", + AuthorityID: "my-brand", + BaseID: "msg1", + Kind: "unknown-kind", + Devices: []string{"serial-1.my-model.my-brand"}, + ValidSince: fixedTestTime, + ValidUntil: fixedTestTime.Add(24 * time.Hour), + Body: `{"action": "get"}`, + }, + }, + }, + }, + } + s.mgr.SetState(ms) + + chg := s.st.NewChange("test", "test change") + t := s.st.NewTask("apply-mgmt-message", "apply message with unknown kind") + t.Set("message-id", "msg1") + chg.AddTask(t) s.settle(c) @@ -1231,7 +1763,7 @@ func (s *deviceMgmtMgrSuite) TestDoApplyMessageApplyError(c *C) { s.mgr.RegisterHandler("test-kind", &mockMessageHandler{ apply: func(st *state.State, msg *devicemgmtstate.RequestMessage) (string, error) { - return "", fmt.Errorf("system in inconsistent state") + return "", fmt.Errorf("cannot apply message: system in inconsistent state") }, }) @@ -1245,7 +1777,7 @@ func (s *deviceMgmtMgrSuite) TestDoApplyMessageApplyError(c *C) { msg := ms.Sequences["msg1"].Messages[0] c.Check(msg.ApplyChangeID, Equals, "") c.Check(msg.ResponseStatus, Equals, asserts.MessageStatusError) - c.Check(msg.ResponseBody["message"], Equals, "system in inconsistent state") + c.Check(msg.ResponseBody["message"], Equals, "cannot apply message: system in inconsistent state") } func (s *deviceMgmtMgrSuite) TestDoApplyMessageIdempotent(c *C) { @@ -1412,31 +1944,77 @@ func (s *deviceMgmtMgrSuite) TestDoApplyMessageMessageNotFound(c *C) { c.Assert(err, ErrorMatches, `cannot find message "seqA-2"`) } -func (s *deviceMgmtMgrSuite) TestDoQueueResponseOK(c *C) { +func (s *deviceMgmtMgrSuite) TestDoApplyMessageConcurrentWriteAfterApply(c *C) { s.st.Lock() defer s.st.Unlock() s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { return &store.MessageExchangeResponse{ Messages: []store.MessageWithToken{ - s.makeStoreRequestMessage(c, "mesg-1", "test-kind", "token-1"), + s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1"), + s.makeStoreRequestMessage(c, "msg2", "test-kind", "token-2"), }, }, nil }) + firstCall := true s.mgr.RegisterHandler("test-kind", &mockMessageHandler{ - validate: func(*state.State, *devicemgmtstate.RequestMessage) error { return nil }, apply: func(st *state.State, msg *devicemgmtstate.RequestMessage) (string, error) { + if firstCall { + firstCall = false + // Wait for the other lane's full write before resuming. + lane2Done := make(chan struct{}) + st.AddTaskStatusChangedHandler(func(t *state.Task, _, new state.Status) (remove bool) { + if t.Kind() == "apply-mgmt-message" && new == state.DoneStatus { + close(lane2Done) + return true + } + return false + }) + + st.Unlock() + <-lane2Done + st.Lock() + } + chg := st.NewChange("subsystem", "apply payload") devicemgmtstate.MarkChangeForMessage(chg, msg) return chg.ID(), nil }, - resultFromChange: func(*state.Change) (map[string]any, error) { - return map[string]any{"values": "ok"}, nil - }, }) - s.mgr.MockSigner(&mockSigner{ + s.runner.AddHandler("queue-mgmt-response", noopTask, nil) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + for i := 1; i <= 2; i++ { + msgID := fmt.Sprintf("msg%d", i) + applyChgID := ms.Sequences[msgID].Messages[0].ApplyChangeID + c.Check(applyChgID, Not(Equals), "") + + applyChg := s.st.Change(applyChgID) + c.Assert(applyChg, Not(IsNil)) + c.Check(applyChg.Has(devicemgmtstate.MgmtMessageIDKey), Equals, true) + } +} + +func (s *deviceMgmtMgrSuite) TestDoQueueResponseOK(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "mesg-1", "test-kind", "token-1"), + }, + }, nil + }) + + s.mgr.MockBackend(&mockDeviceBackend{ + serial: s.makeSerial(c, "serial-1"), sign: func(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) { c.Check(accountID, Equals, "my-brand") c.Check(messageID, Equals, "mesg-1") @@ -1506,7 +2084,8 @@ func (s *deviceMgmtMgrSuite) TestDoQueueResponseStatusAlreadyKnown(c *C) { }, }) - s.mgr.MockSigner(&mockSigner{ + s.mgr.MockBackend(&mockDeviceBackend{ + serial: s.makeSerial(c, "serial-1"), sign: func(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) { c.Check(messageID, Equals, "mesg-1") c.Check(status, Equals, asserts.MessageStatusRejected) @@ -1538,7 +2117,8 @@ func (s *deviceMgmtMgrSuite) TestDoQueueResponseIdempotent(c *C) { defer s.st.Unlock() signCalls := 0 - s.mgr.MockSigner(&mockSigner{ + s.mgr.MockBackend(&mockDeviceBackend{ + serial: s.makeSerial(c, "serial-1"), sign: func(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) { signCalls++ return assertstest.FakeAssertionWithBody(body, map[string]any{ @@ -1615,15 +2195,16 @@ func (s *deviceMgmtMgrSuite) TestDoQueueResponseResultFromChangeError(c *C) { return chg.ID(), nil }, resultFromChange: func(*state.Change) (map[string]any, error) { - return nil, fmt.Errorf("operation failed") + return nil, fmt.Errorf("cannot get result from change: operation failed") }, }) - s.mgr.MockSigner(&mockSigner{ + s.mgr.MockBackend(&mockDeviceBackend{ + serial: s.makeSerial(c, "serial-1"), sign: func(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) { c.Check(messageID, Equals, "mesg-1") c.Check(status, Equals, asserts.MessageStatusError) - c.Check(string(body), Equals, `{"message":"cannot process message: operation failed"}`) + c.Check(string(body), Equals, `{"message":"cannot get result from change: operation failed"}`) return assertstest.FakeAssertionWithBody(body, map[string]any{ "type": "response-message", @@ -1717,7 +2298,8 @@ func (s *deviceMgmtMgrSuite) TestDoQueueResponseNoHandlerForMessageKind(c *C) { } s.mgr.SetState(ms) - s.mgr.MockSigner(&mockSigner{ + s.mgr.MockBackend(&mockDeviceBackend{ + serial: s.makeSerial(c, "serial-1"), sign: func(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) { c.Check(accountID, Equals, "my-brand") c.Check(messageID, Equals, "msg1") @@ -1811,7 +2393,8 @@ func (s *deviceMgmtMgrSuite) TestDoQueueResponseSubsystemChangeNotReady(c *C) { changeReady = true subsysChg.SetStatus(state.DoneStatus) - s.mgr.MockSigner(&mockSigner{ + s.mgr.MockBackend(&mockDeviceBackend{ + serial: s.makeSerial(c, "serial-1"), sign: func(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) { return assertstest.FakeAssertionWithBody(body, map[string]any{ "type": "response-message", @@ -1859,7 +2442,8 @@ func (s *deviceMgmtMgrSuite) TestDoQueueResponseSigningError(c *C) { }, }) - s.mgr.MockSigner(&mockSigner{ + s.mgr.MockBackend(&mockDeviceBackend{ + serial: s.makeSerial(c, "serial-1"), sign: func(_, _ string, _ asserts.MessageStatus, _ []byte) (*asserts.ResponseMessage, error) { return nil, fmt.Errorf("device key not found") }, @@ -1881,6 +2465,72 @@ func (s *deviceMgmtMgrSuite) TestDoQueueResponseSigningError(c *C) { c.Check(ms.Sequences["mesg"].Applied, Equals, 0) } +func (s *deviceMgmtMgrSuite) TestDoQueueResponseConcurrentWriteAfterResultFromChange(c *C) { + s.st.Lock() + defer s.st.Unlock() + + s.mockStore(func(_ context.Context, _ *store.MessageExchangeRequest) (*store.MessageExchangeResponse, error) { + return &store.MessageExchangeResponse{ + Messages: []store.MessageWithToken{ + s.makeStoreRequestMessage(c, "msg1", "test-kind", "token-1"), + s.makeStoreRequestMessage(c, "msg2", "test-kind", "token-2"), + }, + }, nil + }) + + firstCall := true + s.mgr.RegisterHandler("test-kind", &mockMessageHandler{ + apply: func(st *state.State, msg *devicemgmtstate.RequestMessage) (string, error) { + chg := st.NewChange("subsystem", "apply payload") + devicemgmtstate.MarkChangeForMessage(chg, msg) + return chg.ID(), nil + }, + resultFromChange: func(*state.Change) (map[string]any, error) { + if firstCall { + firstCall = false + // Wait for the other lane's full write before resuming. + lane2Done := make(chan struct{}) + s.st.AddTaskStatusChangedHandler(func(t *state.Task, _, new state.Status) (remove bool) { + if t.Kind() == "queue-mgmt-response" && new == state.DoneStatus { + close(lane2Done) + return true + } + return false + }) + + s.st.Unlock() + <-lane2Done + s.st.Lock() + } + + return map[string]any{"values": "ok"}, nil + }, + }) + + s.mgr.MockBackend(&mockDeviceBackend{ + serial: s.makeSerial(c, "serial-1"), + sign: func(accountID, messageID string, status asserts.MessageStatus, body []byte) (*asserts.ResponseMessage, error) { + return assertstest.FakeAssertionWithBody(body, map[string]any{ + "type": "response-message", + "account-id": accountID, + "message-id": messageID, + "device": "serial-1.my-model.my-brand", + "status": string(status), + "body-length": strconv.Itoa(len(body)), + }).(*asserts.ResponseMessage), nil + }, + }) + + s.settle(c) + + ms, err := s.mgr.GetState() + c.Assert(err, IsNil) + + c.Assert(ms.ReadyResponses, HasLen, 2) + c.Check(ms.ReadyResponses["msg1"].Format, Equals, "assertion") + c.Check(ms.ReadyResponses["msg2"].Format, Equals, "assertion") +} + func (s *deviceMgmtMgrSuite) TestParseRequestMessageInvalid(c *C) { type test struct { name string diff --git a/overlord/devicemgmtstate/export_test.go b/overlord/devicemgmtstate/export_test.go index 971afd44a0..887417d180 100644 --- a/overlord/devicemgmtstate/export_test.go +++ b/overlord/devicemgmtstate/export_test.go @@ -37,6 +37,10 @@ var ( MgmtMessageIDKey = mgmtMessageIDKey ) +func MockFetchAccountKey(f func(st *state.State, userID int, signKeyID string) error) func() { + return testutil.Mock(&assertstateFetchAccountKey, f) +} + func MockMaxSequences(n int) func() { return testutil.Mock(&maxSequences, n) } @@ -48,7 +52,7 @@ func MockMaxBlockedMessagesPerSequence(n int) func() { type SequenceState = sequenceState type DeviceMgmtState = deviceMgmtState -type ResponseMessageSigner = responseMessageSigner +type DeviceBackend = deviceBackend func (m *DeviceMgmtManager) GetState() (*DeviceMgmtState, error) { ms, err := m.getState() @@ -59,8 +63,8 @@ func (m *DeviceMgmtManager) SetState(ms *DeviceMgmtState) { m.setState(ms) } -func (m *DeviceMgmtManager) MockSigner(signer responseMessageSigner) { - m.signer = signer +func (m *DeviceMgmtManager) MockBackend(backend deviceBackend) { + m.device = backend } func (m *DeviceMgmtManager) ShouldExchangeMessages(ms *DeviceMgmtState) bool {