feat: add wallets balance
This commit is contained in:
373
src/infrastructure/wallet_balances/service.py
Normal file
373
src/infrastructure/wallet_balances/service.py
Normal file
@@ -0,0 +1,373 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
from dataclasses import dataclass
|
||||
from decimal import Decimal, InvalidOperation
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from src.infrastructure.config import settings
|
||||
from src.infrastructure.wallet_balances.token_registry import (
|
||||
NATIVE_DECIMALS,
|
||||
NATIVE_SYMBOLS,
|
||||
get_coingecko_id,
|
||||
get_evm_tokens,
|
||||
get_sol_tokens,
|
||||
get_trx_tokens,
|
||||
)
|
||||
|
||||
TIMEOUT_SECONDS = 15
|
||||
BLOCKSTREAM_URL = 'https://blockstream.info/api'
|
||||
TRONGRID_URL = 'https://api.trongrid.io'
|
||||
COINGECKO_URL = 'https://api.coingecko.com/api/v3/simple/price'
|
||||
BALANCE_OF_SELECTOR = '70a08231'
|
||||
BASE58_ALPHABET = '123456789ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnopqrstuvwxyz'
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class FormattedAmount:
|
||||
raw: str
|
||||
formatted: str
|
||||
decimals: int
|
||||
usd_price: Decimal | None = None
|
||||
usd_value: Decimal | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class WalletBalance:
|
||||
wallet_id: str
|
||||
chain: str
|
||||
address: str
|
||||
derivation_path: str
|
||||
native_symbol: str
|
||||
native: FormattedAmount
|
||||
tokens: dict[str, FormattedAmount]
|
||||
total_usd: Decimal | None
|
||||
error: str | None = None
|
||||
|
||||
|
||||
class WalletBalanceService:
|
||||
async def get_wallet_balances(self, wallets) -> list[WalletBalance]:
|
||||
prices = await self._get_prices(wallets)
|
||||
results = await asyncio.gather(
|
||||
*(self._get_wallet_balance(wallet, prices) for wallet in wallets),
|
||||
return_exceptions=True,
|
||||
)
|
||||
balances: list[WalletBalance] = []
|
||||
for wallet, result in zip(wallets, results, strict=False):
|
||||
if isinstance(result, Exception):
|
||||
balances.append(self._failed_wallet(wallet, str(result)))
|
||||
else:
|
||||
balances.append(result)
|
||||
return balances
|
||||
|
||||
async def _get_wallet_balance(self, wallet, prices: dict[str, Decimal | None]) -> WalletBalance:
|
||||
chain = wallet.chain.upper()
|
||||
if chain == 'BTC':
|
||||
native_raw, token_raw = await self._btc_balance(wallet.address)
|
||||
elif chain == 'ETH':
|
||||
native_raw, token_raw = await self._evm_balance(
|
||||
settings.ETH_RPC_URL,
|
||||
wallet.address,
|
||||
get_evm_tokens(chain),
|
||||
)
|
||||
elif chain == 'BSC':
|
||||
native_raw, token_raw = await self._evm_balance(
|
||||
settings.BSC_RPC_URL,
|
||||
wallet.address,
|
||||
get_evm_tokens(chain),
|
||||
)
|
||||
elif chain == 'TRX':
|
||||
native_raw, token_raw = await self._trx_balance(wallet.address)
|
||||
elif chain == 'SOL':
|
||||
native_raw, token_raw = await self._sol_balance(wallet.address)
|
||||
else:
|
||||
raise ValueError(f'Unsupported chain: {wallet.chain}')
|
||||
|
||||
native = self._amount(native_raw, NATIVE_DECIMALS[chain], prices.get(f'{chain}:{NATIVE_SYMBOLS[chain]}'))
|
||||
tokens = self._token_amounts(chain, token_raw, prices)
|
||||
total_usd = self._total_usd(native, tokens)
|
||||
return WalletBalance(
|
||||
wallet_id=wallet.id,
|
||||
chain=chain,
|
||||
address=wallet.address,
|
||||
derivation_path=wallet.derivation_path,
|
||||
native_symbol=NATIVE_SYMBOLS[chain],
|
||||
native=native,
|
||||
tokens=tokens,
|
||||
total_usd=total_usd,
|
||||
)
|
||||
|
||||
async def _btc_balance(self, address: str) -> tuple[str, dict[str, str]]:
|
||||
data = await self._get_json(f'{BLOCKSTREAM_URL}/address/{address}')
|
||||
stats = data.get('chain_stats') or {}
|
||||
funded = int(stats.get('funded_txo_sum') or 0)
|
||||
spent = int(stats.get('spent_txo_sum') or 0)
|
||||
return str(funded - spent), {}
|
||||
|
||||
async def _evm_balance(self, rpc_url: str, address: str, tokens: list[dict]) -> tuple[str, dict[str, str]]:
|
||||
native_result = await self._rpc(rpc_url, 'eth_getBalance', [address, 'latest'])
|
||||
native = self._hex_to_int_string(native_result)
|
||||
token_items = await asyncio.gather(
|
||||
*(self._evm_token_balance(rpc_url, address, token) for token in tokens),
|
||||
return_exceptions=True,
|
||||
)
|
||||
result: dict[str, str] = {}
|
||||
for token, item in zip(tokens, token_items, strict=False):
|
||||
result[token['symbol']] = '0' if isinstance(item, Exception) else item
|
||||
return native, result
|
||||
|
||||
async def _evm_token_balance(self, rpc_url: str, address: str, token: dict) -> str:
|
||||
data = '0x' + BALANCE_OF_SELECTOR + address.lower().removeprefix('0x').rjust(64, '0')
|
||||
result = await self._rpc(rpc_url, 'eth_call', [{'to': token['contract'], 'data': data}, 'latest'])
|
||||
return self._hex_to_int_string(result)
|
||||
|
||||
async def _trx_balance(self, address: str) -> tuple[str, dict[str, str]]:
|
||||
headers = self._tron_headers()
|
||||
account = await self._get_json(f'{TRONGRID_URL}/v1/accounts/{address}', headers=headers)
|
||||
trx = str(((account.get('data') or [{}])[0] or {}).get('balance') or 0)
|
||||
address_hex = self._tron_address_to_hex(address).rjust(64, '0')
|
||||
token_config = self._tokens_for_chain('TRX')
|
||||
token_items = await asyncio.gather(
|
||||
*(self._trx_token_balance(address, address_hex, token, headers) for token in token_config),
|
||||
return_exceptions=True,
|
||||
)
|
||||
tokens: dict[str, str] = {}
|
||||
for token, item in zip(token_config, token_items, strict=False):
|
||||
tokens[token['symbol']] = '0' if isinstance(item, Exception) else item
|
||||
return trx, tokens
|
||||
|
||||
async def _trx_token_balance(self, address: str, address_hex: str, token: dict, headers: dict[str, str]) -> str:
|
||||
data = await self._post_json(
|
||||
f'{TRONGRID_URL}/wallet/triggerconstantcontract',
|
||||
{
|
||||
'owner_address': address,
|
||||
'contract_address': token['contract'],
|
||||
'function_selector': 'balanceOf(address)',
|
||||
'parameter': address_hex,
|
||||
'visible': True,
|
||||
},
|
||||
headers={**headers, 'Content-Type': 'application/json'},
|
||||
)
|
||||
value = (data.get('constant_result') or [None])[0]
|
||||
if not value or set(value) == {'0'}:
|
||||
return '0'
|
||||
return str(int(value, 16))
|
||||
|
||||
async def _sol_balance(self, address: str) -> tuple[str, dict[str, str]]:
|
||||
native_data = await self._rpc(settings.SOL_RPC_URL, 'getBalance', [address])
|
||||
native = str((native_data or {}).get('value') or 0)
|
||||
tokens = {token['symbol']: '0' for token in get_sol_tokens()}
|
||||
try:
|
||||
accounts = await self._rpc(
|
||||
settings.SOL_RPC_URL,
|
||||
'getTokenAccountsByOwner',
|
||||
[
|
||||
address,
|
||||
{'programId': 'TokenkegQfeZyiNwAJbNbGKPFXCWuBvf9Ss623VQ5DA'},
|
||||
{'encoding': 'jsonParsed'},
|
||||
],
|
||||
)
|
||||
known = {token['mint']: token['symbol'] for token in get_sol_tokens()}
|
||||
for item in accounts.get('value') or []:
|
||||
info = (((item.get('account') or {}).get('data') or {}).get('parsed') or {}).get('info') or {}
|
||||
mint = info.get('mint')
|
||||
amount = ((info.get('tokenAmount') or {}).get('amount')) or '0'
|
||||
if mint in known:
|
||||
symbol = known[mint]
|
||||
tokens[symbol] = str(int(tokens[symbol]) + int(amount))
|
||||
except Exception:
|
||||
pass
|
||||
return native, tokens
|
||||
|
||||
async def _get_prices(self, wallets) -> dict[str, Decimal | None]:
|
||||
pairs: dict[str, str] = {}
|
||||
for wallet in wallets:
|
||||
chain = wallet.chain.upper()
|
||||
symbol = NATIVE_SYMBOLS.get(chain)
|
||||
coin_id = get_coingecko_id(chain, symbol or chain)
|
||||
if symbol and coin_id:
|
||||
pairs[f'{chain}:{symbol}'] = coin_id
|
||||
for token in self._tokens_for_chain(chain):
|
||||
coin_id = get_coingecko_id(chain, token['symbol'])
|
||||
if coin_id:
|
||||
pairs[f'{chain}:{token['symbol']}'] = coin_id
|
||||
if not pairs:
|
||||
return {}
|
||||
|
||||
try:
|
||||
response = await self._get_json(
|
||||
COINGECKO_URL,
|
||||
params={
|
||||
'ids': ','.join(sorted(set(pairs.values()))),
|
||||
'vs_currencies': 'usd',
|
||||
},
|
||||
headers=self._coingecko_headers(),
|
||||
timeout=5,
|
||||
)
|
||||
except Exception:
|
||||
return {key: None for key in pairs}
|
||||
|
||||
return {
|
||||
key: self._decimal_or_none((response.get(coin_id) or {}).get('usd'))
|
||||
for key, coin_id in pairs.items()
|
||||
}
|
||||
|
||||
def _token_amounts(
|
||||
self,
|
||||
chain: str,
|
||||
raw: dict[str, str],
|
||||
prices: dict[str, Decimal | None],
|
||||
) -> dict[str, FormattedAmount]:
|
||||
decimals = {token['symbol']: token['decimals'] for token in self._tokens_for_chain(chain)}
|
||||
return {
|
||||
symbol: self._amount(value, decimals.get(symbol, 0), prices.get(f'{chain}:{symbol}'))
|
||||
for symbol, value in raw.items()
|
||||
}
|
||||
|
||||
def _amount(self, raw: str, decimals: int, usd_price: Decimal | None) -> FormattedAmount:
|
||||
formatted = self._format_units(raw, decimals)
|
||||
usd_value = None
|
||||
if usd_price is not None:
|
||||
usd_value = self._round_usd(Decimal(formatted) * usd_price)
|
||||
return FormattedAmount(
|
||||
raw=raw,
|
||||
formatted=formatted,
|
||||
decimals=decimals,
|
||||
usd_price=usd_price,
|
||||
usd_value=usd_value,
|
||||
)
|
||||
|
||||
def _total_usd(self, native: FormattedAmount, tokens: dict[str, FormattedAmount]) -> Decimal | None:
|
||||
values = [native.usd_value, *(token.usd_value for token in tokens.values())]
|
||||
present = [value for value in values if value is not None]
|
||||
if not present:
|
||||
return None
|
||||
return self._round_usd(sum(present, Decimal('0')))
|
||||
|
||||
def _failed_wallet(self, wallet, error: str) -> WalletBalance:
|
||||
chain = wallet.chain.upper()
|
||||
native = FormattedAmount(
|
||||
raw='0',
|
||||
formatted='0',
|
||||
decimals=NATIVE_DECIMALS.get(chain, 0),
|
||||
)
|
||||
return WalletBalance(
|
||||
wallet_id=wallet.id,
|
||||
chain=chain,
|
||||
address=wallet.address,
|
||||
derivation_path=wallet.derivation_path,
|
||||
native_symbol=NATIVE_SYMBOLS.get(chain, chain),
|
||||
native=native,
|
||||
tokens={},
|
||||
total_usd=None,
|
||||
error=error,
|
||||
)
|
||||
|
||||
async def _rpc(self, url: str, method: str, params: list[Any]) -> Any:
|
||||
payload = {
|
||||
'jsonrpc': '2.0',
|
||||
'id': 1,
|
||||
'method': method,
|
||||
'params': params,
|
||||
}
|
||||
response = await self._post_json(url, payload)
|
||||
if response.get('error'):
|
||||
raise RuntimeError(str(response['error']))
|
||||
return response.get('result')
|
||||
|
||||
async def _get_json(
|
||||
self,
|
||||
url: str,
|
||||
*,
|
||||
params: dict[str, str] | None = None,
|
||||
headers: dict[str, str] | None = None,
|
||||
timeout: int = TIMEOUT_SECONDS,
|
||||
) -> Any:
|
||||
async with httpx.AsyncClient(timeout=timeout) as client:
|
||||
response = await client.get(url, params=params, headers=headers)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def _post_json(
|
||||
self,
|
||||
url: str,
|
||||
payload: dict[str, Any],
|
||||
*,
|
||||
headers: dict[str, str] | None = None,
|
||||
timeout: int = TIMEOUT_SECONDS,
|
||||
) -> Any:
|
||||
async with httpx.AsyncClient(timeout=timeout) as client:
|
||||
response = await client.post(url, json=payload, headers=headers)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
def _tokens_for_chain(self, chain: str) -> list[dict]:
|
||||
if chain in {'ETH', 'BSC'}:
|
||||
return get_evm_tokens(chain)
|
||||
if chain == 'TRX':
|
||||
tokens = [dict(token) for token in get_trx_tokens()]
|
||||
if settings.USDT_CONTRACT_ADDRESS:
|
||||
for token in tokens:
|
||||
if token['symbol'] == 'USDT':
|
||||
token['contract'] = settings.USDT_CONTRACT_ADDRESS
|
||||
return tokens
|
||||
if chain == 'SOL':
|
||||
return get_sol_tokens()
|
||||
return []
|
||||
|
||||
def _format_units(self, raw: str, decimals: int) -> str:
|
||||
if not raw or not raw.lstrip('-').isdigit():
|
||||
return '0'
|
||||
if decimals == 0:
|
||||
return raw
|
||||
negative = raw.startswith('-')
|
||||
value = raw[1:] if negative else raw
|
||||
padded = value.rjust(decimals + 1, '0')
|
||||
whole = padded[:-decimals]
|
||||
fraction = padded[-decimals:].rstrip('0')
|
||||
result = f'{whole}.{fraction}' if fraction else whole
|
||||
return f'-{result}' if negative else result
|
||||
|
||||
def _round_usd(self, value: Decimal) -> Decimal:
|
||||
return value.quantize(Decimal('0.00000001'))
|
||||
|
||||
def _decimal_or_none(self, value: Any) -> Decimal | None:
|
||||
try:
|
||||
return Decimal(str(value))
|
||||
except (InvalidOperation, TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def _hex_to_int_string(self, value: str | None) -> str:
|
||||
if not value:
|
||||
return '0'
|
||||
return str(int(value, 16))
|
||||
|
||||
def _tron_headers(self) -> dict[str, str]:
|
||||
headers = {'Accept': 'application/json'}
|
||||
if settings.TRON_API_KEY:
|
||||
headers['TRON-PRO-API-KEY'] = settings.TRON_API_KEY
|
||||
return headers
|
||||
|
||||
def _coingecko_headers(self) -> dict[str, str]:
|
||||
headers = {'Accept': 'application/json'}
|
||||
if settings.COINGECKO_API_KEY:
|
||||
headers['x-cg-demo-api-key'] = settings.COINGECKO_API_KEY
|
||||
return headers
|
||||
|
||||
def _tron_address_to_hex(self, address: str) -> str:
|
||||
num = 0
|
||||
for char in address:
|
||||
index = BASE58_ALPHABET.find(char)
|
||||
if index == -1:
|
||||
raise ValueError('Invalid TRON address')
|
||||
num = num * 58 + index
|
||||
raw = num.to_bytes(25, 'big')
|
||||
payload = raw[:-4]
|
||||
checksum = raw[-4:]
|
||||
expected = hashlib.sha256(hashlib.sha256(payload).digest()).digest()[:4]
|
||||
if checksum != expected or payload[0] != 0x41:
|
||||
raise ValueError('Invalid TRON address')
|
||||
return payload[1:].hex()
|
||||
Reference in New Issue
Block a user