diff --git a/app/__tests__/components/trade/DepositWithdrawCard.wallet-switch.test.tsx b/app/__tests__/components/trade/DepositWithdrawCard.wallet-switch.test.tsx new file mode 100644 index 00000000..c91dad62 --- /dev/null +++ b/app/__tests__/components/trade/DepositWithdrawCard.wallet-switch.test.tsx @@ -0,0 +1,344 @@ +import { act, fireEvent, render, screen, waitFor } from '@testing-library/react'; +import { PublicKey } from '@solana/web3.js'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; + +import { DepositWithdrawCard } from '@/components/trade/DepositWithdrawCard'; + +const mocks = vi.hoisted(() => ({ + useWalletCompat: vi.fn(), + useConnectionCompat: vi.fn(), + getTokenAccountBalance: vi.fn(), + getAssociatedTokenAddressSync: vi.fn(), + useUserAccount: vi.fn(), + useSlabState: vi.fn(), + useTokenMeta: vi.fn(), + initUser: vi.fn(), + deposit: vi.fn(), +})); + +vi.mock('@/hooks/useWalletCompat', () => ({ + useWalletCompat: mocks.useWalletCompat, + useConnectionCompat: mocks.useConnectionCompat, +})); + +vi.mock('@solana/spl-token', () => ({ + getAssociatedTokenAddressSync: mocks.getAssociatedTokenAddressSync, +})); + +vi.mock('@/hooks/useUserAccount', () => ({ + useUserAccount: mocks.useUserAccount, +})); + +vi.mock('@/hooks/useDeposit', () => ({ + useDeposit: () => ({ + deposit: mocks.deposit, + loading: false, + error: null, + }), +})); + +vi.mock('@/hooks/useWithdraw', () => ({ + useWithdraw: () => ({ + withdraw: vi.fn(), + loading: false, + error: null, + }), +})); + +vi.mock('@/hooks/useInitUser', () => ({ + useInitUser: () => ({ + initUser: mocks.initUser, + loading: false, + error: null, + }), +})); + +vi.mock('@/components/providers/SlabProvider', () => ({ + useSlabState: mocks.useSlabState, +})); + +vi.mock('@/hooks/useTokenMeta', () => ({ + useTokenMeta: mocks.useTokenMeta, +})); + +vi.mock('@/hooks/useLivePrice', () => ({ + useLivePrice: () => ({ + priceE6: null, + }), +})); + +vi.mock('@/lib/mock-mode', () => ({ + isMockMode: () => false, +})); + +vi.mock('@/lib/mock-trade-data', () => ({ + isMockSlab: () => false, + getMockUserAccount: () => null, +})); + +vi.mock('@/lib/tx', () => ({ + prewarmTxLanding: vi.fn(), +})); + +vi.mock('@/components/trade/DevnetTokenFaucetButton', () => ({ + DevnetTokenFaucetButton: () => null, +})); + +describe('DepositWithdrawCard wallet-scoped balance lifecycle', () => { + const walletA = new PublicKey('11111111111111111111111111111111'); + const walletB = new PublicKey('So11111111111111111111111111111111111111112'); + const collateralMint = new PublicKey('SysvarRent111111111111111111111111111111111'); + + const connection = { + getTokenAccountBalance: mocks.getTokenAccountBalance, + }; + + let activeWallet = walletA; + + beforeEach(() => { + vi.clearAllMocks(); + + activeWallet = walletA; + + mocks.useWalletCompat.mockImplementation(() => ({ + connected: true, + publicKey: activeWallet, + })); + + mocks.useConnectionCompat.mockReturnValue({ + connection, + }); + + mocks.getAssociatedTokenAddressSync.mockReturnValue(collateralMint); + + mocks.useUserAccount.mockReturnValue(null); + + mocks.useSlabState.mockReturnValue({ + config: { + collateralMint, + }, + params: null, + }); + + mocks.useTokenMeta.mockReturnValue({ + symbol: 'USDC', + decimals: 6, + }); + + mocks.deposit.mockResolvedValue('deposit-signature'); + + mocks.initUser.mockResolvedValue({ + sig: 'test-signature', + }); + }); + + it('withholds wallet A balance and prevents wallet A-derived account creation while wallet B balance is pending', async () => { + mocks.getTokenAccountBalance + .mockResolvedValueOnce({ + value: { + amount: '900000000', + decimals: 6, + }, + }) + .mockImplementationOnce( + () => + new Promise<{ + value: { + amount: string; + decimals: number; + }; + }>(() => {}), + ); + + const { rerender } = render(); + + await waitFor(() => { + expect(screen.getByText(/^Wallet:/)).toBeInTheDocument(); + }); + + expect( + screen.getByRole('button', { + name: 'Create Trading Account', + }), + ).toBeInTheDocument(); + + expect(mocks.getTokenAccountBalance).toHaveBeenCalledTimes(1); + + activeWallet = walletB; + + await act(async () => { + rerender(); + + await Promise.resolve(); + await Promise.resolve(); + }); + + await waitFor(() => { + expect(mocks.getTokenAccountBalance).toHaveBeenCalledTimes(2); + }); + + // Wallet A state must not remain visible or actionable after wallet B + // becomes active, even while wallet B's balance request is unresolved. + expect.soft(screen.queryByText(/^Wallet:/)).not.toBeInTheDocument(); + + expect + .soft( + screen.queryByRole('button', { + name: 'Create Trading Account', + }), + ) + .not.toBeInTheDocument(); + + // Confirm that the second balance request belongs to wallet B. + expect(mocks.getAssociatedTokenAddressSync).toHaveBeenNthCalledWith(2, collateralMint, walletB); + + const staleCreateAccountButton = screen.queryByRole('button', { + name: 'Create Trading Account', + }); + + // On a fixed implementation the stale button is absent and this block + // does not execute. On the vulnerable baseline it remains actionable. + if (staleCreateAccountButton) { + await act(async () => { + fireEvent.click(staleCreateAccountButton); + await Promise.resolve(); + }); + } + + // 500 USDC is the starter-deposit cap. It must never be derived from + // wallet A after wallet B becomes active but remains unverified. + expect.soft(mocks.initUser).not.toHaveBeenCalledWith(500_000_000n); + }); + + it('updates to wallet B balance after the replacement request resolves', async () => { + mocks.getTokenAccountBalance + .mockResolvedValueOnce({ + value: { + amount: '900000000', + decimals: 6, + }, + }) + .mockResolvedValueOnce({ + value: { + amount: '25000000', + decimals: 6, + }, + }); + + const { rerender } = render(); + + await waitFor(() => { + expect(screen.getByText(/^Wallet:/)).toHaveTextContent('900'); + }); + + expect(mocks.getTokenAccountBalance).toHaveBeenCalledTimes(1); + + activeWallet = walletB; + + await act(async () => { + rerender(); + + await Promise.resolve(); + }); + + await waitFor(() => { + expect(mocks.getTokenAccountBalance).toHaveBeenCalledTimes(2); + }); + + await waitFor(() => { + expect(screen.getByText(/^Wallet:/)).toHaveTextContent('25'); + }); + + expect(mocks.getAssociatedTokenAddressSync).toHaveBeenNthCalledWith(2, collateralMint, walletB); + }); + it('clears a wallet A Max amount and blocks deposit while wallet B balance is unverified', async () => { + let resolveWalletBBalance: + | ((value: { value: { amount: string; decimals: number } }) => void) + | undefined; + + mocks.useUserAccount.mockImplementation(() => ({ + idx: 7, + pubkey: activeWallet, + account: { + capital: 100_000_000n, + positionSize: 0n, + entryPrice: 0n, + pnl: 0n, + }, + })); + + mocks.getTokenAccountBalance + .mockResolvedValueOnce({ + value: { + amount: '900000000', + decimals: 6, + }, + }) + .mockImplementationOnce( + () => + new Promise<{ value: { amount: string; decimals: number } }>((resolve) => { + resolveWalletBBalance = resolve; + }), + ); + + const { rerender } = render(); + + const maxButton = await screen.findByRole('button', { name: 'Max' }); + fireEvent.click(maxButton); + + const amountInput = screen.getByPlaceholderText('Amount (USDC)'); + expect(amountInput).toHaveValue('900'); + + activeWallet = walletB; + + await act(async () => { + rerender(); + await Promise.resolve(); + await Promise.resolve(); + }); + + await waitFor(() => { + expect(mocks.getTokenAccountBalance).toHaveBeenCalledTimes(2); + }); + + expect(amountInput).toHaveValue(''); + + const pendingDepositButton = screen.getByRole('button', { + name: 'Deposit USDC', + }); + + expect(pendingDepositButton).toBeDisabled(); + fireEvent.click(pendingDepositButton); + expect(mocks.deposit).not.toHaveBeenCalled(); + + await act(async () => { + resolveWalletBBalance?.({ + value: { + amount: '25000000', + decimals: 6, + }, + }); + await Promise.resolve(); + }); + + const walletBMaxButton = await screen.findByRole('button', { name: 'Max' }); + fireEvent.click(walletBMaxButton); + + expect(amountInput).toHaveValue('25'); + + fireEvent.click( + screen.getByRole('button', { + name: 'Deposit USDC', + }), + ); + + await waitFor(() => { + expect(mocks.deposit).toHaveBeenCalledWith({ + userIdx: 7, + amount: 25_000_000n, + accountExists: true, + portfolioPk: walletB, + }); + }); + }); +}); diff --git a/app/components/trade/DepositWithdrawCard.tsx b/app/components/trade/DepositWithdrawCard.tsx index c1ff3c4e..1e3c4281 100644 --- a/app/components/trade/DepositWithdrawCard.tsx +++ b/app/components/trade/DepositWithdrawCard.tsx @@ -66,9 +66,43 @@ export const DepositWithdrawCard: FC = ({ slabAddress, if (!mockMode && walletConnected) prewarmTxLanding(connection); }, [connection, mockMode, walletConnected]); const [lastSig, setLastSig] = useState(null); - const [walletBalance, setWalletBalance] = useState(mockMode ? 500_000_000n : null); + + type WalletBalanceSnapshot = { + scopeKey: string; + amount: bigint | null; + decimals: number | null; + }; + + const walletBalanceScopeKey = + publicKey && mktConfig?.collateralMint + ? `${publicKey.toBase58()}:${mktConfig.collateralMint.toBase58()}` + : null; + + const [walletBalanceSnapshot, setWalletBalanceSnapshot] = useState( + null, + ); + + const walletBalance = mockMode + ? 500_000_000n + : walletBalanceSnapshot?.scopeKey === walletBalanceScopeKey + ? walletBalanceSnapshot.amount + : null; + const maxRawRef = useRef(null); - const [onChainDecimals, setOnChainDecimals] = useState(null); + + // A typed or Max-derived amount belongs to the wallet/mint scope that + // produced it. Clear both representations immediately after that scope + // changes so a replacement wallet cannot submit the previous value. + useEffect(() => { + maxRawRef.current = null; + setAmount(""); + }, [walletBalanceScopeKey]); + + const onChainDecimals = + !mockMode && walletBalanceSnapshot?.scopeKey === walletBalanceScopeKey + ? walletBalanceSnapshot.decimals + : null; + const decimals = onChainDecimals ?? tokenMeta?.decimals ?? 6; // Keep the mode-specific MAX raw value from leaking across Deposit/Withdraw. @@ -76,33 +110,68 @@ export const DepositWithdrawCard: FC = ({ slabAddress, // both the display amount and the raw ref before the opposite action can submit. useEffect(() => { setMode(initialMode); - setAmount(""); + setAmount(''); maxRawRef.current = null; }, [initialMode]); - const switchMode = (nextMode: "deposit" | "withdraw") => { + const switchMode = (nextMode: 'deposit' | 'withdraw') => { if (nextMode === mode) return; maxRawRef.current = null; - setAmount(""); + setAmount(''); setMode(nextMode); }; useEffect(() => { - if (!publicKey || !mktConfig?.collateralMint) { setWalletBalance(null); setOnChainDecimals(null); return; } + if (mockMode) return; + + if (!publicKey || !mktConfig?.collateralMint || !walletBalanceScopeKey) { + setWalletBalanceSnapshot(null); + return; + } + + const requestScopeKey = walletBalanceScopeKey; let cancelled = false; + + // Immediately invalidate a snapshot owned by a different wallet or mint. + // Same-scope refreshes retain their verified value while a post-transaction + // balance refresh is pending. + setWalletBalanceSnapshot((current) => + current?.scopeKey === requestScopeKey + ? current + : { + scopeKey: requestScopeKey, + amount: null, + decimals: null, + }, + ); + (async () => { try { const ata = getAssociatedTokenAddressSync(mktConfig.collateralMint, publicKey); + const info = await connection.getTokenAccountBalance(ata); + if (!cancelled && info.value.amount) { - setWalletBalance(BigInt(info.value.amount)); - if (info.value.decimals !== undefined) { - setOnChainDecimals(info.value.decimals); - } + setWalletBalanceSnapshot({ + scopeKey: requestScopeKey, + amount: BigInt(info.value.amount), + decimals: info.value.decimals ?? null, + }); } - } catch { if (!cancelled) { setWalletBalance(null); setOnChainDecimals(null); } } + } catch { + if (!cancelled) { + setWalletBalanceSnapshot({ + scopeKey: requestScopeKey, + amount: null, + decimals: null, + }); + } + } })(); - return () => { cancelled = true; }; - }, [publicKey, mktConfig?.collateralMint, connection, lastSig]); + + return () => { + cancelled = true; + }; + }, [mockMode, publicKey, mktConfig?.collateralMint, walletBalanceScopeKey, connection, lastSig]); // Pre-fill deposit: the FIRST time this card is open for a brand-new // (0-capital) account with a known wallet balance, default the amount @@ -248,6 +317,8 @@ export const DepositWithdrawCard: FC = ({ slabAddress, const freeMargin = capital > lockedMargin ? capital - lockedMargin : 0n; const loading = mode === "deposit" ? depositLoading : withdrawLoading; const error = mode === "deposit" ? depositError : withdrawError; + const isDepositBalanceUnverified = + !mockMode && mode === "deposit" && walletBalance === null; let parsedAmount: bigint = 0n; let parseError: string | null = null; @@ -271,7 +342,7 @@ export const DepositWithdrawCard: FC = ({ slabAddress, : null; async function handleSubmit() { - if (!amount || !userAccount || validationError) return; + if (!amount || !userAccount || validationError || isDepositBalanceUnverified) return; if (mockMode) { setAmount(""); return; } try { const amtNative = maxRawRef.current ?? parseHumanAmount(amount, decimals); @@ -408,7 +479,7 @@ export const DepositWithdrawCard: FC = ({ slabAddress,