import type { ExplorerFactory } from "@starknet-start/explorers"; import type { ChainProviderFactory } from "@starknet-start/providers"; import type { AccountInterface, PaymasterRpc, ProviderInterface } from "starknet"; import { StandardEvents, type StandardEventsListeners } from "@starknet-io/get-starknet-core"; import { GetStarknetProvider, type UseConnect, useConnect as useGetStarknetConnect, useStarknetProvider, } from "@starknet-io/get-starknet-modal"; import { StarknetWalletApi } from "@starknet-io/get-starknet-wallet-standard/features"; import { type Address, type Chain, mainnet, sepolia } from "@starknet-start/chains"; import { avnuPaymasterProvider, type ChainPaymasterFactory } from "@starknet-start/providers/paymaster"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { createContext, useCallback, useContext, useEffect, useMemo, useState } from "react"; import { constants, WalletAccountV5 } from "starknet"; import { AccountProvider } from "./account"; type Simplify = { [K in keyof T]: T[K] } & {}; type GetStarknetState = ReturnType; type GetStarknetProviderProps = Parameters[0]; const defaultQueryClient = new QueryClient(); export type StarknetState = Simplify< { chains: Chain[]; chain: Chain; explorer?: ExplorerFactory; provider: ProviderInterface; paymasterProvider?: PaymasterRpc; error?: Error; } & UseConnect & GetStarknetState >; const StarknetContext = createContext(undefined); export type StarknetProviderProps = Simplify< StarknetProviderInnerProps & Omit & { children?: React.ReactNode } >; type StarknetProviderInnerProps = { /** Chains supported by the app. */ chains: Chain[]; /** Provider to use. */ provider: ChainProviderFactory; /** Paymaster provider to use. */ paymasterProvider?: ChainPaymasterFactory; /** Explorer to use. */ explorer?: ExplorerFactory; /** Connect the first available connector on page load. */ autoConnect?: boolean; /** React-query client to use. */ queryClient?: QueryClient; /** Application. */ children?: React.ReactNode; /** Default chain to use when wallet is not connected */ defaultChainId?: bigint; }; export function StarknetProvider(props: StarknetProviderProps) { const { recommendedWallets, extraWallets, store, children, ...rest } = props; return ( {children} ); } function StarknetProviderInner({ chains, provider, // autoConnect, children, defaultChainId, explorer, paymasterProvider, queryClient, }: StarknetProviderInnerProps) { const { connect, disconnect, isConnecting, isError, connected } = useGetStarknetConnect(); const { extraWallets, injectedWallets, onSelectedChange, recommendedWallets, wallets, selected } = useStarknetProvider(); const defaultChain = defaultChainId ? (chains.find((c) => c.id === defaultChainId) ?? chains[0]) : chains[0]; if (defaultChain === undefined) { throw new Error("Must provide at least one chain."); } // check for duplicated ids in the chains list const seen = new Set(); for (const chain of chains) { if (seen.has(chain.id)) { throw new Error(`Duplicated chain id found: ${chain.id}`); } seen.add(chain.id); } const { provider: defaultProvider } = useMemo( () => providerForChain(defaultChain, provider), [defaultChain, provider], ); const _paymasterProvider = useMemo(() => paymasterProvider ?? avnuPaymasterProvider({}), [paymasterProvider]); const { paymasterProvider: defaultPaymasterProvider } = useMemo( () => paymasterProviderForChain(defaultChain, _paymasterProvider), [defaultChain, _paymasterProvider], ); const [currentChain, setCurrentChain] = useState(defaultChain); const [currentProvider, setCurrentProvider] = useState(defaultProvider); const [currentPaymasterProvider, setCurrentPaymasterProvider] = useState( defaultPaymasterProvider, ); const [address, setAddress] = useState
(); const [account, setAccount] = useState(); const updateChainAndProvider = useCallback( (chainId: bigint) => { const targetChain = chains.find((c) => c.id === chainId); if (!targetChain) { return; } const { provider: newProvider } = providerForChain(targetChain, provider); const { paymasterProvider: newPaymasterProvider } = paymasterProviderForChain(targetChain, _paymasterProvider); setCurrentChain(targetChain); setCurrentProvider(newProvider); setCurrentPaymasterProvider(newPaymasterProvider); }, [chains, provider, _paymasterProvider], ); const handleChange: StandardEventsListeners["change"] = useCallback( (change) => { if (change.accounts && change.accounts.length > 0) { const account = change.accounts[0]; setAddress(account.address as Address); try { const chainIdentifier = account.chains[0]; const parts = chainIdentifier.split(":"); const chainIdHex = parts[parts.length - 1]; const chainId = BigInt(chainIdHex); updateChainAndProvider(chainId); } catch (error) { console.error("Failed to parse chain ID:", error); } } }, [updateChainAndProvider], ); useEffect(() => { let cleanup: (() => void) | undefined; if (connected) { const walletAddress = connected.accounts?.[0]?.address as Address; setAddress(walletAddress); if (connected.accounts?.[0]?.chains?.[0]) { try { const chainIdentifier = connected.accounts[0].chains[0]; const parts = chainIdentifier.split(":"); const chainIdHex = parts[parts.length - 1]; const chainId = BigInt(chainIdHex); let targetChainId: bigint; if (defaultChainId) { targetChainId = defaultChainId; } else if (chains.length > 0) { targetChainId = chains[0].id; } else { targetChainId = chainId; } if (chainId !== targetChainId) { const targetChain = chains.find((c) => c.id === targetChainId); if (targetChain) { updateChainAndProvider(targetChainId); const targetStarknetChainId = starknetChainId(targetChainId); if (targetStarknetChainId) { connected.features[StarknetWalletApi] .request({ type: "wallet_switchStarknetChain", params: { chainId: targetStarknetChainId }, }) .catch((error: Error) => { console.warn("Failed to switch wallet to target chain:", error); updateChainAndProvider(chainId); }); } } else { updateChainAndProvider(chainId); } } else { updateChainAndProvider(chainId); } } catch (error) { console.error("Failed to parse chain ID:", error); } } else if (defaultChainId) { updateChainAndProvider(defaultChainId); } else if (chains.length > 0) { updateChainAndProvider(chains[0].id); } cleanup = connected.features[StandardEvents].on("change", handleChange); } else { setAddress(undefined); setAccount(undefined); setCurrentChain(defaultChain); setCurrentProvider(defaultProvider); setCurrentPaymasterProvider(defaultPaymasterProvider); } return () => { cleanup?.(); }; }, [ defaultChain, defaultChainId, defaultPaymasterProvider, defaultProvider, connected, chains, updateChainAndProvider, handleChange, ]); useEffect(() => { if (connected && address) { setAccount( new WalletAccountV5({ address, provider: currentProvider, walletProvider: connected, paymaster: currentPaymasterProvider, }), ); } }, [connected, address, currentProvider, currentPaymasterProvider]); const state: StarknetState = useMemo(() => { return { connect, disconnect, isConnecting, isError, connected, chains, chain: currentChain, explorer, provider: currentProvider, paymasterProvider: currentPaymasterProvider, error: undefined, extraWallets, injectedWallets, onSelectedChange, recommendedWallets, wallets, selected, }; }, [ connect, disconnect, isConnecting, isError, connected, chains, currentChain, explorer, currentProvider, currentPaymasterProvider, extraWallets, injectedWallets, onSelectedChange, recommendedWallets, wallets, selected, ]); return ( {children} ); } export function useStarknet(): StarknetState { const context = useContext(StarknetContext); if (!context) { throw new Error("useStarknet must be used within a StarknetProvider"); } return context; } export function useStarknetManager() { const { connect, disconnect } = useStarknet(); return { connect, disconnect }; } function providerForChain(chain: Chain, factory: ChainProviderFactory): { chain: Chain; provider: ProviderInterface } { const provider = factory(chain); if (provider) { return { chain, provider }; } throw new Error(`No provider found for chain ${chain.name}`); } function paymasterProviderForChain( chain: Chain, factory: ChainPaymasterFactory, ): { chain: Chain; paymasterProvider: PaymasterRpc } { const paymasterProvider = factory(chain); if (paymasterProvider) { return { chain, paymasterProvider }; } throw new Error(`No paymaster provider found for chain ${chain.name}`); } export function starknetChainId(chainId: bigint): constants.StarknetChainId | undefined { switch (chainId) { case mainnet.id: return constants.StarknetChainId.SN_MAIN; case sepolia.id: return constants.StarknetChainId.SN_SEPOLIA; default: return undefined; } }