File size: 5,980 Bytes
6993919 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 | """
Input Validation Middleware
===========================
Comprehensive Pydantic-based validation for security endpoints.
Includes:
- Request sanitization
- Type coercion
- Length/range validation
- Regex pattern validation
- Custom validators
"""
import logging
import re
from pydantic import BaseModel, Field, validator
logger = logging.getLogger(__name__)
# βββ COMMON VALIDATION CONSTANTS βββββββββββββββββββββββββββββββββ
ADDRESS_PATTERN = re.compile(r"^(0x)?[0-9a-fA-F]{38,64}$")
HASH_PATTERN = re.compile(r"^(0x)?[0-9a-fA-F]{64}$")
WALLET_PATTERN = re.compile(r"^(0x)?[0-9a-fA-F]{38,64}$")
CHAIN_NAME_PATTERN = re.compile(r"^[a-z][a-z0-9-]+$")
MIN_ADDRESS_LENGTH = 38 # 0x + 38 hex chars = 20 bytes (EVM)
MAX_ADDRESS_LENGTH = 66 # 0x + 66 hex chars = 32 bytes (some chains)
MIN_CHAIN_LENGTH = 2
MAX_CHAIN_LENGTH = 20
# Chain names whitelist
VALID_CHAINS = {
"ethereum",
"base",
"arbitrum",
"optimism",
"polygon",
"bsc",
"avalanche",
"fantom",
"celo",
"moonbeam",
"polygon_zkevm",
"linea",
"scroll",
"zksync",
"blast",
"Mode",
"sepolia",
"holesky",
}
# βββ BASE REQUEST MODELS ββββββββββββββββββββββββββββββββββββββββ
class WalletRequest(BaseModel):
"""Base model for wallet address requests."""
address: str = Field(
...,
description="Wallet address (EVM format: 0x...)",
min_length=MIN_ADDRESS_LENGTH,
max_length=MAX_ADDRESS_LENGTH,
)
chain: str = Field(default="ethereum", description="Blockchain network name")
@validator("address")
def validate_address(self, v: str) -> str:
"""Validate wallet address format."""
if not v.startswith("0x"):
v = "0x" + v
if not ADDRESS_PATTERN.match(v):
raise ValueError(f"Invalid address format: {v}")
return v.lower()
@validator("chain")
def validate_chain(self, v: str) -> str:
"""Validate chain name."""
if v not in VALID_CHAINS:
logger.warning(f"Unknown chain '{v}', allowing anyway for flexibility")
return v.lower()
class BatchWalletRequest(BaseModel):
"""Request for batch wallet operations."""
addresses: list[str] = Field(..., min_length=1, max_length=50, description="List of wallet addresses")
chain: str = Field(default="ethereum", description="Blockchain network")
@validator("addresses", each_item=True)
def validate_address(self, v: str) -> str:
"""Validate each address in batch."""
if not v.startswith("0x"):
v = "0x" + v
if not ADDRESS_PATTERN.match(v):
raise ValueError(f"Invalid address format: {v}")
return v.lower()
@validator("chain")
def validate_chain(self, v: str) -> str:
"""Validate chain name."""
if v not in VALID_CHAINS:
logger.warning(f"Unknown chain '{v}', allowing anyway")
return v.lower()
class ContractRequest(WalletRequest):
"""Request for contract operations."""
deep: bool = Field(default=False, description="Run deep scan with Slither/Mythril")
class TextSearchRequest(BaseModel):
"""Request for text-based searches."""
query: str = Field(..., min_length=1, max_length=200, description="Search query")
limit: int = Field(default=10, ge=1, le=100, description="Maximum results to return")
# βββ SECURITY-FOCUSED MODELS ββββββββββββββββββββββββββββββββββββ
class ThreatCheckRequest(WalletRequest):
"""Request for threat intelligence checks."""
check_sanctions: bool = Field(default=True, description="Check OFAC sanctions list")
check_risk: bool = Field(default=True, description="Check risk score")
check_exchange: bool = Field(default=True, description="Check if address is exchange/deposit")
class MempoolMonitorRequest(WalletRequest):
"""Request for mempool monitoring."""
threshold_gas_gwei: float = Field(default=50.0, ge=1.0, le=1000.0, description="Gas price threshold in Gwei")
alert_on_bot: bool = Field(default=True, description="Alert on bot detection")
class PortfolioRequest(WalletRequest):
"""Request for portfolio analysis."""
include_tokens: bool = Field(default=True, description="Include ERC-20 tokens")
include_nfts: bool = Field(default=False, description="Include NFTs")
history_days: int = Field(default=30, ge=1, le=365, description="Days of transaction history")
# βββ UTILITY VALIDATORS βββββββββββββββββββββββββββββββββββββββββ
def tier_validator(tier: str, rate_type: str = "requests_per_hour") -> dict[str, int]:
"""
Get rate limit for authentication tier.
Args:
tier: Authentication tier (FREE, BASIC, PREMIUM, ENTERPRISE)
rate_type: Rate type to check
Returns:
Rate limit configuration
"""
tiers = {
"FREE": {"requests_per_hour": 10, "requests_per_day": 100},
"BASIC": {"requests_per_hour": 50, "requests_per_day": 500},
"PREMIUM": {"requests_per_hour": 200, "requests_per_day": 2000},
"ENTERPRISE": {"requests_per_hour": 1000, "requests_per_day": 10000},
}
return tiers.get(tier.upper(), tiers["FREE"])
def validate_x402_payment(required: int, tier: str) -> bool:
"""
Validate x402 micropayment capability based on tier.
Args:
required: Required x402 tokens
tier: User's authentication tier
Returns:
True if payment capability is sufficient
"""
tier_caps = {
"FREE": 10,
"BASIC": 100,
"PREMIUM": 1000,
"ENTERPRISE": 10000,
}
available = tier_caps.get(tier.upper(), 0)
return available >= required
|