diff --git a/package.json b/package.json index b2615bed4ce..090bb5d7d2f 100644 --- a/package.json +++ b/package.json @@ -1,6 +1,6 @@ { "name": "@metamask/core-monorepo", - "version": "675.0.0", + "version": "678.0.0", "private": true, "description": "Monorepo for packages shared between MetaMask clients", "repository": { diff --git a/packages/notification-services-controller/CHANGELOG.md b/packages/notification-services-controller/CHANGELOG.md index c1a2031f875..46c5a2677f2 100644 --- a/packages/notification-services-controller/CHANGELOG.md +++ b/packages/notification-services-controller/CHANGELOG.md @@ -7,6 +7,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [20.0.0] + ### Changed - **BREAKING:** Moved Notification API from v2 to v3 ([#7102](https://github.com/MetaMask/core/pull/7102)) @@ -618,7 +620,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Initial release -[Unreleased]: https://github.com/MetaMask/core/compare/@metamask/notification-services-controller@19.0.0...HEAD +[Unreleased]: https://github.com/MetaMask/core/compare/@metamask/notification-services-controller@20.0.0...HEAD +[20.0.0]: https://github.com/MetaMask/core/compare/@metamask/notification-services-controller@19.0.0...@metamask/notification-services-controller@20.0.0 [19.0.0]: https://github.com/MetaMask/core/compare/@metamask/notification-services-controller@18.3.1...@metamask/notification-services-controller@19.0.0 [18.3.1]: https://github.com/MetaMask/core/compare/@metamask/notification-services-controller@18.3.0...@metamask/notification-services-controller@18.3.1 [18.3.0]: https://github.com/MetaMask/core/compare/@metamask/notification-services-controller@18.2.0...@metamask/notification-services-controller@18.3.0 diff --git a/packages/notification-services-controller/package.json b/packages/notification-services-controller/package.json index 2999150edf4..fcf5b4bec37 100644 --- a/packages/notification-services-controller/package.json +++ b/packages/notification-services-controller/package.json @@ -1,6 +1,6 @@ { "name": "@metamask/notification-services-controller", - "version": "19.0.0", + "version": "20.0.0", "description": "Manages New MetaMask decentralized Notification system", "keywords": [ "MetaMask", diff --git a/packages/shield-controller/CHANGELOG.md b/packages/shield-controller/CHANGELOG.md index fcd14bc3abe..30fa896757b 100644 --- a/packages/shield-controller/CHANGELOG.md +++ b/packages/shield-controller/CHANGELOG.md @@ -7,6 +7,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [2.1.0] + +### Added + +- Added metrics in the Shield coverage response to track the latency ( [#7133](https://github.com/MetaMask/core/pull/7133)) + ## [2.0.0] ### Changed @@ -124,7 +130,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Initial release of the shield-controller package ([#6137](https://github.com/MetaMask/core/pull/6137) -[Unreleased]: https://github.com/MetaMask/core/compare/@metamask/shield-controller@2.0.0...HEAD +[Unreleased]: https://github.com/MetaMask/core/compare/@metamask/shield-controller@2.1.0...HEAD +[2.1.0]: https://github.com/MetaMask/core/compare/@metamask/shield-controller@2.0.0...@metamask/shield-controller@2.1.0 [2.0.0]: https://github.com/MetaMask/core/compare/@metamask/shield-controller@1.2.0...@metamask/shield-controller@2.0.0 [1.2.0]: https://github.com/MetaMask/core/compare/@metamask/shield-controller@1.1.0...@metamask/shield-controller@1.2.0 [1.1.0]: https://github.com/MetaMask/core/compare/@metamask/shield-controller@1.0.0...@metamask/shield-controller@1.1.0 diff --git a/packages/shield-controller/package.json b/packages/shield-controller/package.json index e220f2f009d..2bc2b31ec71 100644 --- a/packages/shield-controller/package.json +++ b/packages/shield-controller/package.json @@ -1,6 +1,6 @@ { "name": "@metamask/shield-controller", - "version": "2.0.0", + "version": "2.1.0", "description": "Controller handling shield transaction coverage logic", "keywords": [ "MetaMask", diff --git a/packages/shield-controller/src/backend.test.ts b/packages/shield-controller/src/backend.test.ts index 79adef1327a..b5777f6bbc7 100644 --- a/packages/shield-controller/src/backend.test.ts +++ b/packages/shield-controller/src/backend.test.ts @@ -51,6 +51,10 @@ function setup({ } describe('ShieldRemoteBackend', () => { + beforeEach(() => { + jest.clearAllMocks(); + }); + it('should check coverage', async () => { const { backend, fetchMock, getAccessToken } = setup(); @@ -70,14 +74,24 @@ describe('ShieldRemoteBackend', () => { const txMeta = generateMockTxMeta(); const coverageResult = await backend.checkCoverage({ txMeta }); - expect(coverageResult).toStrictEqual({ coverageId, ...result }); + expect({ + coverageId: coverageResult.coverageId, + message: result.message, + reasonCode: result.reasonCode, + status: result.status, + }).toStrictEqual({ + coverageId, + ...result, + }); + expect(typeof coverageResult.metrics.latency).toBe('number'); expect(fetchMock).toHaveBeenCalledTimes(2); expect(getAccessToken).toHaveBeenCalledTimes(2); }); it('should check coverage with delay', async () => { + const pollInterval = 100; const { backend, fetchMock, getAccessToken } = setup({ - getCoverageResultPollInterval: 100, + getCoverageResultPollInterval: pollInterval, }); // Mock init coverage check. @@ -101,13 +115,38 @@ describe('ShieldRemoteBackend', () => { } as unknown as Response); const txMeta = generateMockTxMeta(); + + // generateMockTxMeta also use Date.now() to set the time, only do this after generateMockTxMeta + // Mock Date.now() to control latency measurement + // Simulate latency that includes the retry delay (poll interval + processing time) + let callCount = 0; + const startTime = 1000; + const expectedLatency = pollInterval + 50; // poll interval + processing time + const nowSpy = jest.spyOn(Date, 'now').mockImplementation(() => { + callCount += 1; + // First call: start of #getCoverageResult + if (callCount === 1) { + return startTime; + } + // Final call: end of #getCoverageResult (after retry delay) + return startTime + expectedLatency; + }); + const coverageResult = await backend.checkCoverage({ txMeta }); - expect(coverageResult).toStrictEqual({ + + expect(coverageResult).toMatchObject({ coverageId, - ...result, + status: result.status, + message: result.message, + reasonCode: result.reasonCode, }); + expect(coverageResult.metrics.latency).toBe(expectedLatency); + // Latency should include the retry delay (at least the poll interval) + expect(coverageResult.metrics.latency).toBeGreaterThanOrEqual(pollInterval); expect(fetchMock).toHaveBeenCalledTimes(3); expect(getAccessToken).toHaveBeenCalledTimes(2); + + nowSpy.mockRestore(); }); it('should throw on init coverage check failure', async () => { @@ -187,6 +226,66 @@ describe('ShieldRemoteBackend', () => { await new Promise((resolve) => setTimeout(resolve, 10)); }); + it('returns latency in coverageResult', async () => { + const { backend, fetchMock } = setup(); + + fetchMock.mockResolvedValueOnce({ + status: 200, + json: jest.fn().mockResolvedValue({ coverageId: 'coverageId' }), + } as unknown as Response); + + const result = { status: 'covered', message: 'ok', reasonCode: 'E104' }; + fetchMock.mockResolvedValueOnce({ + status: 200, + json: jest.fn().mockResolvedValue(result), + } as unknown as Response); + + let nowValue = 1000; + const latencyMs = 123; + const nowSpy = jest.spyOn(Date, 'now').mockImplementation(() => { + const val = nowValue; + nowValue += latencyMs; + return val; + }); + + const txMeta = generateMockTxMeta(); + const coverageResult = await backend.checkCoverage({ txMeta }); + expect(coverageResult.metrics.latency).toBe(latencyMs); + + nowSpy.mockRestore(); + }); + + it('returns latency in signatureCoverageResult', async () => { + const { backend, fetchMock } = setup(); + + fetchMock.mockResolvedValueOnce({ + status: 200, + json: jest.fn().mockResolvedValue({ coverageId: 'coverageId' }), + } as unknown as Response); + + const result = { status: 'covered', message: 'ok', reasonCode: 'E104' }; + fetchMock.mockResolvedValueOnce({ + status: 200, + json: jest.fn().mockResolvedValue(result), + } as unknown as Response); + + let nowValue = 2000; + const latencyMs = 456; + const nowSpy = jest.spyOn(Date, 'now').mockImplementation(() => { + const val = nowValue; + nowValue += latencyMs; + return val; + }); + + const signatureRequest = generateMockSignatureRequest(); + const coverageResult = await backend.checkSignatureCoverage({ + signatureRequest, + }); + expect(coverageResult.metrics.latency).toBe(latencyMs); + + nowSpy.mockRestore(); + }); + describe('checkSignatureCoverage', () => { it('should check signature coverage', async () => { const { backend, fetchMock, getAccessToken } = setup(); @@ -209,10 +308,16 @@ describe('ShieldRemoteBackend', () => { const coverageResult = await backend.checkSignatureCoverage({ signatureRequest, }); - expect(coverageResult).toStrictEqual({ + expect({ + coverageId: coverageResult.coverageId, + message: result.message, + reasonCode: result.reasonCode, + status: result.status, + }).toStrictEqual({ coverageId, ...result, }); + expect(typeof coverageResult.metrics.latency).toBe('number'); expect(fetchMock).toHaveBeenCalledTimes(2); expect(getAccessToken).toHaveBeenCalledTimes(2); }); diff --git a/packages/shield-controller/src/backend.ts b/packages/shield-controller/src/backend.ts index a748bedcde0..0db6e677ef2 100644 --- a/packages/shield-controller/src/backend.ts +++ b/packages/shield-controller/src/backend.ts @@ -57,6 +57,9 @@ export type GetCoverageResultResponse = { message?: string; reasonCode?: string; status: CoverageStatus; + metrics: { + latency?: number; + }; }; export class ShieldRemoteBackend implements ShieldBackend { @@ -117,6 +120,7 @@ export class ShieldRemoteBackend implements ShieldBackend { message: coverageResult.message, reasonCode: coverageResult.reasonCode, status: coverageResult.status, + metrics: coverageResult.metrics, }; } @@ -143,6 +147,7 @@ export class ShieldRemoteBackend implements ShieldBackend { message: coverageResult.message, reasonCode: coverageResult.reasonCode, status: coverageResult.status, + metrics: coverageResult.metrics, }; } @@ -220,6 +225,9 @@ export class ShieldRemoteBackend implements ShieldBackend { const headers = await this.#createHeaders(); + // Start measuring total end-to-end latency including retries and delays + const startTime = Date.now(); + const getCoverageResultFn = async (signal: AbortSignal) => { const res = await this.#fetch(coverageResultUrl, { method: 'POST', @@ -227,8 +235,10 @@ export class ShieldRemoteBackend implements ShieldBackend { body: JSON.stringify(reqBody), signal, }); + if (res.status === 200) { - return (await res.json()) as GetCoverageResultResponse; + // Return the result without latency here - we'll add total latency after polling completes + return (await res.json()) as Omit; } // parse the error message from the response body @@ -242,7 +252,19 @@ export class ShieldRemoteBackend implements ShieldBackend { throw new HttpError(res.status, errorMessage); }; - return this.#pollingPolicy.start(requestId, getCoverageResultFn); + const result = await this.#pollingPolicy.start( + requestId, + getCoverageResultFn, + ); + + // Calculate total end-to-end latency including all retries and delays + const now = Date.now(); + const totalLatency = now - startTime; + + return { + ...result, + metrics: { latency: totalLatency }, + } as GetCoverageResultResponse; } async #createHeaders() { diff --git a/packages/shield-controller/src/types.ts b/packages/shield-controller/src/types.ts index ab00daabd85..1c9a911c691 100644 --- a/packages/shield-controller/src/types.ts +++ b/packages/shield-controller/src/types.ts @@ -6,6 +6,9 @@ export type CoverageResult = { message?: string; reasonCode?: string; status: CoverageStatus; + metrics: { + latency?: number; + }; }; export const coverageStatuses = ['covered', 'malicious', 'unknown'] as const; diff --git a/packages/subscription-controller/CHANGELOG.md b/packages/subscription-controller/CHANGELOG.md index 5e45964d221..f31c132d8ed 100644 --- a/packages/subscription-controller/CHANGELOG.md +++ b/packages/subscription-controller/CHANGELOG.md @@ -7,6 +7,13 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [4.2.2] + +### Changed + +- Trigger `triggerAccessTokenRefresh` everytime subscription state change instead of only when polling ([#7149](https://github.com/MetaMask/core/pull/7149)) +- Remove `triggerAccessTokenRefresh` after `startShieldSubscriptionWithCard` ([#7149](https://github.com/MetaMask/core/pull/7149)) + ## [4.2.1] ### Added @@ -188,7 +195,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Bump `@metamask/controller-utils` from `^11.12.0` to `^11.14.0` ([#6620](https://github.com/MetaMask/core/pull/6620), [#6629](https://github.com/MetaMask/core/pull/6629)) - Bump `@metamask/utils` from `^11.4.2` to `^11.8.0` ([#6588](https://github.com/MetaMask/core/pull/6588)) -[Unreleased]: https://github.com/MetaMask/core/compare/@metamask/subscription-controller@4.2.1...HEAD +[Unreleased]: https://github.com/MetaMask/core/compare/@metamask/subscription-controller@4.2.2...HEAD +[4.2.2]: https://github.com/MetaMask/core/compare/@metamask/subscription-controller@4.2.1...@metamask/subscription-controller@4.2.2 [4.2.1]: https://github.com/MetaMask/core/compare/@metamask/subscription-controller@4.2.0...@metamask/subscription-controller@4.2.1 [4.2.0]: https://github.com/MetaMask/core/compare/@metamask/subscription-controller@4.1.0...@metamask/subscription-controller@4.2.0 [4.1.0]: https://github.com/MetaMask/core/compare/@metamask/subscription-controller@4.0.0...@metamask/subscription-controller@4.1.0 diff --git a/packages/subscription-controller/package.json b/packages/subscription-controller/package.json index fbe3c4536f4..2d2d32c8f1a 100644 --- a/packages/subscription-controller/package.json +++ b/packages/subscription-controller/package.json @@ -1,6 +1,6 @@ { "name": "@metamask/subscription-controller", - "version": "4.2.1", + "version": "4.2.2", "description": "Handle user subscription", "keywords": [ "MetaMask", diff --git a/packages/subscription-controller/src/SubscriptionController.ts b/packages/subscription-controller/src/SubscriptionController.ts index a3ebf862394..6a654d39fed 100644 --- a/packages/subscription-controller/src/SubscriptionController.ts +++ b/packages/subscription-controller/src/SubscriptionController.ts @@ -241,8 +241,6 @@ export class SubscriptionController extends StaticIntervalPollingController()< > { readonly #subscriptionService: ISubscriptionService; - #shouldCallRefreshAuthToken: boolean = false; - /** * Creates a new SubscriptionController instance. * @@ -391,7 +389,8 @@ export class SubscriptionController extends StaticIntervalPollingController()< state.trialedProducts = newTrialedProducts; state.lastSubscription = newLastSubscription; }); - this.#shouldCallRefreshAuthToken = true; + // trigger access token refresh to ensure the user has the latest access token if subscription state change + this.triggerAccessTokenRefresh(); } return newSubscriptions; @@ -466,8 +465,7 @@ export class SubscriptionController extends StaticIntervalPollingController()< const response = await this.#subscriptionService.startSubscriptionWithCard(request); - - this.triggerAccessTokenRefresh(); + // note: no need to trigger access token refresh after startSubscriptionWithCard request because this only return stripe checkout session url, subscription not created yet return response; } @@ -476,7 +474,7 @@ export class SubscriptionController extends StaticIntervalPollingController()< this.#assertIsUserNotSubscribed({ products: request.products }); const response = await this.#subscriptionService.startSubscriptionWithCrypto(request); - this.triggerAccessTokenRefresh(); + return response; } @@ -730,10 +728,6 @@ export class SubscriptionController extends StaticIntervalPollingController()< async _executePoll(): Promise { await this.getSubscriptions(); - if (this.#shouldCallRefreshAuthToken) { - this.triggerAccessTokenRefresh(); - this.#shouldCallRefreshAuthToken = false; - } } /** diff --git a/packages/transaction-controller/CHANGELOG.md b/packages/transaction-controller/CHANGELOG.md index d1bccf9fb33..64f82485329 100644 --- a/packages/transaction-controller/CHANGELOG.md +++ b/packages/transaction-controller/CHANGELOG.md @@ -7,6 +7,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- Add optional `gasFeeToken` property to `addTransaction` and `addTransactionBatch` methods ([#7123](https://github.com/MetaMask/core/pull/7123)) + - Also add optional `gasFeeToken` and `isGasFeeTokenIgnoredIfBalance` properties to `TransactionMeta`. + ## [61.2.0] ### Added diff --git a/packages/transaction-controller/jest.config.js b/packages/transaction-controller/jest.config.js index 4a6ed4accc0..4b75427e9d2 100644 --- a/packages/transaction-controller/jest.config.js +++ b/packages/transaction-controller/jest.config.js @@ -18,7 +18,7 @@ module.exports = merge(baseConfig, { coverageThreshold: { global: { branches: 91.76, - functions: 93.24, + functions: 92.76, lines: 96.83, statements: 96.82, }, diff --git a/packages/transaction-controller/package.json b/packages/transaction-controller/package.json index 22e5c16a5a1..ab734742f42 100644 --- a/packages/transaction-controller/package.json +++ b/packages/transaction-controller/package.json @@ -64,6 +64,7 @@ "@metamask/rpc-errors": "^7.0.2", "@metamask/utils": "^11.8.1", "async-mutex": "^0.5.0", + "bignumber.js": "^9.1.2", "bn.js": "^5.2.1", "eth-method-registry": "^4.0.0", "fast-json-patch": "^3.1.1", diff --git a/packages/transaction-controller/src/TransactionController.test.ts b/packages/transaction-controller/src/TransactionController.test.ts index e9c865d50e7..3591a166ed3 100644 --- a/packages/transaction-controller/src/TransactionController.test.ts +++ b/packages/transaction-controller/src/TransactionController.test.ts @@ -1693,10 +1693,12 @@ describe('TransactionController', () => { isFirstTimeInteraction: undefined, isGasFeeIncluded: undefined, isGasFeeSponsored: undefined, + isGasFeeTokenIgnoredIfBalance: false, nestedTransactions: undefined, networkClientId: NETWORK_CLIENT_ID_MOCK, origin: undefined, securityAlertResponse: undefined, + selectedGasFeeToken: undefined, sendFlowHistory: expect.any(Array), status: TransactionStatus.unapproved as const, time: expect.any(Number), diff --git a/packages/transaction-controller/src/TransactionController.ts b/packages/transaction-controller/src/TransactionController.ts index b89a718bafc..3fa98332ac1 100644 --- a/packages/transaction-controller/src/TransactionController.ts +++ b/packages/transaction-controller/src/TransactionController.ts @@ -137,7 +137,10 @@ import { import { validateConfirmedExternalTransaction } from './utils/external-transactions'; import { updateFirstTimeInteraction } from './utils/first-time-interaction'; import { addGasBuffer, estimateGas, updateGas } from './utils/gas'; -import { getGasFeeTokens } from './utils/gas-fee-tokens'; +import { + checkGasFeeTokenBeforePublish, + getGasFeeTokens, +} from './utils/gas-fee-tokens'; import { updateGasFees } from './utils/gas-fees'; import { getGasFeeFlow } from './utils/gas-flow'; import { @@ -1215,6 +1218,7 @@ export class TransactionController extends BaseController< batchId, deviceConfirmedOn, disableGasBuffer, + gasFeeToken, isGasFeeIncluded, isGasFeeSponsored, method, @@ -1315,6 +1319,7 @@ export class TransactionController extends BaseController< deviceConfirmedOn, disableGasBuffer, id: random(), + isGasFeeTokenIgnoredIfBalance: Boolean(gasFeeToken), isGasFeeIncluded, isGasFeeSponsored, isFirstTimeInteraction: undefined, @@ -1322,6 +1327,7 @@ export class TransactionController extends BaseController< networkClientId, origin, securityAlertResponse, + selectedGasFeeToken: gasFeeToken, status: TransactionStatus.unapproved as const, time: Date.now(), txParams, @@ -3114,6 +3120,20 @@ export class TransactionController extends BaseController< clearApprovingTransactionId = () => this.#approvingTransactionIds.delete(transactionId); + const { networkClientId } = transactionMeta; + const ethQuery = this.#getEthQuery({ networkClientId }); + + await checkGasFeeTokenBeforePublish({ + ethQuery, + fetchGasFeeTokens: async (tx) => + (await this.#getGasFeeTokens(tx)).gasFeeTokens, + transaction: transactionMeta, + updateTransaction: (txId, fn) => + this.#updateTransactionInternal({ transactionId: txId }, fn), + }); + + transactionMeta = this.#getTransactionOrThrow(transactionId); + const [nonce, releaseNonce] = await getNextNonce( transactionMeta, (address: string) => @@ -3165,9 +3185,6 @@ export class TransactionController extends BaseController< return ApprovalState.NotApproved; } - const { networkClientId } = transactionMeta; - const ethQuery = this.#getEthQuery({ networkClientId }); - let preTxBalance: string | undefined; const shouldUpdatePreTxBalance = transactionMeta.type === TransactionType.swap; @@ -4255,14 +4272,8 @@ export class TransactionController extends BaseController< }; } - const gasFeeTokensResponse = await getGasFeeTokens({ - chainId, - getSimulationConfig: this.#getSimulationConfig, - isEIP7702GasFeeTokensEnabled: this.#isEIP7702GasFeeTokensEnabled, - messenger: this.messenger, - publicKeyEIP7702: this.#publicKeyEIP7702, - transactionMeta, - }); + const gasFeeTokensResponse = await this.#getGasFeeTokens(transactionMeta); + gasFeeTokens = gasFeeTokensResponse?.gasFeeTokens ?? []; isGasFeeSponsored = gasFeeTokensResponse?.isGasFeeSponsored ?? false; } @@ -4637,4 +4648,17 @@ export class TransactionController extends BaseController< return { transactionHash }; } + + async #getGasFeeTokens(transaction: TransactionMeta) { + const { chainId } = transaction; + + return await getGasFeeTokens({ + chainId, + getSimulationConfig: this.#getSimulationConfig, + isEIP7702GasFeeTokensEnabled: this.#isEIP7702GasFeeTokensEnabled, + messenger: this.messenger, + publicKeyEIP7702: this.#publicKeyEIP7702, + transactionMeta: transaction, + }); + } } diff --git a/packages/transaction-controller/src/types.ts b/packages/transaction-controller/src/types.ts index b23a567785f..b41c9462b16 100644 --- a/packages/transaction-controller/src/types.ts +++ b/packages/transaction-controller/src/types.ts @@ -268,6 +268,9 @@ export type TransactionMeta = { /** Whether MetaMask will be compensated for the gas fee by the transaction. */ isGasFeeIncluded?: boolean; + /** Whether the `selectedGasFeeToken` is only used if the user has insufficient native balance. */ + isGasFeeTokenIgnoredIfBalance?: boolean; + /** Whether the intent of the transaction was achieved via an alternate route or chain. */ isIntentComplete?: boolean; @@ -1723,6 +1726,9 @@ export type TransactionBatchRequest = { /** Address of the account to submit the transaction batch. */ from: Hex; + /** Address of an ERC-20 token to pay for the gas fee, if the user has insufficient native balance. */ + gasFeeToken?: Hex; + /** Whether MetaMask will be compensated for the gas fee by the transaction. */ isGasFeeIncluded?: boolean; @@ -2061,6 +2067,9 @@ export type AddTransactionOptions = { /** Whether to disable the gas estimation buffer. */ disableGasBuffer?: boolean; + /** Address of an ERC-20 token to pay for the gas fee, if the user has insufficient native balance. */ + gasFeeToken?: Hex; + /** Whether MetaMask will be compensated for the gas fee by the transaction. */ isGasFeeIncluded?: boolean; diff --git a/packages/transaction-controller/src/utils/balance-changes.test.ts b/packages/transaction-controller/src/utils/balance-changes.test.ts index 9e92e619209..ed507af2796 100644 --- a/packages/transaction-controller/src/utils/balance-changes.test.ts +++ b/packages/transaction-controller/src/utils/balance-changes.test.ts @@ -282,7 +282,7 @@ function mockParseLog({ } } -describe('Simulation Utils', () => { +describe('Balance Change Utils', () => { const simulateTransactionsMock = jest.mocked(simulateTransactions); const queryMock = jest.mocked(query); diff --git a/packages/transaction-controller/src/utils/balance-changes.ts b/packages/transaction-controller/src/utils/balance-changes.ts index c8c01ddc243..a49404535a9 100644 --- a/packages/transaction-controller/src/utils/balance-changes.ts +++ b/packages/transaction-controller/src/utils/balance-changes.ts @@ -1,11 +1,12 @@ import type { Fragment, LogDescription, Result } from '@ethersproject/abi'; import { Interface } from '@ethersproject/abi'; -import { hexToBN, query, toHex } from '@metamask/controller-utils'; +import { hexToBN, toHex } from '@metamask/controller-utils'; import type EthQuery from '@metamask/eth-query'; import { abiERC20, abiERC721, abiERC1155 } from '@metamask/metamask-eth-abis'; import { createModuleLogger, type Hex } from '@metamask/utils'; import BN from 'bn.js'; +import { getNativeBalance } from './balance'; import { simulateTransactions } from '../api/simulation-api'; import type { SimulationResponseLog, @@ -726,11 +727,8 @@ async function baseRequest({ log('Required balance', requiredBalanceHex); - const currentBalanceHex = (await query(ethQuery, 'getBalance', [ - from, - 'latest', - ])) as Hex; - + const { balanceRaw } = await getNativeBalance(from, ethQuery); + const currentBalanceHex = toHex(balanceRaw); const currentBalanceBN = hexToBN(currentBalanceHex); log('Current balance', currentBalanceHex); diff --git a/packages/transaction-controller/src/utils/balance.test.ts b/packages/transaction-controller/src/utils/balance.test.ts new file mode 100644 index 00000000000..01800c54d78 --- /dev/null +++ b/packages/transaction-controller/src/utils/balance.test.ts @@ -0,0 +1,68 @@ +import { query, toHex } from '@metamask/controller-utils'; +import type EthQuery from '@metamask/eth-query'; + +import { getNativeBalance, isNativeBalanceSufficientForGas } from './balance'; +import type { TransactionMeta } from '..'; + +jest.mock('@metamask/controller-utils', () => ({ + ...jest.requireActual('@metamask/controller-utils'), + query: jest.fn(), +})); + +const ETH_QUERY_MOCK = {} as EthQuery; +const BALANCE_MOCK = '21000000000000'; + +const TRANSACTION_META_MOCK = { + txParams: { + from: '0x1234', + gas: toHex(21000), + maxFeePerGas: toHex(1000000000), // 1 Gwei + }, +} as TransactionMeta; + +describe('Balance Utils', () => { + const queryMock = jest.mocked(query); + + beforeEach(() => { + jest.resetAllMocks(); + + queryMock.mockResolvedValue(toHex(BALANCE_MOCK)); + }); + + describe('getNativeBalance', () => { + it('returns native balance', async () => { + const result = await getNativeBalance('0x1234', ETH_QUERY_MOCK); + + expect(result).toStrictEqual({ + balanceRaw: BALANCE_MOCK, + balanceHuman: '0.000021', + }); + }); + }); + + describe('isNativeBalanceSufficientForGas', () => { + it('returns true if balance is sufficient for gas', async () => { + const result = await isNativeBalanceSufficientForGas( + TRANSACTION_META_MOCK, + ETH_QUERY_MOCK, + ); + + expect(result).toBe(true); + }); + + it('returns false if balance is insufficient for gas', async () => { + const result = await isNativeBalanceSufficientForGas( + { + ...TRANSACTION_META_MOCK, + txParams: { + ...TRANSACTION_META_MOCK.txParams, + gas: toHex(21001), + }, + }, + ETH_QUERY_MOCK, + ); + + expect(result).toBe(false); + }); + }); +}); diff --git a/packages/transaction-controller/src/utils/balance.ts b/packages/transaction-controller/src/utils/balance.ts new file mode 100644 index 00000000000..1a362d60c36 --- /dev/null +++ b/packages/transaction-controller/src/utils/balance.ts @@ -0,0 +1,52 @@ +import { query } from '@metamask/controller-utils'; +import type EthQuery from '@metamask/eth-query'; +import type { Hex } from '@metamask/utils'; +import { BigNumber } from 'bignumber.js'; + +import type { TransactionMeta } from '..'; + +/** + * Get the native balance for an address. + * + * @param address - Address to get the balance for. + * @param ethQuery - EthQuery instance to use. + * @returns Balance in both human-readable and raw format. + */ +export async function getNativeBalance(address: Hex, ethQuery: EthQuery) { + const balanceRawHex = (await query(ethQuery, 'getBalance', [ + address, + 'latest', + ])) as Hex; + + const balanceRaw = new BigNumber(balanceRawHex).toString(10); + const balanceHuman = new BigNumber(balanceRaw).shiftedBy(-18).toString(10); + + return { + balanceHuman, + balanceRaw, + }; +} + +/** + * Determine if the native balance is sufficient to cover max gas cost. + * + * @param transaction - Transaction metadata. + * @param ethQuery - EthQuery instance. + * @returns True if the native balance is sufficient, false otherwise. + */ +export async function isNativeBalanceSufficientForGas( + transaction: TransactionMeta, + ethQuery: EthQuery, +): Promise { + const from = transaction.txParams.from as Hex; + + const gasCostRawValue = new BigNumber( + transaction.txParams.gas ?? '0x0', + ).multipliedBy( + transaction.txParams.maxFeePerGas ?? transaction.txParams.gasPrice ?? '0x0', + ); + + const { balanceRaw } = await getNativeBalance(from, ethQuery); + + return gasCostRawValue.isLessThanOrEqualTo(balanceRaw); +} diff --git a/packages/transaction-controller/src/utils/batch.ts b/packages/transaction-controller/src/utils/batch.ts index 3e5ead32db8..b0b3385d3b7 100644 --- a/packages/transaction-controller/src/utils/batch.ts +++ b/packages/transaction-controller/src/utils/batch.ts @@ -289,6 +289,7 @@ async function addTransactionBatchWith7702( const { batchId: batchIdOverride, from, + gasFeeToken, networkClientId, origin, requireApproval, @@ -400,6 +401,7 @@ async function addTransactionBatchWith7702( const { result } = await addTransaction(txParams, { batchId, + gasFeeToken, isGasFeeIncluded: userRequest.isGasFeeIncluded, isGasFeeSponsored: userRequest.isGasFeeSponsored, nestedTransactions, 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 febdf38b2c4..50579834418 100644 --- a/packages/transaction-controller/src/utils/gas-fee-tokens.test.ts +++ b/packages/transaction-controller/src/utils/gas-fee-tokens.test.ts @@ -1,10 +1,17 @@ +import type EthQuery from '@metamask/eth-query'; +import type { Hex } from '@metamask/utils'; import { cloneDeep } from 'lodash'; +import { isNativeBalanceSufficientForGas } from './balance'; import { doesChainSupportEIP7702 } from './eip7702'; import { getEIP7702UpgradeContractAddress } from './feature-flags'; import type { GetGasFeeTokensRequest } from './gas-fee-tokens'; -import { getGasFeeTokens } from './gas-fee-tokens'; +import { + checkGasFeeTokenBeforePublish, + getGasFeeTokens, +} from './gas-fee-tokens'; import type { + GasFeeToken, GetSimulationConfig, TransactionControllerMessenger, TransactionMeta, @@ -14,32 +21,41 @@ import { simulateTransactions } from '../api/simulation-api'; jest.mock('../api/simulation-api'); jest.mock('./eip7702'); jest.mock('./feature-flags'); +jest.mock('./balance'); const CHAIN_ID_MOCK = '0x1'; -const TOKEN_ADDRESS_1_MOCK = '0x1234567890abcdef1234567890abcdef12345678'; +const TOKEN_ADDRESS_1_MOCK = + '0x1234567890abcdef1234567890abcdef12345678' as Hex; const TOKEN_ADDRESS_2_MOCK = '0xabcdef1234567890abcdef1234567890abcdef12'; const UPGRADE_CONTRACT_ADDRESS_MOCK = '0xabcdefabcdefabcdefabcdefabcdefabcdefabcdef'; +const TRANSACTION_META_MOCK = { + txParams: { + from: '0xabcdefabcdefabcdefabcdefabcdefabcdefabcdef', + to: '0x1234567890abcdef1234567890abcdef1234567a', + value: '0x1000000000000000000', + data: '0x', + }, +} as TransactionMeta; + const REQUEST_MOCK: GetGasFeeTokensRequest = { chainId: CHAIN_ID_MOCK, isEIP7702GasFeeTokensEnabled: jest.fn().mockResolvedValue(true), getSimulationConfig: jest.fn(), messenger: {} as TransactionControllerMessenger, publicKeyEIP7702: '0x123', - transactionMeta: { - txParams: { - from: '0xabcdefabcdefabcdefabcdefabcdefabcdefabcdef', - to: '0x1234567890abcdef1234567890abcdef1234567a', - value: '0x1000000000000000000', - data: '0x', - }, - } as TransactionMeta, + transactionMeta: TRANSACTION_META_MOCK, }; describe('Gas Fee Tokens Utils', () => { const simulateTransactionsMock = jest.mocked(simulateTransactions); const doesChainSupportEIP7702Mock = jest.mocked(doesChainSupportEIP7702); + + const isNativeBalanceSufficientForGasMock = jest.mocked( + isNativeBalanceSufficientForGas, + ); + const getEIP7702UpgradeContractAddressMock = jest.mocked( getEIP7702UpgradeContractAddress, ); @@ -50,6 +66,8 @@ describe('Gas Fee Tokens Utils', () => { getEIP7702UpgradeContractAddressMock.mockReturnValue( UPGRADE_CONTRACT_ADDRESS_MOCK, ); + + isNativeBalanceSufficientForGasMock.mockResolvedValue(false); }); describe('getGasFeeTokens', () => { @@ -376,4 +394,75 @@ describe('Gas Fee Tokens Utils', () => { ); }); }); + + describe('checkGasFeeTokenBeforePublish', () => { + let request: Parameters[0]; + + beforeEach(() => { + request = { + ethQuery: {} as EthQuery, + fetchGasFeeTokens: jest.fn(), + transaction: cloneDeep(TRANSACTION_META_MOCK), + updateTransaction: jest.fn(), + }; + }); + + it('throws if gas fee token not found in gas fee tokens', async () => { + request.transaction.isGasFeeTokenIgnoredIfBalance = true; + request.transaction.selectedGasFeeToken = TOKEN_ADDRESS_1_MOCK; + request.transaction.gasFeeTokens = []; + + await expect(checkGasFeeTokenBeforePublish(request)).rejects.toThrow( + 'Gas fee token not found and insufficient native balance', + ); + }); + + it('updates gas fee tokens if undefined', async () => { + request.transaction.isGasFeeTokenIgnoredIfBalance = true; + request.transaction.selectedGasFeeToken = TOKEN_ADDRESS_1_MOCK; + request.transaction.gasFeeTokens = undefined; + + jest.mocked(request.fetchGasFeeTokens).mockResolvedValueOnce([ + { + tokenAddress: TOKEN_ADDRESS_1_MOCK, + } as GasFeeToken, + ]); + + await checkGasFeeTokenBeforePublish(request); + + expect(request.fetchGasFeeTokens).toHaveBeenCalledTimes(1); + }); + + it('removes selected gas fee token if native balance sufficient', async () => { + request.transaction.isGasFeeTokenIgnoredIfBalance = true; + request.transaction.selectedGasFeeToken = TOKEN_ADDRESS_1_MOCK; + request.transaction.isExternalSign = true; + + isNativeBalanceSufficientForGasMock.mockResolvedValueOnce(true); + + await checkGasFeeTokenBeforePublish(request); + + jest + .mocked(request.updateTransaction) + .mock.calls[0][1](request.transaction); + + expect(request.transaction.selectedGasFeeToken).toBeUndefined(); + expect(request.transaction.isExternalSign).toBe(false); + }); + + it('does nothing if no selected gas fee token', async () => { + await checkGasFeeTokenBeforePublish(request); + + expect(request.updateTransaction).not.toHaveBeenCalled(); + }); + + it('does nothing if not ignoring gas fee token when native balance sufficient', async () => { + request.transaction.selectedGasFeeToken = TOKEN_ADDRESS_1_MOCK; + request.transaction.isGasFeeTokenIgnoredIfBalance = false; + + await checkGasFeeTokenBeforePublish(request); + + expect(request.updateTransaction).not.toHaveBeenCalled(); + }); + }); }); diff --git a/packages/transaction-controller/src/utils/gas-fee-tokens.ts b/packages/transaction-controller/src/utils/gas-fee-tokens.ts index b317bec336e..73d10e6c86a 100644 --- a/packages/transaction-controller/src/utils/gas-fee-tokens.ts +++ b/packages/transaction-controller/src/utils/gas-fee-tokens.ts @@ -1,7 +1,9 @@ +import type EthQuery from '@metamask/eth-query'; import { rpcErrors } from '@metamask/rpc-errors'; import type { Hex } from '@metamask/utils'; import { createModuleLogger } from '@metamask/utils'; +import { isNativeBalanceSufficientForGas } from './balance'; import { ERROR_MESSAGE_NO_UPGRADE_CONTRACT } from './batch'; import { ERROR_MESSGE_PUBLIC_KEY, doesChainSupportEIP7702 } from './eip7702'; import { getEIP7702UpgradeContractAddress } from './feature-flags'; @@ -115,6 +117,81 @@ export async function getGasFeeTokens({ } } +/** + * Check and update gas fee token selection before publishing a transaction. + * + * @param request - Request object. + * @param request.ethQuery - EthQuery instance. + * @param request.fetchGasFeeTokens - Function to fetch gas fee tokens. + * @param request.transaction - Transaction metadata. + * @param request.updateTransaction - Function to update the transaction. + */ +export async function checkGasFeeTokenBeforePublish({ + ethQuery, + fetchGasFeeTokens, + transaction, + updateTransaction, +}: { + ethQuery: EthQuery; + fetchGasFeeTokens: (transaction: TransactionMeta) => Promise; + transaction: TransactionMeta; + updateTransaction: ( + transactionId: string, + fn: (tx: TransactionMeta) => void, + ) => void; +}) { + const { gasFeeTokens, isGasFeeTokenIgnoredIfBalance, selectedGasFeeToken } = + transaction; + + if (!selectedGasFeeToken || !isGasFeeTokenIgnoredIfBalance) { + return; + } + + const hasNativeBalance = await isNativeBalanceSufficientForGas( + transaction, + ethQuery, + ); + + if (hasNativeBalance) { + log( + 'Ignoring gas fee token before publish due to sufficient native balance', + ); + + updateTransaction(transaction.id, (tx) => { + tx.isExternalSign = false; + tx.selectedGasFeeToken = undefined; + }); + + return; + } + + updateTransaction(transaction.id, (tx) => { + tx.isExternalSign = true; + }); + + let finalGasFeeTokens = gasFeeTokens; + + if (finalGasFeeTokens === undefined) { + const newGasFeeTokens = await fetchGasFeeTokens(transaction); + + updateTransaction(transaction.id, (tx) => { + tx.gasFeeTokens = newGasFeeTokens; + }); + + log('Updated gas fee tokens before publish', newGasFeeTokens); + + finalGasFeeTokens = newGasFeeTokens; + } + + if ( + !finalGasFeeTokens?.some( + (t) => t.tokenAddress.toLowerCase() === selectedGasFeeToken.toLowerCase(), + ) + ) { + throw new Error('Gas fee token not found and insufficient native balance'); + } +} + /** * Extract gas fee tokens from a simulation response. * diff --git a/packages/transaction-pay-controller/CHANGELOG.md b/packages/transaction-pay-controller/CHANGELOG.md index 46245d32f82..6eb3fb4642e 100644 --- a/packages/transaction-pay-controller/CHANGELOG.md +++ b/packages/transaction-pay-controller/CHANGELOG.md @@ -7,6 +7,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Fixed + +- **BREAKING:** Always retrieve quote if using Relay strategy and required token is Arbitrum USDC, even if payment token matches ([#7146](https://github.com/MetaMask/core/pull/7146)) + - Change `getStrategy` constructor option from asynchronous to synchronous. + ## [5.0.0] ### Added diff --git a/packages/transaction-pay-controller/src/TransactionPayController.test.ts b/packages/transaction-pay-controller/src/TransactionPayController.test.ts index 59095e59bc4..473fc919c12 100644 --- a/packages/transaction-pay-controller/src/TransactionPayController.test.ts +++ b/packages/transaction-pay-controller/src/TransactionPayController.test.ts @@ -77,7 +77,7 @@ describe('TransactionPayController', () => { createController(); expect( - await messenger.call( + messenger.call( 'TransactionPayController:getStrategy', TRANSACTION_META_MOCK, ), @@ -87,12 +87,12 @@ describe('TransactionPayController', () => { it('returns callback value if provided', async () => { new TransactionPayController({ getDelegationTransaction: jest.fn(), - getStrategy: async () => TransactionPayStrategy.Test, + getStrategy: () => TransactionPayStrategy.Test, messenger, }); expect( - await messenger.call( + messenger.call( 'TransactionPayController:getStrategy', TRANSACTION_META_MOCK, ), diff --git a/packages/transaction-pay-controller/src/TransactionPayController.ts b/packages/transaction-pay-controller/src/TransactionPayController.ts index eac0d916eb2..be80ea3f7d5 100644 --- a/packages/transaction-pay-controller/src/TransactionPayController.ts +++ b/packages/transaction-pay-controller/src/TransactionPayController.ts @@ -41,7 +41,7 @@ export class TransactionPayController extends BaseController< readonly #getStrategy?: ( transaction: TransactionMeta, - ) => Promise; + ) => TransactionPayStrategy; constructor({ getDelegationTransaction, @@ -139,7 +139,7 @@ export class TransactionPayController extends BaseController< this.messenger.registerActionHandler( 'TransactionPayController:getStrategy', - this.#getStrategy ?? (async () => TransactionPayStrategy.Relay), + this.#getStrategy ?? (() => TransactionPayStrategy.Relay), ); this.messenger.registerActionHandler( diff --git a/packages/transaction-pay-controller/src/helpers/TransactionPayPublishHook.test.ts b/packages/transaction-pay-controller/src/helpers/TransactionPayPublishHook.test.ts index b87ab8f2a22..681847534cc 100644 --- a/packages/transaction-pay-controller/src/helpers/TransactionPayPublishHook.test.ts +++ b/packages/transaction-pay-controller/src/helpers/TransactionPayPublishHook.test.ts @@ -61,7 +61,7 @@ describe('TransactionPayPublishHook', () => { }, } as TransactionPayControllerState); - getStrategyMock.mockResolvedValue(TransactionPayStrategy.Test); + getStrategyMock.mockReturnValue(TransactionPayStrategy.Test); }); it('executes strategy with quotes', async () => { diff --git a/packages/transaction-pay-controller/src/helpers/TransactionPayPublishHook.ts b/packages/transaction-pay-controller/src/helpers/TransactionPayPublishHook.ts index 7bbf7c7e8ac..c49cbcf252f 100644 --- a/packages/transaction-pay-controller/src/helpers/TransactionPayPublishHook.ts +++ b/packages/transaction-pay-controller/src/helpers/TransactionPayPublishHook.ts @@ -68,7 +68,7 @@ export class TransactionPayPublishHook { return EMPTY_RESULT; } - const strategy = await getStrategy(this.#messenger, transactionMeta); + const strategy = getStrategy(this.#messenger, transactionMeta); return await strategy.execute({ isSmartTransaction: this.#isSmartTransaction, diff --git a/packages/transaction-pay-controller/src/types.ts b/packages/transaction-pay-controller/src/types.ts index bb7571dac92..530abe7d4a3 100644 --- a/packages/transaction-pay-controller/src/types.ts +++ b/packages/transaction-pay-controller/src/types.ts @@ -68,7 +68,7 @@ export type TransactionPayControllerGetDelegationTransactionAction = { /** Action to get the pay strategy type used for a transaction. */ export type TransactionPayControllerGetStrategyAction = { type: `${typeof CONTROLLER_NAME}:getStrategy`; - handler: (transaction: TransactionMeta) => Promise; + handler: (transaction: TransactionMeta) => TransactionPayStrategy; }; /** Action to update the payment token for a transaction. */ @@ -104,9 +104,7 @@ export type TransactionPayControllerOptions = { getDelegationTransaction: GetDelegationTransactionCallback; /** Callback to select the PayStrategy for a transaction. */ - getStrategy?: ( - transaction: TransactionMeta, - ) => Promise; + getStrategy?: (transaction: TransactionMeta) => TransactionPayStrategy; /** Controller messenger. */ messenger: TransactionPayControllerMessenger; diff --git a/packages/transaction-pay-controller/src/utils/quotes.test.ts b/packages/transaction-pay-controller/src/utils/quotes.test.ts index 576e282c5c1..af60914f3d5 100644 --- a/packages/transaction-pay-controller/src/utils/quotes.test.ts +++ b/packages/transaction-pay-controller/src/utils/quotes.test.ts @@ -114,7 +114,7 @@ describe('Quotes Utils', () => { jest.resetAllMocks(); jest.clearAllTimers(); - getStrategyMock.mockResolvedValue({ + getStrategyMock.mockReturnValue({ execute: jest.fn(), getQuotes: getQuotesMock, getBatchTransactions: getBatchTransactionsMock, diff --git a/packages/transaction-pay-controller/src/utils/quotes.ts b/packages/transaction-pay-controller/src/utils/quotes.ts index 95c498414b2..d5d94581200 100644 --- a/packages/transaction-pay-controller/src/utils/quotes.ts +++ b/packages/transaction-pay-controller/src/utils/quotes.ts @@ -272,7 +272,7 @@ async function getQuotes( messenger: TransactionPayControllerMessenger, ) { const { id: transactionId } = transaction; - const strategy = await getStrategy(messenger as never, transaction); + const strategy = getStrategy(messenger as never, transaction); let quotes: TransactionPayQuote[] | undefined = []; try { diff --git a/packages/transaction-pay-controller/src/utils/source-amounts.test.ts b/packages/transaction-pay-controller/src/utils/source-amounts.test.ts index 5784a49170b..7a8bfc90921 100644 --- a/packages/transaction-pay-controller/src/utils/source-amounts.test.ts +++ b/packages/transaction-pay-controller/src/utils/source-amounts.test.ts @@ -1,9 +1,16 @@ import { updateSourceAmounts } from './source-amounts'; import { getTokenFiatRate } from './token'; -import type { TransactionPaymentToken } from '..'; +import { getTransaction } from './transaction'; +import { TransactionPayStrategy, type TransactionPaymentToken } from '..'; +import { + ARBITRUM_USDC_ADDRESS, + CHAIN_ID_ARBITRUM, +} from '../strategy/relay/constants'; +import { getMessengerMock } from '../tests/messenger-mock'; import type { TransactionData, TransactionPayRequiredToken } from '../types'; jest.mock('./token'); +jest.mock('./transaction'); const PAYMENT_TOKEN_MOCK: TransactionPaymentToken = { address: '0x123', @@ -37,11 +44,15 @@ const TRANSACTION_ID_MOCK = '123-456'; describe('Source Amounts Utils', () => { const getTokenFiatRateMock = jest.mocked(getTokenFiatRate); + const getTransactionMock = jest.mocked(getTransaction); + const { messenger, getStrategyMock } = getMessengerMock(); beforeEach(() => { jest.resetAllMocks(); getTokenFiatRateMock.mockReturnValue({ fiatRate: '2.0', usdRate: '3.0' }); + getStrategyMock.mockReturnValue(TransactionPayStrategy.Test); + getTransactionMock.mockReturnValue({ id: TRANSACTION_ID_MOCK } as never); }); describe('updateSourceAmounts', () => { @@ -52,7 +63,7 @@ describe('Source Amounts Utils', () => { tokens: [TRANSACTION_TOKEN_MOCK], }; - updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, {} as never); + updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, messenger); expect(transactionData.sourceAmounts).toStrictEqual([ { @@ -76,11 +87,35 @@ describe('Source Amounts Utils', () => { ], }; - updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, {} as never); + updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, messenger); expect(transactionData.sourceAmounts).toStrictEqual([]); }); + it('does not return empty array if payment token matches but hyperliquid deposit and relay strategy', () => { + getStrategyMock.mockReturnValue(TransactionPayStrategy.Relay); + + const transactionData: TransactionData = { + isLoading: false, + paymentToken: { + ...PAYMENT_TOKEN_MOCK, + address: ARBITRUM_USDC_ADDRESS, + chainId: CHAIN_ID_ARBITRUM, + }, + tokens: [ + { + ...TRANSACTION_TOKEN_MOCK, + address: ARBITRUM_USDC_ADDRESS, + chainId: CHAIN_ID_ARBITRUM, + }, + ], + }; + + updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, messenger); + + expect(transactionData.sourceAmounts).toHaveLength(1); + }); + it('returns empty array if skipIfBalance and has balance', () => { const transactionData: TransactionData = { isLoading: false, @@ -94,7 +129,7 @@ describe('Source Amounts Utils', () => { ], }; - updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, {} as never); + updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, messenger); expect(transactionData.sourceAmounts).toStrictEqual([]); }); @@ -108,7 +143,7 @@ describe('Source Amounts Utils', () => { getTokenFiatRateMock.mockReturnValue(undefined); - updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, {} as never); + updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, messenger); expect(transactionData.sourceAmounts).toStrictEqual([]); }); @@ -125,7 +160,7 @@ describe('Source Amounts Utils', () => { ], }; - updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, {} as never); + updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, messenger); expect(transactionData.sourceAmounts).toStrictEqual([]); }); @@ -136,7 +171,7 @@ describe('Source Amounts Utils', () => { tokens: [TRANSACTION_TOKEN_MOCK], }; - updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, {} as never); + updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, messenger); expect(transactionData.sourceAmounts).toBeUndefined(); }); @@ -148,14 +183,14 @@ describe('Source Amounts Utils', () => { tokens: [], }; - updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, {} as never); + updateSourceAmounts(TRANSACTION_ID_MOCK, transactionData, messenger); expect(transactionData.sourceAmounts).toBeUndefined(); }); // eslint-disable-next-line jest/expect-expect it('does nothing if no transaction data', () => { - updateSourceAmounts(TRANSACTION_ID_MOCK, undefined, {} as never); + updateSourceAmounts(TRANSACTION_ID_MOCK, undefined, messenger); }); }); }); diff --git a/packages/transaction-pay-controller/src/utils/source-amounts.ts b/packages/transaction-pay-controller/src/utils/source-amounts.ts index fe5203b852c..9907dba2dab 100644 --- a/packages/transaction-pay-controller/src/utils/source-amounts.ts +++ b/packages/transaction-pay-controller/src/utils/source-amounts.ts @@ -2,11 +2,18 @@ import { createModuleLogger } from '@metamask/utils'; import { BigNumber } from 'bignumber.js'; import { getTokenFiatRate } from './token'; +import { getTransaction } from './transaction'; import type { TransactionPayControllerMessenger, TransactionPaymentToken, } from '..'; +import { TransactionPayStrategy } from '..'; +import type { TransactionMeta } from '../../../transaction-controller/src'; import { projectLogger } from '../logger'; +import { + ARBITRUM_USDC_ADDRESS, + CHAIN_ID_ARBITRUM, +} from '../strategy/relay/constants'; import type { TransactionPaySourceAmount, TransactionData, @@ -38,7 +45,9 @@ export function updateSourceAmounts( } const sourceAmounts = tokens - .map((t) => calculateSourceAmount(paymentToken, t, messenger)) + .map((t) => + calculateSourceAmount(paymentToken, t, messenger, transactionId), + ) .filter(Boolean) as TransactionPaySourceAmount[]; log('Updated source amounts', { transactionId, sourceAmounts }); @@ -52,12 +61,14 @@ export function updateSourceAmounts( * @param paymentToken - Selected payment token. * @param token - Target token to cover. * @param messenger - Controller messenger. + * @param transactionId - ID of the transaction. * @returns The source amount or undefined if calculation failed. */ function calculateSourceAmount( paymentToken: TransactionPaymentToken, token: TransactionPayRequiredToken, messenger: TransactionPayControllerMessenger, + transactionId: string, ): TransactionPaySourceAmount | undefined { const paymentTokenFiatRate = getTokenFiatRate( messenger, @@ -71,10 +82,6 @@ function calculateSourceAmount( const hasBalance = new BigNumber(token.balanceRaw).gte(token.amountRaw); - const isSameTokenSelected = - token.address.toLowerCase() === paymentToken.address.toLowerCase() && - token.chainId === paymentToken.chainId; - if (token.skipIfBalance && hasBalance) { log('Skipping token as sufficient balance', { tokenAddress: token.address, @@ -82,7 +89,15 @@ function calculateSourceAmount( return undefined; } - if (isSameTokenSelected) { + const strategy = getStrategyType(transactionId, messenger); + + const isSameTokenSelected = + token.address.toLowerCase() === paymentToken.address.toLowerCase() && + token.chainId === paymentToken.chainId; + + const isAlwaysRequired = isQuoteAlwaysRequired(token, strategy); + + if (isSameTokenSelected && !isAlwaysRequired) { log('Skipping token as same as payment token'); return undefined; } @@ -108,3 +123,40 @@ function calculateSourceAmount( targetTokenAddress: token.address, }; } + +/** + * Determine if a quote is always required for a token and strategy. + * + * @param token - Target token. + * @param strategy - Payment strategy. + * @returns True if a quote is always required, false otherwise. + */ +function isQuoteAlwaysRequired( + token: TransactionPayRequiredToken, + strategy: TransactionPayStrategy, +) { + const isHyperliquidDeposit = + token.chainId === CHAIN_ID_ARBITRUM && + token.address.toLowerCase() === ARBITRUM_USDC_ADDRESS.toLowerCase(); + + return strategy === TransactionPayStrategy.Relay && isHyperliquidDeposit; +} + +/** + * Get the strategy type for a transaction. + * + * @param transactionId - ID of the transaction. + * @param messenger - Controller messenger. + * @returns Payment strategy type. + */ +function getStrategyType( + transactionId: string, + messenger: TransactionPayControllerMessenger, +) { + const transaction = getTransaction( + transactionId, + messenger, + ) as TransactionMeta; + + return messenger.call('TransactionPayController:getStrategy', transaction); +} diff --git a/packages/transaction-pay-controller/src/utils/strategy.test.ts b/packages/transaction-pay-controller/src/utils/strategy.test.ts index 5af1a6694ee..7b2c65593e2 100644 --- a/packages/transaction-pay-controller/src/utils/strategy.test.ts +++ b/packages/transaction-pay-controller/src/utils/strategy.test.ts @@ -18,35 +18,35 @@ describe('Strategy Utils', () => { describe('getStrategy', () => { it('returns TestStrategy if strategy name is Test', async () => { - getStrategyMock.mockResolvedValue(TransactionPayStrategy.Test); + getStrategyMock.mockReturnValue(TransactionPayStrategy.Test); - const strategy = await getStrategy(messenger, TRANSACTION_META_MOCK); + const strategy = getStrategy(messenger, TRANSACTION_META_MOCK); expect(strategy).toBeInstanceOf(TestStrategy); }); it('returns BridgeStrategy if strategy name is Bridge', async () => { - getStrategyMock.mockResolvedValue(TransactionPayStrategy.Bridge); + getStrategyMock.mockReturnValue(TransactionPayStrategy.Bridge); - const strategy = await getStrategy(messenger, TRANSACTION_META_MOCK); + const strategy = getStrategy(messenger, TRANSACTION_META_MOCK); expect(strategy).toBeInstanceOf(BridgeStrategy); }); it('returns RelayStrategy if strategy name is Relay', async () => { - getStrategyMock.mockResolvedValue(TransactionPayStrategy.Relay); + getStrategyMock.mockReturnValue(TransactionPayStrategy.Relay); - const strategy = await getStrategy(messenger, TRANSACTION_META_MOCK); + const strategy = getStrategy(messenger, TRANSACTION_META_MOCK); expect(strategy).toBeInstanceOf(RelayStrategy); }); it('throws if strategy name is unknown', async () => { - getStrategyMock.mockResolvedValue('UnknownStrategy' as never); + getStrategyMock.mockReturnValue('UnknownStrategy' as never); - await expect( - getStrategy(messenger, TRANSACTION_META_MOCK), - ).rejects.toThrow('Unknown strategy: UnknownStrategy'); + expect(() => getStrategy(messenger, TRANSACTION_META_MOCK)).toThrow( + 'Unknown strategy: UnknownStrategy', + ); }); }); diff --git a/packages/transaction-pay-controller/src/utils/strategy.ts b/packages/transaction-pay-controller/src/utils/strategy.ts index 450e2b2cbf5..dc505431f4e 100644 --- a/packages/transaction-pay-controller/src/utils/strategy.ts +++ b/packages/transaction-pay-controller/src/utils/strategy.ts @@ -13,11 +13,11 @@ import type { PayStrategy, TransactionPayControllerMessenger } from '../types'; * @param transaction - Transaction to get the strategy for. * @returns The payment strategy instance. */ -export async function getStrategy( +export function getStrategy( messenger: TransactionPayControllerMessenger, transaction: TransactionMeta, -): Promise> { - const strategyName = await messenger.call( +): PayStrategy { + const strategyName = messenger.call( 'TransactionPayController:getStrategy', transaction, ); diff --git a/yarn.lock b/yarn.lock index 9f7926f4abd..257d3964e0a 100644 --- a/yarn.lock +++ b/yarn.lock @@ -5086,6 +5086,7 @@ __metadata: "@types/jest": "npm:^27.4.1" "@types/node": "npm:^16.18.54" async-mutex: "npm:^0.5.0" + bignumber.js: "npm:^9.1.2" bn.js: "npm:^5.2.1" deepmerge: "npm:^4.2.2" eth-method-registry: "npm:^4.0.0"