diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 360764abba6..d3eaf86b264 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -155,6 +155,7 @@ /packages/wallet/src/initialization/instances/passkey-controller/ @MetaMask/web3auth /packages/wallet/src/initialization/instances/remote-feature-flag-controller/ @MetaMask/extension-platform @MetaMask/mobile-platform @MetaMask/core-platform /packages/wallet/src/initialization/instances/seedless-onboarding-controller/ @MetaMask/web3auth +/packages/wallet/src/initialization/instances/shield-controller/ @MetaMask/web3auth /packages/wallet/src/initialization/instances/storage-service/ @MetaMask/extension-platform @MetaMask/mobile-platform @MetaMask/core-platform /packages/wallet/src/initialization/instances/transaction-controller/ @MetaMask/confirmations diff --git a/README.md b/README.md index e03df83f857..1dbd7e66bb7 100644 --- a/README.md +++ b/README.md @@ -656,6 +656,8 @@ linkStyle default opacity:0.5 wallet --> passkey_controller; wallet --> remote_feature_flag_controller; wallet --> seedless_onboarding_controller; + wallet --> shield_controller; + wallet --> signature_controller; wallet --> storage_service; wallet --> transaction_controller; wallet_cli --> base_controller; diff --git a/codeowners.ts b/codeowners.ts index 2a864d499ba..fcc07c1f5b1 100644 --- a/codeowners.ts +++ b/codeowners.ts @@ -310,6 +310,7 @@ const PACKAGES: Record = { }, 'shield-controller': { teams: ['@MetaMask/web3auth'], + initializationPath: 'shield-controller', }, 'signature-controller': { teams: ['@MetaMask/confirmations'], diff --git a/packages/shield-controller/src/backend.test.ts b/packages/shield-controller/src/backend.test.ts index ab22c5d90b8..8d322289d50 100644 --- a/packages/shield-controller/src/backend.test.ts +++ b/packages/shield-controller/src/backend.test.ts @@ -50,7 +50,7 @@ function setup({ getCoverageResultTimeout, getCoverageResultPollInterval, fetch, - baseUrl: 'https://rule-engine.metamask.io', + baseUrl: 'https://ruleset-engine.api.cx.metamask.io', captureException: mockCaptureException, }); diff --git a/packages/wallet-cli/CHANGELOG.md b/packages/wallet-cli/CHANGELOG.md index e70a1ee28fa..e23f86000d5 100644 --- a/packages/wallet-cli/CHANGELOG.md +++ b/packages/wallet-cli/CHANGELOG.md @@ -9,6 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added +- Wire the `shieldController` slot in the daemon wallet's instance options with the production rule-engine base URL and `fetch`, so the daemon constructs `ShieldController` with explicit host configuration rather than relying on implicit defaults ([#9616](https://github.com/MetaMask/core/pull/9616)) - Wire the `transactionController` slot in the daemon wallet's instance options, so the daemon runs the `TransactionController` with an explicit CLI-appropriate configuration (swaps processing disabled, no client hooks) rather than relying on the controller's implicit defaults ([#9509](https://github.com/MetaMask/core/pull/9509)) - Wire the `gasFeeController` slot in the daemon wallet's instance options, passing `clientId: 'cli'` so the CLI identifies itself to the gas estimation API, now that `@metamask/wallet` requires this option ([#9527](https://github.com/MetaMask/core/pull/9527)) - Add the `mm wallet unlock` command, which dispatches `KeyringController:submitPassword` over the daemon socket, allowing the keyring to be unlocked after a daemon start with no password or after a `mm daemon call KeyringController:setLocked` ([#8821](https://github.com/MetaMask/core/pull/8821)) diff --git a/packages/wallet-cli/src/daemon/wallet-factory.test.ts b/packages/wallet-cli/src/daemon/wallet-factory.test.ts index d50a2b4fd78..c922429769b 100644 --- a/packages/wallet-cli/src/daemon/wallet-factory.test.ts +++ b/packages/wallet-cli/src/daemon/wallet-factory.test.ts @@ -121,6 +121,12 @@ describe('createWallet', () => { ); expect(instanceOptions.transactionController?.disableSwaps).toBe(true); expect(instanceOptions.transactionController?.hooks).toStrictEqual({}); + expect(instanceOptions.shieldController.baseUrl).toBe( + 'https://ruleset-engine.api.cx.metamask.io', + ); + expect(instanceOptions.shieldController.fetchFunction).toBe( + globalThis.fetch, + ); expect(ClientConfigApiService).toHaveBeenCalled(); await dispose(); diff --git a/packages/wallet-cli/src/daemon/wallet-factory.ts b/packages/wallet-cli/src/daemon/wallet-factory.ts index 832c16a155a..6db7f887d7d 100644 --- a/packages/wallet-cli/src/daemon/wallet-factory.ts +++ b/packages/wallet-cli/src/daemon/wallet-factory.ts @@ -65,6 +65,9 @@ export type CreateWalletResult = { * - `transactionController` — swaps processing disabled and no client hooks; * see the slot's inline comment for why the daemon relies on the * controller's defaults for everything else. + * - `shieldController` — production rule-engine base URL and `fetch`; the daemon + * does not register `AuthenticationController` or `SignatureController`, so + * hosts must call `ShieldController:start` only after wiring those peers. * * The optional `keyringController` slot is intentionally omitted so the * controller's built-in defaults (e.g. the PBKDF2 encryptor) apply. @@ -121,6 +124,10 @@ function buildInstanceOptions( // the controller's default. hooks: {}, }, + shieldController: { + baseUrl: 'https://ruleset-engine.api.cx.metamask.io', + fetchFunction: globalThis.fetch, + }, }; } diff --git a/packages/wallet/CHANGELOG.md b/packages/wallet/CHANGELOG.md index 9eaa012500c..6a8e9d00729 100644 --- a/packages/wallet/CHANGELOG.md +++ b/packages/wallet/CHANGELOG.md @@ -7,6 +7,14 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- **BREAKING:** Wire `ShieldController` into the default wallet initialization ([#9616](https://github.com/MetaMask/core/pull/9616)) + - Adds required `shieldController` slot to `instanceOptions` with `baseUrl` and `fetchFunction` (or an injected `backend` override) + - Default backend construction uses `ShieldRemoteBackend` with optional `getAccessToken`, `captureException`, polling, history limits, and `normalizeSignatureRequest` + - Delegates `AuthenticationController:getBearerToken` plus `TransactionController:stateChange` and `SignatureController:stateChange` on the shared messenger bus + - Hosts must register `AuthenticationController` and `SignatureController` on the wallet root messenger and explicitly call `ShieldController:start` after wiring + ## [8.0.0] ### Added diff --git a/packages/wallet/package.json b/packages/wallet/package.json index e9ebc5e43ae..69fe467a194 100644 --- a/packages/wallet/package.json +++ b/packages/wallet/package.json @@ -70,6 +70,8 @@ "@metamask/remote-feature-flag-controller": "^4.2.2", "@metamask/scure-bip39": "^2.1.1", "@metamask/seedless-onboarding-controller": "^10.1.0", + "@metamask/shield-controller": "^5.1.3", + "@metamask/signature-controller": "^39.2.7", "@metamask/storage-service": "^1.0.2", "@metamask/transaction-controller": "^69.2.1", "@metamask/utils": "^11.11.0" diff --git a/packages/wallet/src/Wallet.test.ts b/packages/wallet/src/Wallet.test.ts index e370c699aa0..eba260cd4df 100644 --- a/packages/wallet/src/Wallet.test.ts +++ b/packages/wallet/src/Wallet.test.ts @@ -1,6 +1,10 @@ import { getDefaultAddressBookControllerState } from '@metamask/address-book-controller'; import { CONNECTIVITY_STATUSES } from '@metamask/connectivity-controller'; import { Messenger } from '@metamask/messenger'; +import { + getDefaultShieldControllerState, + ShieldController, +} from '@metamask/shield-controller'; import { InMemoryStorageAdapter } from '@metamask/storage-service'; import { Json } from '@metamask/utils'; import { webcrypto } from 'crypto'; @@ -8,6 +12,7 @@ import { webcrypto } from 'crypto'; import MockEncryptor from '../../keyring-controller/tests/mocks/mockEncryptor.js'; import * as initializationModule from './initialization/initialization.js'; import { AlwaysOnlineAdapter } from './initialization/instances/connectivity-controller/always-online-adapter.js'; +import type { WalletOptions } from './types.js'; import { importSecretRecoveryPhrase } from './utilities.js'; import { Wallet } from './Wallet.js'; @@ -23,23 +28,39 @@ const REMOTE_FEATURE_FLAG_OPTIONS = { }, }; +const SHIELD_CONTROLLER_OPTIONS = { + baseUrl: 'https://ruleset-engine.api.cx.metamask.io', + fetchFunction: jest.fn(), + backend: { + checkCoverage: jest.fn(), + checkSignatureCoverage: jest.fn(), + logSignature: jest.fn(), + logTransaction: jest.fn(), + }, +}; + +function getInstanceOptions(): WalletOptions['instanceOptions'] { + return { + connectivityController: { + connectivityAdapter: new AlwaysOnlineAdapter(), + }, + gasFeeController: { + clientId: 'test', + }, + networkController: { + infuraProjectId: 'fake-infura-project-id', + }, + storageService: { + storage: new InMemoryStorageAdapter(), + }, + remoteFeatureFlagController: REMOTE_FEATURE_FLAG_OPTIONS, + shieldController: SHIELD_CONTROLLER_OPTIONS, + }; +} + async function setupWallet(): Promise { const wallet = new Wallet({ - instanceOptions: { - connectivityController: { - connectivityAdapter: new AlwaysOnlineAdapter(), - }, - gasFeeController: { - clientId: 'test', - }, - networkController: { - infuraProjectId: 'fake-infura-project-id', - }, - storageService: { - storage: new InMemoryStorageAdapter(), - }, - remoteFeatureFlagController: REMOTE_FEATURE_FLAG_OPTIONS, - }, + instanceOptions: getInstanceOptions(), }); await importSecretRecoveryPhrase(wallet, TEST_PASSWORD, TEST_SRP); @@ -89,22 +110,10 @@ describe('Wallet', () => { it('supports passing instance options', async () => { const wallet = new Wallet({ instanceOptions: { - connectivityController: { - connectivityAdapter: new AlwaysOnlineAdapter(), - }, - gasFeeController: { - clientId: 'test', - }, + ...getInstanceOptions(), keyringController: { encryptor: new MockEncryptor(), }, - networkController: { - infuraProjectId: 'fake-infura-project-id', - }, - storageService: { - storage: new InMemoryStorageAdapter(), - }, - remoteFeatureFlagController: REMOTE_FEATURE_FLAG_OPTIONS, }, }); @@ -143,21 +152,7 @@ describe('Wallet', () => { init: (): DummyService => new DummyService(), }, ], - instanceOptions: { - connectivityController: { - connectivityAdapter: new AlwaysOnlineAdapter(), - }, - gasFeeController: { - clientId: 'test', - }, - networkController: { - infuraProjectId: 'fake-infura-project-id', - }, - storageService: { - storage: new InMemoryStorageAdapter(), - }, - remoteFeatureFlagController: REMOTE_FEATURE_FLAG_OPTIONS, - }, + instanceOptions: getInstanceOptions(), }); const { state } = wallet; @@ -189,21 +184,7 @@ describe('Wallet', () => { }); const wallet = new Wallet({ - instanceOptions: { - connectivityController: { - connectivityAdapter: new AlwaysOnlineAdapter(), - }, - gasFeeController: { - clientId: 'test', - }, - networkController: { - infuraProjectId: 'fake-infura-project-id', - }, - storageService: { - storage: new InMemoryStorageAdapter(), - }, - remoteFeatureFlagController: REMOTE_FEATURE_FLAG_OPTIONS, - }, + instanceOptions: getInstanceOptions(), }); expect(wallet.controllerMetadata).toStrictEqual({ @@ -302,21 +283,7 @@ describe('Wallet', () => { addressBook: { '0x1': { [ADDRESS]: entry } }, }, }, - instanceOptions: { - connectivityController: { - connectivityAdapter: new AlwaysOnlineAdapter(), - }, - gasFeeController: { - clientId: 'test', - }, - networkController: { - infuraProjectId: 'fake-infura-project-id', - }, - storageService: { - storage: new InMemoryStorageAdapter(), - }, - remoteFeatureFlagController: REMOTE_FEATURE_FLAG_OPTIONS, - }, + instanceOptions: getInstanceOptions(), }); expect( @@ -337,21 +304,7 @@ describe('Wallet', () => { describe('ConnectivityController', () => { it('reports online connectivity status', () => { const wallet = new Wallet({ - instanceOptions: { - connectivityController: { - connectivityAdapter: new AlwaysOnlineAdapter(), - }, - gasFeeController: { - clientId: 'test', - }, - networkController: { - infuraProjectId: 'fake-infura-project-id', - }, - storageService: { - storage: new InMemoryStorageAdapter(), - }, - remoteFeatureFlagController: REMOTE_FEATURE_FLAG_OPTIONS, - }, + instanceOptions: getInstanceOptions(), }); expect(wallet.state.ConnectivityController.connectivityStatus).toBe( @@ -380,21 +333,7 @@ describe('Wallet', () => { vault, }, }, - instanceOptions: { - connectivityController: { - connectivityAdapter: new AlwaysOnlineAdapter(), - }, - gasFeeController: { - clientId: 'test', - }, - networkController: { - infuraProjectId: 'fake-infura-project-id', - }, - storageService: { - storage: new InMemoryStorageAdapter(), - }, - remoteFeatureFlagController: REMOTE_FEATURE_FLAG_OPTIONS, - }, + instanceOptions: getInstanceOptions(), }); await wallet.messenger.call( @@ -454,6 +393,19 @@ describe('Wallet', () => { }); }); + describe('ShieldController', () => { + it('is wired and exposes its state on the wallet messenger', async () => { + const wallet = await setupWallet(); + + expect(wallet.getInstance('ShieldController')).toBeInstanceOf( + ShieldController, + ); + expect(wallet.messenger.call('ShieldController:getState')).toStrictEqual( + getDefaultShieldControllerState(), + ); + }); + }); + describe('RemoteFeatureFlagController', () => { it('is wired and exposes its state on the wallet messenger', async () => { const wallet = await setupWallet(); @@ -472,17 +424,8 @@ describe('Wallet', () => { it('routes injected instanceOptions through to the controller', async () => { const wallet = new Wallet({ instanceOptions: { - connectivityController: { - connectivityAdapter: new AlwaysOnlineAdapter(), - }, - gasFeeController: { - clientId: 'test', - }, - networkController: { - infuraProjectId: 'fake-infura-project-id', - }, + ...getInstanceOptions(), keyringController: { encryptor: new MockEncryptor() }, - storageService: { storage: new InMemoryStorageAdapter() }, remoteFeatureFlagController: { clientConfigApiService: { fetchRemoteFeatureFlags: async (): Promise<{ diff --git a/packages/wallet/src/initialization/instances/index.ts b/packages/wallet/src/initialization/instances/index.ts index d45bb0917b2..ec4b987ea89 100644 --- a/packages/wallet/src/initialization/instances/index.ts +++ b/packages/wallet/src/initialization/instances/index.ts @@ -8,5 +8,6 @@ export { networkController } from './network-controller/network-controller.js'; export { passkeyController } from './passkey-controller/passkey-controller.js'; export { remoteFeatureFlagController } from './remote-feature-flag-controller/remote-feature-flag-controller.js'; export { seedlessOnboardingController } from './seedless-onboarding-controller/seedless-onboarding-controller.js'; +export { shieldController } from './shield-controller/shield-controller.js'; export { storageService } from './storage-service/storage-service.js'; export { transactionController } from './transaction-controller/transaction-controller.js'; diff --git a/packages/wallet/src/initialization/instances/shield-controller/shield-controller.test.ts b/packages/wallet/src/initialization/instances/shield-controller/shield-controller.test.ts new file mode 100644 index 00000000000..82d66894ee9 --- /dev/null +++ b/packages/wallet/src/initialization/instances/shield-controller/shield-controller.test.ts @@ -0,0 +1,366 @@ +import { Messenger } from '@metamask/messenger'; +import { + getDefaultShieldControllerState, + ShieldController, +} from '@metamask/shield-controller'; +import type { TransactionControllerState } from '@metamask/transaction-controller'; +import { TransactionStatus } from '@metamask/transaction-controller'; + +import { defaultConfigurations } from '../../defaults.js'; +import type { + DefaultActions, + DefaultEvents, + RootMessenger, +} from '../../defaults.js'; +import { shieldController } from './shield-controller.js'; +import type { ShieldBackend } from './types.js'; + +const MOCK_COVERAGE_ID = 'coverage-id-1'; +const SHIELD_BASE_URL = 'https://ruleset-engine.api.cx.metamask.io'; + +type ActionHandler = (...args: unknown[]) => unknown; + +type AnyMessenger = Messenger; + +const SHIELD_OPTIONS = { + baseUrl: SHIELD_BASE_URL, + fetchFunction: globalThis.fetch, +}; + +function getRootMessenger(): RootMessenger { + return new Messenger({ namespace: 'Root' }); +} + +function registerActionHandler( + parent: RootMessenger, + namespace: string, + actionType: string, + handler: ActionHandler, +): void { + const messenger = new Messenger({ + namespace, + parent: parent as unknown as AnyMessenger, + }); + + ( + messenger as unknown as { + registerActionHandler(type: string, handler: ActionHandler): void; + } + ).registerActionHandler(actionType, handler); +} + +function createMockBackend(): jest.Mocked { + return { + checkCoverage: jest.fn().mockResolvedValue({ + coverageId: MOCK_COVERAGE_ID, + status: 'covered', + metrics: {}, + }), + checkSignatureCoverage: jest.fn().mockResolvedValue({ + coverageId: MOCK_COVERAGE_ID, + status: 'covered', + metrics: {}, + }), + logSignature: jest.fn(), + logTransaction: jest.fn(), + }; +} + +function createMockSignatureRequest(): Parameters< + ShieldController['checkSignatureCoverage'] +>[0] { + return { + chainId: '0x1', + id: 'signature-request-1', + type: 'personal_sign', + messageParams: { + data: '0x00', + from: '0x0000000000000000000000000000000000000000', + }, + networkClientId: 'mainnet', + status: 'unapproved', + time: Date.now(), + }; +} + +describe('shieldController', () => { + it('is registered as a default initialization configuration', () => { + expect(Object.values(defaultConfigurations)).toContain(shieldController); + }); + + it('initializes a ShieldController with default state', () => { + const messenger = shieldController.getMessenger(getRootMessenger()); + + const instance = shieldController.init({ + state: undefined, + messenger, + options: SHIELD_OPTIONS, + }); + + expect(instance).toBeInstanceOf(ShieldController); + expect(instance.state).toStrictEqual(getDefaultShieldControllerState()); + }); + + it('forwards the provided state to the controller', () => { + const messenger = shieldController.getMessenger(getRootMessenger()); + + const instance = shieldController.init({ + state: { + orderedTransactionHistory: ['tx-1'], + }, + messenger, + options: SHIELD_OPTIONS, + }); + + expect(instance.state.orderedTransactionHistory).toStrictEqual(['tx-1']); + }); + + it('uses a provided backend override', () => { + const messenger = shieldController.getMessenger(getRootMessenger()); + const mockBackend = createMockBackend(); + + const instance = shieldController.init({ + state: undefined, + messenger, + options: { + ...SHIELD_OPTIONS, + backend: mockBackend, + }, + }); + + expect(instance).toBeInstanceOf(ShieldController); + }); + + it('forwards transactionHistoryLimit and coverageHistoryLimit', () => { + const messenger = shieldController.getMessenger(getRootMessenger()); + const mockBackend = createMockBackend(); + + const instance = shieldController.init({ + state: undefined, + messenger, + options: { + ...SHIELD_OPTIONS, + backend: mockBackend, + transactionHistoryLimit: 5, + coverageHistoryLimit: 2, + }, + }); + + expect(instance).toBeInstanceOf(ShieldController); + }); + + it('forwards normalizeSignatureRequest to the controller', async () => { + const rootMessenger = getRootMessenger(); + const messenger = shieldController.getMessenger(rootMessenger); + const mockBackend = createMockBackend(); + const signatureRequest = createMockSignatureRequest(); + const normalizedSignatureRequest = { + ...signatureRequest, + messageParams: { + ...signatureRequest.messageParams, + data: 'normalized data', + }, + }; + const normalizeSignatureRequest = jest + .fn() + .mockReturnValue(normalizedSignatureRequest); + + const instance = shieldController.init({ + state: undefined, + messenger, + options: { + ...SHIELD_OPTIONS, + backend: mockBackend, + normalizeSignatureRequest, + }, + }); + + await instance.checkSignatureCoverage(signatureRequest); + + expect(normalizeSignatureRequest).toHaveBeenCalledWith(signatureRequest); + expect(mockBackend.checkSignatureCoverage).toHaveBeenCalledWith({ + signatureRequest: normalizedSignatureRequest, + }); + }); + + it('wires default getAccessToken to AuthenticationController:getBearerToken', async () => { + const rootMessenger = getRootMessenger(); + registerActionHandler( + rootMessenger, + 'AuthenticationController', + 'AuthenticationController:getBearerToken', + async () => 'test-bearer-token', + ); + const messenger = shieldController.getMessenger(rootMessenger); + const fetchFunction = jest.fn(async () => { + const callCount = fetchFunction.mock.calls.length; + if (callCount === 1) { + return new globalThis.Response( + JSON.stringify({ coverageId: MOCK_COVERAGE_ID }), + { status: 200 }, + ); + } + + return new globalThis.Response( + JSON.stringify({ + status: 'covered', + metrics: {}, + }), + { status: 200 }, + ); + }); + + const instance = shieldController.init({ + state: undefined, + messenger, + options: { + baseUrl: SHIELD_BASE_URL, + fetchFunction, + }, + }); + + await instance.checkCoverage({ + id: 'tx-1', + chainId: '0x1', + status: TransactionStatus.Unapproved, + time: Date.now(), + txParams: { + from: '0x0000000000000000000000000000000000000000', + }, + } as never); + + expect(fetchFunction).toHaveBeenCalled(); + const firstCall = fetchFunction.mock.calls[0] as unknown as [ + string, + RequestInit, + ]; + const [, requestInit] = firstCall; + const headers = new globalThis.Headers(requestInit.headers); + expect(headers.get('Authorization')).toBe('Bearer test-bearer-token'); + }); + + it('uses a provided getAccessToken override', async () => { + const rootMessenger = getRootMessenger(); + const messenger = shieldController.getMessenger(rootMessenger); + const fetchFunction = jest.fn(async () => { + const callCount = fetchFunction.mock.calls.length; + if (callCount === 1) { + return new globalThis.Response( + JSON.stringify({ coverageId: MOCK_COVERAGE_ID }), + { status: 200 }, + ); + } + + return new globalThis.Response( + JSON.stringify({ + status: 'covered', + metrics: {}, + }), + { status: 200 }, + ); + }); + const getAccessToken = jest.fn().mockResolvedValue('override-token'); + + const instance = shieldController.init({ + state: undefined, + messenger, + options: { + baseUrl: SHIELD_BASE_URL, + fetchFunction, + getAccessToken, + }, + }); + + await instance.checkCoverage({ + id: 'tx-1', + chainId: '0x1', + status: TransactionStatus.Unapproved, + time: Date.now(), + txParams: { + from: '0x0000000000000000000000000000000000000000', + }, + } as never); + + expect(getAccessToken).toHaveBeenCalled(); + const firstCall = fetchFunction.mock.calls[0] as unknown as [ + string, + RequestInit, + ]; + const [, requestInit] = firstCall; + const headers = new globalThis.Headers(requestInit.headers); + expect(headers.get('Authorization')).toBe('Bearer override-token'); + }); + + it('delegates AuthenticationController:getBearerToken and controller state-change events', () => { + const parent = getRootMessenger(); + const delegateSpy = jest.spyOn(parent, 'delegate'); + const messenger = shieldController.getMessenger(parent); + + expect(delegateSpy).toHaveBeenCalledWith({ + messenger, + actions: ['AuthenticationController:getBearerToken'], + events: [ + 'TransactionController:stateChange', + 'SignatureController:stateChange', + ], + }); + }); + + it('exposes its actions through the root messenger', () => { + const rootMessenger = getRootMessenger(); + const messenger = shieldController.getMessenger(rootMessenger); + + shieldController.init({ + state: undefined, + messenger, + options: { + ...SHIELD_OPTIONS, + backend: createMockBackend(), + }, + }); + + expect(rootMessenger.call('ShieldController:getState')).toStrictEqual( + getDefaultShieldControllerState(), + ); + }); + + it('does not auto-start on initialization', () => { + const rootMessenger = getRootMessenger(); + const messenger = shieldController.getMessenger(rootMessenger); + const mockBackend = createMockBackend(); + + shieldController.init({ + state: undefined, + messenger, + options: { + ...SHIELD_OPTIONS, + backend: mockBackend, + }, + }); + + const transactionMessenger = new Messenger({ + namespace: 'TransactionController', + parent: rootMessenger as unknown as AnyMessenger, + }); + + transactionMessenger.publish( + 'TransactionController:stateChange', + { + transactions: [ + { + id: 'tx-1', + chainId: '0x1', + status: TransactionStatus.Unapproved, + time: Date.now(), + txParams: { + from: '0x0000000000000000000000000000000000000000', + }, + }, + ], + } as TransactionControllerState, + undefined as never, + ); + + expect(mockBackend.checkCoverage).not.toHaveBeenCalled(); + }); +}); diff --git a/packages/wallet/src/initialization/instances/shield-controller/shield-controller.ts b/packages/wallet/src/initialization/instances/shield-controller/shield-controller.ts new file mode 100644 index 00000000000..58d606c1121 --- /dev/null +++ b/packages/wallet/src/initialization/instances/shield-controller/shield-controller.ts @@ -0,0 +1,88 @@ +import { Messenger } from '@metamask/messenger'; +import type { ShieldControllerMessenger } from '@metamask/shield-controller'; +import { + ShieldController, + ShieldRemoteBackend, +} from '@metamask/shield-controller'; +import type { ShieldControllerState } from '@metamask/shield-controller'; + +import type { InitializationConfiguration } from '../../types.js'; +import type { + ShieldBackend, + ShieldControllerInitializationMessenger, + ShieldControllerInstanceOptions, +} from './types.js'; + +export type { + ShieldControllerInitializationMessenger, + ShieldControllerInstanceOptions, +} from './types.js'; + +function resolveShieldBackend( + messenger: ShieldControllerInitializationMessenger, + options: ShieldControllerInstanceOptions, +): ShieldBackend { + if (options.backend) { + return options.backend; + } + + const getAccessToken = + options.getAccessToken ?? + ((): Promise => + messenger.call('AuthenticationController:getBearerToken')); + + return new ShieldRemoteBackend({ + baseUrl: options.baseUrl, + fetch: options.fetchFunction, + getAccessToken, + captureException: options.captureException, + getCoverageResultTimeout: options.getCoverageResultTimeout, + getCoverageResultPollInterval: options.getCoverageResultPollInterval, + }); +} + +export const shieldController: InitializationConfiguration< + ShieldController, + ShieldControllerInitializationMessenger +> = { + name: 'ShieldController', + init: ({ + state, + messenger, + options, + }: { + state: Partial | undefined; + messenger: ShieldControllerInitializationMessenger; + options: ShieldControllerInstanceOptions; + }) => { + return new ShieldController({ + messenger: messenger as unknown as ShieldControllerMessenger, + state, + backend: resolveShieldBackend(messenger, options), + transactionHistoryLimit: options.transactionHistoryLimit, + coverageHistoryLimit: options.coverageHistoryLimit, + normalizeSignatureRequest: options.normalizeSignatureRequest, + }); + }, + getMessenger: (parent) => { + const messenger: ShieldControllerInitializationMessenger = new Messenger({ + namespace: 'ShieldController', + parent, + }); + + parent.delegate({ + messenger, + actions: ['AuthenticationController:getBearerToken'], + events: [ + // ShieldController subscribes to :stateChange internally; the + // delegation must match until those controllers migrate to :stateChanged. + // eslint-disable-next-line no-restricted-syntax + 'TransactionController:stateChange', + // eslint-disable-next-line no-restricted-syntax + 'SignatureController:stateChange', + ], + }); + + return messenger; + }, +}; diff --git a/packages/wallet/src/initialization/instances/shield-controller/types.ts b/packages/wallet/src/initialization/instances/shield-controller/types.ts new file mode 100644 index 00000000000..cdf70454288 --- /dev/null +++ b/packages/wallet/src/initialization/instances/shield-controller/types.ts @@ -0,0 +1,45 @@ +import type { Messenger } from '@metamask/messenger'; +import type { + NormalizeSignatureRequestFn, + ShieldControllerActions, + ShieldControllerEvents, + ShieldRemoteBackend, +} from '@metamask/shield-controller'; +import type { SignatureStateChange } from '@metamask/signature-controller'; +import type { TransactionControllerStateChangeEvent } from '@metamask/transaction-controller'; + +export type ShieldBackend = Pick< + ShieldRemoteBackend, + 'checkCoverage' | 'checkSignatureCoverage' | 'logSignature' | 'logTransaction' +>; + +type AuthenticationControllerGetBearerTokenAction = { + type: 'AuthenticationController:getBearerToken'; + handler: (entropySourceId?: string) => Promise; +}; + +export type ShieldControllerInitializationMessenger = Messenger< + 'ShieldController', + ShieldControllerActions | AuthenticationControllerGetBearerTokenAction, + | ShieldControllerEvents + | SignatureStateChange + | TransactionControllerStateChangeEvent +>; + +export type ShieldControllerInstanceOptions = { + /** + * When set, used as-is; `baseUrl`, `fetchFunction`, `getAccessToken`, and + * `captureException` are ignored for backend construction. + */ + backend?: ShieldBackend; + /** Required when building the default `ShieldRemoteBackend`. */ + baseUrl: string; + fetchFunction: typeof fetch; + getAccessToken?: () => Promise; + captureException?: (error: Error) => void; + getCoverageResultTimeout?: number; + getCoverageResultPollInterval?: number; + transactionHistoryLimit?: number; + coverageHistoryLimit?: number; + normalizeSignatureRequest?: NormalizeSignatureRequestFn; +}; diff --git a/packages/wallet/src/initialization/instances/transaction-controller/transaction-controller.test.ts b/packages/wallet/src/initialization/instances/transaction-controller/transaction-controller.test.ts index a3831f1b973..587a0fa2cda 100644 --- a/packages/wallet/src/initialization/instances/transaction-controller/transaction-controller.test.ts +++ b/packages/wallet/src/initialization/instances/transaction-controller/transaction-controller.test.ts @@ -125,6 +125,16 @@ function getInstanceOptions(): WalletOptions['instanceOptions'] { storage: new InMemoryStorageAdapter(), }, remoteFeatureFlagController: REMOTE_FEATURE_FLAG_OPTIONS, + shieldController: { + baseUrl: 'https://ruleset-engine.api.cx.metamask.io', + fetchFunction: jest.fn(), + backend: { + checkCoverage: jest.fn(), + checkSignatureCoverage: jest.fn(), + logSignature: jest.fn(), + logTransaction: jest.fn(), + }, + }, }; } diff --git a/packages/wallet/src/types.ts b/packages/wallet/src/types.ts index 978bf1b1afd..d927cd61d71 100644 --- a/packages/wallet/src/types.ts +++ b/packages/wallet/src/types.ts @@ -13,6 +13,7 @@ import type { NetworkControllerInstanceOptions } from './initialization/instance import type { PasskeyControllerInstanceOptions } from './initialization/instances/passkey-controller/types.js'; import type { RemoteFeatureFlagControllerInstanceOptions } from './initialization/instances/remote-feature-flag-controller/types.js'; import type { SeedlessOnboardingControllerInstanceOptions } from './initialization/instances/seedless-onboarding-controller/types.js'; +import type { ShieldControllerInstanceOptions } from './initialization/instances/shield-controller/types.js'; import type { StorageServiceInstanceOptions } from './initialization/instances/storage-service/types.js'; import type { TransactionControllerInstanceOptions } from './initialization/instances/transaction-controller/types.js'; import type { InitializationConfiguration } from './initialization/types.js'; @@ -38,4 +39,5 @@ export type InstanceSpecificOptions = { transactionController?: TransactionControllerInstanceOptions; passkeyController?: PasskeyControllerInstanceOptions; seedlessOnboardingController?: SeedlessOnboardingControllerInstanceOptions; + shieldController: ShieldControllerInstanceOptions; }; diff --git a/packages/wallet/tsconfig.build.json b/packages/wallet/tsconfig.build.json index dae6d88285a..13467e2a1c2 100644 --- a/packages/wallet/tsconfig.build.json +++ b/packages/wallet/tsconfig.build.json @@ -19,6 +19,8 @@ { "path": "../passkey-controller/tsconfig.build.json" }, { "path": "../remote-feature-flag-controller/tsconfig.build.json" }, { "path": "../seedless-onboarding-controller/tsconfig.build.json" }, + { "path": "../shield-controller/tsconfig.build.json" }, + { "path": "../signature-controller/tsconfig.build.json" }, { "path": "../storage-service/tsconfig.build.json" }, { "path": "../transaction-controller/tsconfig.build.json" } ], diff --git a/packages/wallet/tsconfig.json b/packages/wallet/tsconfig.json index 10221177468..c4e99c0c861 100644 --- a/packages/wallet/tsconfig.json +++ b/packages/wallet/tsconfig.json @@ -43,6 +43,12 @@ { "path": "../seedless-onboarding-controller" }, + { + "path": "../shield-controller" + }, + { + "path": "../signature-controller" + }, { "path": "../storage-service" }, diff --git a/yarn.lock b/yarn.lock index 97d8dc6c6d9..26015d30233 100644 --- a/yarn.lock +++ b/yarn.lock @@ -8617,7 +8617,7 @@ __metadata: languageName: unknown linkType: soft -"@metamask/shield-controller@workspace:packages/shield-controller": +"@metamask/shield-controller@npm:^5.1.3, @metamask/shield-controller@workspace:packages/shield-controller": version: 0.0.0-use.local resolution: "@metamask/shield-controller@workspace:packages/shield-controller" dependencies: @@ -9270,6 +9270,8 @@ __metadata: "@metamask/remote-feature-flag-controller": "npm:^4.2.2" "@metamask/scure-bip39": "npm:^2.1.1" "@metamask/seedless-onboarding-controller": "npm:^10.1.0" + "@metamask/shield-controller": "npm:^5.1.3" + "@metamask/signature-controller": "npm:^39.2.7" "@metamask/storage-service": "npm:^1.0.2" "@metamask/transaction-controller": "npm:^69.2.1" "@metamask/utils": "npm:^11.11.0"