diff --git a/packages/transaction-controller/CHANGELOG.md b/packages/transaction-controller/CHANGELOG.md index 6bd369e5234..94d8adb7e4e 100644 --- a/packages/transaction-controller/CHANGELOG.md +++ b/packages/transaction-controller/CHANGELOG.md @@ -11,6 +11,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Bump `uuid` from `^8.3.2` to `^9.0.1` ([#10117](https://github.com/MetaMask/core/pull/10117)) - Bump `@metamask/core-backend` from `^10.0.0` to `^10.0.1` ([#10166](https://github.com/MetaMask/core/pull/10166)) +- Add approval-time sponsorship/signing hooks to `TransactionController` and keep gas-fee-token preflight in the approval flow ([#10109](https://github.com/MetaMask/core/pull/10109)) ## [70.0.0] diff --git a/packages/transaction-controller/src/TransactionController.test.ts b/packages/transaction-controller/src/TransactionController.test.ts index b32d8f575f5..3ef38d9de73 100644 --- a/packages/transaction-controller/src/TransactionController.test.ts +++ b/packages/transaction-controller/src/TransactionController.test.ts @@ -964,6 +964,39 @@ describe('TransactionController', () => { ); }); + it('updates transaction batch gas fee estimates when the poller emits a batch update', async () => { + const batchId = BATCH_ID_MOCK; + const { controller } = setupController({ + options: { + state: { + transactionBatches: [{ id: batchId } as never], + }, + }, + }); + const batchUpdateHandler = gasFeePollerMock.hub.on.mock.calls.find( + ([event]) => event === 'transaction-batch-updated', + )?.[1] as (request: { + transactionBatchId: Hex; + gasFeeEstimates?: GasFeeEstimates; + }) => void; + + batchUpdateHandler({ + transactionBatchId: batchId, + gasFeeEstimates: { + type: GasFeeEstimateType.FeeMarket, + } as GasFeeEstimates, + }); + + expect(controller.state.transactionBatches).toContainEqual( + expect.objectContaining({ + id: batchId, + gasFeeEstimates: { + type: GasFeeEstimateType.FeeMarket, + }, + }), + ); + }); + it('provides only test flow if option set', () => { setupController({ options: { @@ -2409,6 +2442,143 @@ describe('TransactionController', () => { }); }); + describe('with sponsored approval hooks', () => { + it('calls isSponsored hook before reserving a nonce', async () => { + const callOrder: string[] = []; + + const isSponsoredHook = jest + .fn() + .mockImplementation(async (): Promise => { + callOrder.push('isSponsored'); + expect(getNonceLockSpy).not.toHaveBeenCalled(); + return false; + }); + + const shouldSignHook = jest + .fn() + .mockImplementation(async (): Promise => { + callOrder.push('shouldSign'); + expect(getNonceLockSpy).not.toHaveBeenCalled(); + return true; + }); + + getNonceLockSpy.mockImplementation( + async (): Promise<{ + nextNonce: Hex; + releaseLock: () => Promise; + }> => { + callOrder.push('getNonceLock'); + return { + nextNonce: NONCE_MOCK, + releaseLock: () => Promise.resolve(), + }; + }, + ); + + const { controller } = setupController({ + messengerOptions: { + addTransactionApprovalRequest: { + state: 'approved', + }, + }, + options: { + hooks: { + isSponsored: isSponsoredHook, + shouldSign: shouldSignHook, + }, + }, + }); + + await controller.addTransaction( + { + from: ACCOUNT_MOCK, + to: ACCOUNT_MOCK, + }, + { + networkClientId: NETWORK_CLIENT_ID_MOCK, + }, + ); + + await flushPromises(); + + expect(isSponsoredHook).toHaveBeenCalledTimes(1); + expect(shouldSignHook).toHaveBeenCalledTimes(1); + expect(callOrder).toStrictEqual([ + 'isSponsored', + 'shouldSign', + 'getNonceLock', + ]); + }); + + it('skips nonce reservation when shouldSign resolves false', async () => { + const isSponsoredHook = jest.fn().mockResolvedValue(false); + const shouldSignHook = jest.fn().mockResolvedValue(false); + + const { controller } = setupController({ + messengerOptions: { + addTransactionApprovalRequest: { + state: 'approved', + }, + }, + options: { + hooks: { + isSponsored: isSponsoredHook, + shouldSign: shouldSignHook, + }, + }, + }); + + await controller.addTransaction( + { + from: ACCOUNT_MOCK, + to: ACCOUNT_MOCK, + }, + { + networkClientId: NETWORK_CLIENT_ID_MOCK, + }, + ); + + await flushPromises(); + + expect(isSponsoredHook).toHaveBeenCalledTimes(1); + expect(shouldSignHook).toHaveBeenCalledTimes(1); + expect(getNonceLockSpy).not.toHaveBeenCalled(); + }); + + it('still runs beforeSign when shouldSign resolves false', async () => { + const beforeSignHook = jest.fn().mockResolvedValueOnce({}); + + const { controller } = setupController({ + messengerOptions: { + addTransactionApprovalRequest: { + state: 'approved', + }, + }, + options: { + hooks: { + beforeSign: beforeSignHook, + shouldSign: jest.fn().mockResolvedValue(false), + }, + }, + }); + + await controller.addTransaction( + { + from: ACCOUNT_MOCK, + to: ACCOUNT_MOCK, + }, + { + networkClientId: NETWORK_CLIENT_ID_MOCK, + }, + ); + + await flushPromises(); + + expect(beforeSignHook).toHaveBeenCalledTimes(1); + expect(getNonceLockSpy).not.toHaveBeenCalled(); + }); + }); + describe('with beforeSign hook', () => { it('calls beforeSign hook', async () => { const beforeSignHook = jest.fn().mockResolvedValueOnce({}); @@ -3792,6 +3962,35 @@ describe('TransactionController', () => { providerErrors.unauthorized({ data: { origin: expectedOrigin } }), ); }); + + it('reads internal accounts while validating an approved transaction', async () => { + const { controller, rootMessenger } = setupController({ + messengerOptions: { + addTransactionApprovalRequest: { + state: 'approved', + }, + }, + }); + + rootMessenger.unregisterActionHandler('AccountsController:getState'); + rootMessenger.registerActionHandler( + 'AccountsController:getState', + () => ({ + internalAccounts: { + accounts: { + [INTERNAL_ACCOUNT_MOCK.id]: INTERNAL_ACCOUNT_MOCK, + }, + }, + }), + ); + + const { result } = await controller.addTransaction( + { from: ACCOUNT_MOCK, to: ACCOUNT_MOCK }, + { networkClientId: NETWORK_CLIENT_ID_MOCK }, + ); + + await result; + }); }); describe('updates submit history', () => { @@ -9040,6 +9239,40 @@ describe('TransactionController', () => { expect(approvedEventListener).not.toHaveBeenCalled(); }); + + it('publishes transactionApproved with a nonce after signing approval', async () => { + const { controller, messenger, mockTransactionApprovalRequest } = + setupController(); + + const approvedEventListener = jest.fn(); + + messenger.subscribe( + 'TransactionController:transactionApproved', + approvedEventListener, + ); + + const { result } = await controller.addTransaction( + { + from: ACCOUNT_MOCK, + gas: '0x21000', + gasPrice: '0x1', + to: ACCOUNT_MOCK, + value: '0x0', + }, + { + networkClientId: NETWORK_CLIENT_ID_MOCK, + }, + ); + + mockTransactionApprovalRequest.approve(); + + await result; + + expect(approvedEventListener).toHaveBeenCalledTimes(1); + expect( + approvedEventListener.mock.calls[0][0].transactionMeta.txParams.nonce, + ).toBeDefined(); + }); }); describe('TransactionController:estimateGasBatch', () => { diff --git a/packages/transaction-controller/src/TransactionController.ts b/packages/transaction-controller/src/TransactionController.ts index 778d5f56e36..aa20e434404 100644 --- a/packages/transaction-controller/src/TransactionController.ts +++ b/packages/transaction-controller/src/TransactionController.ts @@ -124,6 +124,8 @@ import type { AfterAddHook, GasFeeEstimateLevel as GasFeeEstimateLevelType, TransactionBatchMeta, + IsSponsoredHook, + ShouldSignHook, BeforeSignHook, GetSimulationConfig, AddTransactionOptions, @@ -413,6 +415,16 @@ export type TransactionControllerOptions = { transactionMeta: TransactionMeta, ) => Promise; + /** + * Additional logic to determine whether a transaction is sponsored. + */ + isSponsored?: IsSponsoredHook; + + /** + * Additional logic to determine whether a transaction should be signed locally. + */ + shouldSign?: ShouldSignHook; + /** * Additional logic to execute before publishing a transaction. * Return false to prevent the broadcast of the transaction. @@ -732,6 +744,10 @@ export class TransactionController extends BaseController< transactionMeta: TransactionMeta, ) => Promise; + readonly #isSponsored: IsSponsoredHook; + + readonly #shouldSign: ShouldSignHook; + readonly #beforePublish: ( transactionMeta: TransactionMeta, ) => Promise; @@ -833,6 +849,22 @@ export class TransactionController extends BaseController< /* istanbul ignore next */ hooks?.beforeCheckPendingTransaction ?? ((): Promise => Promise.resolve(true)); + this.#isSponsored = + hooks?.isSponsored ?? + (async ({ + transactionMeta, + }: { + transactionMeta: TransactionMeta; + }): Promise => Boolean(transactionMeta.isGasFeeSponsored)); + this.#shouldSign = + hooks?.shouldSign ?? + (async ({ + transactionMeta, + isSponsored: _isSponsored, + }: { + transactionMeta: TransactionMeta; + isSponsored: boolean; + }): Promise => !transactionMeta.isExternalSign); this.#beforePublish = hooks?.beforePublish ?? ((): Promise => Promise.resolve(true)); this.#beforeSign = @@ -1240,7 +1272,8 @@ export class TransactionController extends BaseController< } else { const newTransactionMeta = cloneDeep(addedTransactionMeta); - this.#updateGasProperties(newTransactionMeta) + // eslint-disable-next-line no-void + void this.#updateGasProperties(newTransactionMeta) .then(() => { this.#updateTransactionInternal( { @@ -1276,7 +1309,8 @@ export class TransactionController extends BaseController< this.#addMetadata(addedTransactionMeta); - delegationAddressPromise + // eslint-disable-next-line no-void + void delegationAddressPromise .then((delegationAddress) => { this.#updateTransactionInternal( { @@ -2386,21 +2420,20 @@ export class TransactionController extends BaseController< pickBy(transactionsToFilter, (transaction) => { // iterate over the predicateMethods keys to check if the transaction // matches the searchCriteria + const txParams = transaction.txParams as Record; + const txMeta = transaction as Record; + for (const [key, predicate] of Object.entries(predicateMethods)) { // We return false early as soon as we know that one of the specified // search criteria do not match the transaction. This prevents // needlessly checking all criteria when we already know the criteria // are not fully satisfied. We check both txParams and the base // object as predicate keys can be either. - if (key in transaction.txParams) { - // TODO: Replace `any` with type - // eslint-disable-next-line @typescript-eslint/no-explicit-any - if (predicate((transaction.txParams as any)[key]) === false) { + if (key in txParams) { + if (predicate(txParams[key]) === false) { return false; } - // TODO: Replace `any` with type - // eslint-disable-next-line @typescript-eslint/no-explicit-any - } else if (predicate((transaction as any)[key]) === false) { + } else if (predicate(txMeta[key]) === false) { return false; } } @@ -3112,20 +3145,6 @@ export class TransactionController extends BaseController< clearApprovingTransactionId = (): boolean => this.#approvingTransactionIds.delete(transactionId); - const { networkClientId } = transactionMeta; - - const [nonce, releaseNonce] = await getNextNonce( - transactionMeta, - (address: string) => - this.#multichainTrackingHelper.getNonceLock( - address, - transactionMeta.networkClientId, - ), - ); - - clearNonceLock = releaseNonce; - - // eslint-disable-next-line require-atomic-updates transactionMeta = this.#updateTransactionInternal( { transactionId, @@ -3137,7 +3156,6 @@ export class TransactionController extends BaseController< draftTxMeta.status = TransactionStatus.approved; draftTxMeta.txParams.chainId = chainId; draftTxMeta.txParams.gasLimit = gas; - draftTxMeta.txParams.nonce = nonce; if (!type && isEIP1559Transaction(txParams)) { draftTxMeta.txParams.type = TransactionEnvelopeType.feeMarket; @@ -3145,16 +3163,80 @@ export class TransactionController extends BaseController< }, ); - this.#onTransactionStatusChange(transactionMeta); + // eslint-disable-next-line require-atomic-updates + transactionMeta = await this.#applyBeforeSignHook(transactionMeta); - const rawTx = await this.#trace( - { name: 'Sign', parentContext: traceContext }, - () => this.#signTransaction(transactionMeta), - ); + const { networkClientId } = transactionMeta; + + await checkGasFeeTokenBeforePublish({ + messenger: this.messenger, + networkClientId, + fetchGasFeeTokens: async (tx) => + (await this.#getGasFeeTokens(tx)).gasFeeTokens, + transaction: transactionMeta, + updateTransaction: (txId, fn) => + this.#updateTransactionInternal({ transactionId: txId }, fn), + }); // eslint-disable-next-line require-atomic-updates transactionMeta = this.#getTransactionOrThrow(transactionId); + const isSponsored = await this.#isSponsored({ transactionMeta }); + const shouldSign = await this.#shouldSign({ + transactionMeta, + isSponsored, + }); + + // eslint-disable-next-line require-atomic-updates + transactionMeta = this.#updateTransactionInternal( + { + transactionId, + }, + (draftTxMeta) => { + draftTxMeta.isGasFeeSponsored = isSponsored; + draftTxMeta.isExternalSign = !shouldSign; + + if (!shouldSign) { + draftTxMeta.txParams.nonce = undefined; + } + }, + ); + + let rawTx: string | undefined; + + if (shouldSign) { + const [nonce, releaseNonce] = await getNextNonce( + transactionMeta, + (address: string) => + this.#multichainTrackingHelper.getNonceLock( + address, + transactionMeta.networkClientId, + ), + ); + + clearNonceLock = releaseNonce; + + // eslint-disable-next-line require-atomic-updates + transactionMeta = this.#updateTransactionInternal( + { + transactionId, + }, + (draftTxMeta) => { + draftTxMeta.txParams.nonce = nonce; + }, + ); + + rawTx = await this.#trace( + { name: 'Sign', parentContext: traceContext }, + () => this.#signTransaction(transactionMeta, true, true), + ); + + // eslint-disable-next-line require-atomic-updates + transactionMeta = this.#getTransactionOrThrow(transactionId); + } + + this.#onTransactionStatusChange(transactionMeta); + if (!(await this.#beforePublish(transactionMeta))) { log('Skipping publishing transaction based on hook'); this.messenger.publish( @@ -3671,10 +3753,9 @@ export class TransactionController extends BaseController< ); } - async #signTransaction( - originalTransactionMeta: TransactionMeta, - ): Promise { - let transactionMeta = originalTransactionMeta; + async #applyBeforeSignHook( + transactionMeta: TransactionMeta, + ): Promise { const { id: transactionId } = transactionMeta; log('Calling before sign hook', transactionMeta); @@ -3691,21 +3772,37 @@ export class TransactionController extends BaseController< log('Updated transaction after before sign hook'); } - transactionMeta = this.#getTransactionOrThrow(transactionId); + return this.#getTransactionOrThrow(transactionId); + } - const { networkClientId } = transactionMeta; + async #signTransaction( + originalTransactionMeta: TransactionMeta, + skipGasFeeTokenCheck = false, + skipBeforeSign = false, + ): Promise { + let transactionMeta = originalTransactionMeta; + const { id: transactionId } = transactionMeta; - await checkGasFeeTokenBeforePublish({ - messenger: this.messenger, - networkClientId, - fetchGasFeeTokens: async (tx) => - (await this.#getGasFeeTokens(tx)).gasFeeTokens, - transaction: transactionMeta, - updateTransaction: (txId, fn) => - this.#updateTransactionInternal({ transactionId: txId }, fn), - }); + if (!skipBeforeSign) { + transactionMeta = await this.#applyBeforeSignHook(transactionMeta); + } + + if (!skipGasFeeTokenCheck) { + const { networkClientId } = transactionMeta; + + await checkGasFeeTokenBeforePublish({ + messenger: this.messenger, + networkClientId, + fetchGasFeeTokens: async (tx) => + (await this.#getGasFeeTokens(tx)).gasFeeTokens, + transaction: transactionMeta, + updateTransaction: (txId, fn) => + this.#updateTransactionInternal({ transactionId: txId }, fn), + }); + + transactionMeta = this.#getTransactionOrThrow(transactionId); + } - transactionMeta = this.#getTransactionOrThrow(transactionId); const { chainId, isExternalSign, txParams } = transactionMeta; if (isExternalSign) { diff --git a/packages/transaction-controller/src/api/simulation-api.test.ts b/packages/transaction-controller/src/api/simulation-api.test.ts index 919af70645a..4336623cdac 100644 --- a/packages/transaction-controller/src/api/simulation-api.test.ts +++ b/packages/transaction-controller/src/api/simulation-api.test.ts @@ -194,5 +194,26 @@ describe('Simulation API Utils', () => { code: expect.any(String), }); }); + + it('creates overrides when they are not already present', async () => { + const request = cloneDeep(REQUEST_MOCK); + request.overrides = undefined; + request.transactions[0].to = + DELEGATION_MANAGER_ADDRESSES[0].toUpperCase() as Hex; + + await simulateTransactions(CHAIN_ID_MOCK, request); + + expect(fetchMock).toHaveBeenCalledTimes(2); + + const requestBody = JSON.parse( + fetchMock.mock.calls[1][1]?.body as string, + ); + + expect(requestBody.params[0].overrides).toStrictEqual({ + [DELEGATION_MANAGER_ADDRESSES[0]]: { + code: expect.any(String), + }, + }); + }); }); }); diff --git a/packages/transaction-controller/src/index.ts b/packages/transaction-controller/src/index.ts index b7b17038a97..4905020bd28 100644 --- a/packages/transaction-controller/src/index.ts +++ b/packages/transaction-controller/src/index.ts @@ -71,6 +71,8 @@ export type { BatchTransaction, BatchTransactionParams, BeforeSignHook, + IsSponsoredHook, + ShouldSignHook, DappSuggestedGasFees, DefaultGasEstimates, FeeMarketEIP1559Values, diff --git a/packages/transaction-controller/src/types.ts b/packages/transaction-controller/src/types.ts index 83b19561560..bd9b9e4cafa 100644 --- a/packages/transaction-controller/src/types.ts +++ b/packages/transaction-controller/src/types.ts @@ -2130,6 +2130,21 @@ export type AfterAddHook = (request: { updateTransaction?: (transaction: TransactionMeta) => void; }>; +/** + * Custom logic to determine whether a transaction should be treated as sponsored. + */ +export type IsSponsoredHook = (request: { + transactionMeta: TransactionMeta; +}) => Promise; + +/** + * Custom logic to determine whether a transaction should be signed locally. + */ +export type ShouldSignHook = (request: { + transactionMeta: TransactionMeta; + isSponsored: boolean; +}) => Promise; + /** * Custom logic to be executed before a transaction is signed. * Can optionally update the transaction by returning the `updateTransaction` callback. diff --git a/packages/transaction-controller/src/utils/gas-fee-tokens.test.ts b/packages/transaction-controller/src/utils/gas-fee-tokens.test.ts index ea9db48a3dc..99ee4fd3fe5 100644 --- a/packages/transaction-controller/src/utils/gas-fee-tokens.test.ts +++ b/packages/transaction-controller/src/utils/gas-fee-tokens.test.ts @@ -374,6 +374,44 @@ describe('Gas Fee Tokens Utils', () => { ); }); + it('returns empty gas fee tokens if the EIP-7702 public key is not provided', async () => { + const request = cloneDeep(REQUEST_MOCK); + request.publicKeyEIP7702 = undefined; + + doesChainSupportEIP7702Mock.mockReturnValueOnce(true); + simulateTransactionsMock.mockResolvedValueOnce({ + transactions: [], + sponsorship: { + isSponsored: false, + error: null, + }, + }); + + expect(await getGasFeeTokens(request)).toStrictEqual({ + gasFeeTokens: [], + isGasFeeSponsored: false, + }); + }); + + it('returns empty gas fee tokens if the upgrade contract address cannot be resolved', async () => { + const request = cloneDeep(REQUEST_MOCK); + + doesChainSupportEIP7702Mock.mockReturnValueOnce(true); + getEIP7702UpgradeContractAddressMock.mockReturnValueOnce(undefined); + simulateTransactionsMock.mockResolvedValueOnce({ + transactions: [], + sponsorship: { + isSponsored: false, + error: null, + }, + }); + + expect(await getGasFeeTokens(request)).toStrictEqual({ + gasFeeTokens: [], + isGasFeeSponsored: false, + }); + }); + it('forwards simulation config', async () => { const getSimulationConfigMock: GetSimulationConfig = jest.fn(); diff --git a/packages/transaction-controller/src/utils/prepare.ts b/packages/transaction-controller/src/utils/prepare.ts index 1586f750984..e437695d730 100644 --- a/packages/transaction-controller/src/utils/prepare.ts +++ b/packages/transaction-controller/src/utils/prepare.ts @@ -104,13 +104,7 @@ function normalizeAuthorizationList( * @returns The processed hexadecimal string. */ function removeLeadingZeroes(value: Hex | undefined): Hex | undefined { - if (!value) { - return value; - } - - if (value === '0x0') { - return '0x'; - } - - return (value.replace?.(/^0x(00)+/u, '0x') as Hex) ?? value; + return value === '0x0' + ? '0x' + : ((value?.replace?.(/^0x(00)+/u, '0x') as Hex | undefined) ?? value); }