File size: 5,980 Bytes
bde2f3a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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