wip(core): platform aware trading

This commit is contained in:
smypmsa
2025-08-02 13:27:07 +00:00
parent 9ec9eeb793
commit 227ef0b3bd
9 changed files with 699 additions and 506 deletions
+137 -99
View File
@@ -12,18 +12,14 @@ from config_loader import (
get_platform_from_config, get_platform_from_config,
load_bot_config, load_bot_config,
print_config_summary, print_config_summary,
validate_platform_listener_combination,
) )
from trading.universal_trader import UniversalTrader from trading.universal_trader import UniversalTrader
from utils.logger import setup_file_logging from utils.logger import setup_file_logging
def setup_logging(bot_name: str): def setup_logging(bot_name: str):
""" """Set up logging to file for a specific bot instance."""
Set up logging to file for a specific bot instance.
Args:
bot_name: Name of the bot for the log file
"""
log_dir = Path("logs") log_dir = Path("logs")
log_dir.mkdir(exist_ok=True) log_dir.mkdir(exist_ok=True)
@@ -34,96 +30,119 @@ def setup_logging(bot_name: str):
async def start_bot(config_path: str): async def start_bot(config_path: str):
""" """Start a trading bot with the configuration from the specified path."""
Start a trading bot with the configuration from the specified path.
Args:
config_path: Path to the YAML configuration file
"""
cfg = load_bot_config(config_path) cfg = load_bot_config(config_path)
setup_logging(cfg["name"]) setup_logging(cfg["name"])
print_config_summary(cfg) print_config_summary(cfg)
# Get platform from configuration # Get and validate platform from configuration
platform = get_platform_from_config(cfg) try:
platform = get_platform_from_config(cfg)
# Initialize universal trader with platform-specific configuration logging.info(f"Detected platform: {platform.value}")
trader = UniversalTrader( except ValueError as e:
# Connection settings logging.exception(f"Platform configuration error: {e}")
rpc_endpoint=cfg["rpc_endpoint"], return
wss_endpoint=cfg["wss_endpoint"],
private_key=cfg["private_key"],
# Platform configuration
platform=platform,
# Trade parameters
buy_amount=cfg["trade"]["buy_amount"],
buy_slippage=cfg["trade"]["buy_slippage"],
sell_slippage=cfg["trade"]["sell_slippage"],
# Extreme fast mode settings
extreme_fast_mode=cfg["trade"].get("extreme_fast_mode", False),
extreme_fast_token_amount=cfg["trade"].get("extreme_fast_token_amount", 30),
# Exit strategy configuration
exit_strategy=cfg["trade"].get("exit_strategy", "time_based"),
take_profit_percentage=cfg["trade"].get("take_profit_percentage"),
stop_loss_percentage=cfg["trade"].get("stop_loss_percentage"),
max_hold_time=cfg["trade"].get("max_hold_time"),
price_check_interval=cfg["trade"].get("price_check_interval", 10),
# Listener configuration
listener_type=cfg["filters"]["listener_type"],
# Geyser configuration (if applicable)
geyser_endpoint=cfg.get("geyser", {}).get("endpoint"),
geyser_api_token=cfg.get("geyser", {}).get("api_token"),
geyser_auth_type=cfg.get("geyser", {}).get("auth_type", "x-token"),
# PumpPortal configuration (if applicable)
pumpportal_url=cfg.get("pumpportal", {}).get(
"url", "wss://pumpportal.fun/api/data"
),
# Priority fee configuration
enable_dynamic_priority_fee=cfg.get("priority_fees", {}).get(
"enable_dynamic", False
),
enable_fixed_priority_fee=cfg.get("priority_fees", {}).get(
"enable_fixed", True
),
fixed_priority_fee=cfg.get("priority_fees", {}).get("fixed_amount", 500000),
extra_priority_fee=cfg.get("priority_fees", {}).get("extra_percentage", 0.0),
hard_cap_prior_fee=cfg.get("priority_fees", {}).get("hard_cap", 500000),
# Retry and timeout settings
max_retries=cfg.get("retries", {}).get("max_attempts", 10),
wait_time_after_creation=cfg.get("retries", {}).get("wait_after_creation", 15),
wait_time_after_buy=cfg.get("retries", {}).get("wait_after_buy", 15),
wait_time_before_new_token=cfg.get("retries", {}).get(
"wait_before_new_token", 15
),
max_token_age=cfg.get("filters", {}).get("max_token_age", 0.001),
token_wait_timeout=cfg.get("timing", {}).get("token_wait_timeout", 120),
# Cleanup settings
cleanup_mode=cfg.get("cleanup", {}).get("mode", "disabled"),
cleanup_force_close_with_burn=cfg.get("cleanup", {}).get(
"force_close_with_burn", False
),
cleanup_with_priority_fee=cfg.get("cleanup", {}).get(
"with_priority_fee", False
),
# Trading filters
match_string=cfg["filters"].get("match_string"),
bro_address=cfg["filters"].get("bro_address"),
marry_mode=cfg["filters"].get("marry_mode", False),
yolo_mode=cfg["filters"].get("yolo_mode", False),
)
await trader.start() # Validate platform support
try:
from platforms import platform_factory
if not platform_factory.registry.is_platform_supported(platform):
logging.error(f"Platform {platform.value} is not supported. Available platforms: {[p.value for p in platform_factory.get_supported_platforms()]}")
return
except Exception as e:
logging.exception(f"Could not validate platform support: {e}")
return
# Validate listener compatibility
listener_type = cfg["filters"]["listener_type"]
if not validate_platform_listener_combination(platform, listener_type):
from config_loader import get_supported_listeners_for_platform
supported = get_supported_listeners_for_platform(platform)
logging.error(f"Listener '{listener_type}' is not compatible with platform '{platform.value}'. Supported listeners: {supported}")
return
# Initialize universal trader with platform-specific configuration
try:
trader = UniversalTrader(
# Connection settings
rpc_endpoint=cfg["rpc_endpoint"],
wss_endpoint=cfg["wss_endpoint"],
private_key=cfg["private_key"],
# Platform configuration - pass platform enum directly
platform=platform,
# Trade parameters
buy_amount=cfg["trade"]["buy_amount"],
buy_slippage=cfg["trade"]["buy_slippage"],
sell_slippage=cfg["trade"]["sell_slippage"],
# Extreme fast mode settings
extreme_fast_mode=cfg["trade"].get("extreme_fast_mode", False),
extreme_fast_token_amount=cfg["trade"].get("extreme_fast_token_amount", 30),
# Exit strategy configuration
exit_strategy=cfg["trade"].get("exit_strategy", "time_based"),
take_profit_percentage=cfg["trade"].get("take_profit_percentage"),
stop_loss_percentage=cfg["trade"].get("stop_loss_percentage"),
max_hold_time=cfg["trade"].get("max_hold_time"),
price_check_interval=cfg["trade"].get("price_check_interval", 10),
# Listener configuration
listener_type=cfg["filters"]["listener_type"],
# Geyser configuration (if applicable)
geyser_endpoint=cfg.get("geyser", {}).get("endpoint"),
geyser_api_token=cfg.get("geyser", {}).get("api_token"),
geyser_auth_type=cfg.get("geyser", {}).get("auth_type", "x-token"),
# PumpPortal configuration (if applicable)
pumpportal_url=cfg.get("pumpportal", {}).get(
"url", "wss://pumpportal.fun/api/data"
),
# Priority fee configuration
enable_dynamic_priority_fee=cfg.get("priority_fees", {}).get(
"enable_dynamic", False
),
enable_fixed_priority_fee=cfg.get("priority_fees", {}).get(
"enable_fixed", True
),
fixed_priority_fee=cfg.get("priority_fees", {}).get("fixed_amount", 500000),
extra_priority_fee=cfg.get("priority_fees", {}).get("extra_percentage", 0.0),
hard_cap_prior_fee=cfg.get("priority_fees", {}).get("hard_cap", 500000),
# Retry and timeout settings
max_retries=cfg.get("retries", {}).get("max_attempts", 10),
wait_time_after_creation=cfg.get("retries", {}).get("wait_after_creation", 15),
wait_time_after_buy=cfg.get("retries", {}).get("wait_after_buy", 15),
wait_time_before_new_token=cfg.get("retries", {}).get(
"wait_before_new_token", 15
),
max_token_age=cfg.get("filters", {}).get("max_token_age", 0.001),
token_wait_timeout=cfg.get("timing", {}).get("token_wait_timeout", 120),
# Cleanup settings
cleanup_mode=cfg.get("cleanup", {}).get("mode", "disabled"),
cleanup_force_close_with_burn=cfg.get("cleanup", {}).get(
"force_close_with_burn", False
),
cleanup_with_priority_fee=cfg.get("cleanup", {}).get(
"with_priority_fee", False
),
# Trading filters
match_string=cfg["filters"].get("match_string"),
bro_address=cfg["filters"].get("bro_address"),
marry_mode=cfg["filters"].get("marry_mode", False),
yolo_mode=cfg["filters"].get("yolo_mode", False),
)
await trader.start()
except Exception as e:
logging.exception(f"Failed to initialize or start trader: {e}")
raise
def run_bot_process(config_path): def run_bot_process(config_path):
@@ -131,11 +150,7 @@ def run_bot_process(config_path):
def run_all_bots(): def run_all_bots():
""" """Run all bots defined in YAML files in the 'bots' directory."""
Run all bots defined in YAML files in the 'bots' directory.
Only runs bots that have enabled=True (or where enabled is not specified).
Bots can be run in separate processes based on their configuration.
"""
bot_dir = Path("bots") bot_dir = Path("bots")
if not bot_dir.exists(): if not bot_dir.exists():
logging.error(f"Bot directory '{bot_dir}' not found") logging.error(f"Bot directory '{bot_dir}' not found")
@@ -162,19 +177,32 @@ def run_all_bots():
skipped_bots += 1 skipped_bots += 1
continue continue
# Validate platform support # Validate platform configuration
try: try:
platform = get_platform_from_config(cfg) platform = get_platform_from_config(cfg)
# Check platform support
from platforms import platform_factory from platforms import platform_factory
if not platform_factory.registry.is_platform_supported(platform): if not platform_factory.registry.is_platform_supported(platform):
logging.error(f"Platform {platform.value} is not supported for bot '{bot_name}'. Skipping...") logging.error(f"Platform {platform.value} is not supported for bot '{bot_name}'. Available platforms: {[p.value for p in platform_factory.get_supported_platforms()]}")
skipped_bots += 1 skipped_bots += 1
continue continue
# Validate listener compatibility
listener_type = cfg["filters"]["listener_type"]
if not validate_platform_listener_combination(platform, listener_type):
from config_loader import get_supported_listeners_for_platform
supported = get_supported_listeners_for_platform(platform)
logging.error(f"Listener '{listener_type}' is not compatible with platform '{platform.value}' for bot '{bot_name}'. Supported listeners: {supported}")
skipped_bots += 1
continue
except Exception as e: except Exception as e:
logging.exception(f"Invalid platform configuration for bot '{bot_name}': {e}. Skipping...") logging.exception(f"Invalid platform configuration for bot '{bot_name}': {e}. Skipping...")
skipped_bots += 1 skipped_bots += 1
continue continue
# Start bot in separate process or main process
if cfg.get("separate_process", False): if cfg.get("separate_process", False):
logging.info(f"Starting bot '{bot_name}' ({platform.value}) in separate process") logging.info(f"Starting bot '{bot_name}' ({platform.value}) in separate process")
p = multiprocessing.Process( p = multiprocessing.Process(
@@ -185,13 +213,16 @@ def run_all_bots():
else: else:
logging.info(f"Starting bot '{bot_name}' ({platform.value}) in main process") logging.info(f"Starting bot '{bot_name}' ({platform.value}) in main process")
asyncio.run(start_bot(str(file))) asyncio.run(start_bot(str(file)))
except Exception as e: except Exception as e:
logging.exception(f"Failed to start bot from {file}: {e}") logging.exception(f"Failed to start bot from {file}: {e}")
skipped_bots += 1
logging.info( logging.info(
f"Started {len(bot_files) - skipped_bots} bots, skipped {skipped_bots} disabled/invalid bots" f"Started {len(bot_files) - skipped_bots} bots, skipped {skipped_bots} disabled/invalid bots"
) )
# Wait for all processes to complete
for p in processes: for p in processes:
p.join() p.join()
logging.info(f"Process {p.name} completed") logging.info(f"Process {p.name} completed")
@@ -203,11 +234,18 @@ def main() -> None:
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
) )
# Log supported platforms # Log supported platforms and listeners
try: try:
from platforms import platform_factory from platforms import platform_factory
supported_platforms = platform_factory.get_supported_platforms() supported_platforms = platform_factory.get_supported_platforms()
logging.info(f"Supported platforms: {[p.value for p in supported_platforms]}") logging.info(f"Supported platforms: {[p.value for p in supported_platforms]}")
# Log listener compatibility for each platform
from config_loader import get_supported_listeners_for_platform
for platform in supported_platforms:
listeners = get_supported_listeners_for_platform(platform)
logging.info(f"Platform {platform.value} supports listeners: {listeners}")
except Exception as e: except Exception as e:
logging.warning(f"Could not load platform information: {e}") logging.warning(f"Could not load platform information: {e}")
+174 -143
View File
@@ -1,8 +1,5 @@
""" """
Updated configuration validation with platform support. Updated configuration validation with comprehensive platform support.
This extends the existing config_loader.py to support platform selection
and validation while maintaining backward compatibility.
""" """
import os import os
@@ -27,77 +24,33 @@ REQUIRED_FIELDS = [
] ]
CONFIG_VALIDATION_RULES = [ CONFIG_VALIDATION_RULES = [
# (path, type, min_value, max_value, error_message) ("trade.buy_amount", (int, float), 0, float("inf"), "trade.buy_amount must be a positive number"),
(
"trade.buy_amount",
(int, float),
0,
float("inf"),
"trade.buy_amount must be a positive number",
),
("trade.buy_slippage", float, 0, 1, "trade.buy_slippage must be between 0 and 1"), ("trade.buy_slippage", float, 0, 1, "trade.buy_slippage must be between 0 and 1"),
("trade.sell_slippage", float, 0, 1, "trade.sell_slippage must be between 0 and 1"), ("trade.sell_slippage", float, 0, 1, "trade.sell_slippage must be between 0 and 1"),
( ("priority_fees.fixed_amount", int, 0, float("inf"), "priority_fees.fixed_amount must be a non-negative integer"),
"priority_fees.fixed_amount", ("priority_fees.extra_percentage", float, 0, 1, "priority_fees.extra_percentage must be between 0 and 1"),
int, ("priority_fees.hard_cap", int, 0, float("inf"), "priority_fees.hard_cap must be a non-negative integer"),
0, ("retries.max_attempts", int, 0, 100, "retries.max_attempts must be between 0 and 100"),
float("inf"), ("filters.max_token_age", (int, float), 0, float("inf"), "filters.max_token_age must be a non-negative number"),
"priority_fees.fixed_amount must be a non-negative integer",
),
(
"priority_fees.extra_percentage",
float,
0,
1,
"priority_fees.extra_percentage must be between 0 and 1",
),
(
"priority_fees.hard_cap",
int,
0,
float("inf"),
"priority_fees.hard_cap must be a non-negative integer",
),
(
"retries.max_attempts",
int,
0,
100,
"retries.max_attempts must be between 0 and 100",
),
(
"filters.max_token_age",
(int, float),
0,
float("inf"),
"filters.max_token_age must be a non-negative number",
),
] ]
# Valid values for enum-like fields (extended with platform support) # Valid values for enum-like fields
VALID_VALUES = { VALID_VALUES = {
"filters.listener_type": ["logs", "blocks", "geyser", "pumpportal"], "filters.listener_type": ["logs", "blocks", "geyser", "pumpportal"],
"cleanup.mode": ["disabled", "on_fail", "after_sell", "post_session"], "cleanup.mode": ["disabled", "on_fail", "after_sell", "post_session"],
"trade.exit_strategy": ["time_based", "tp_sl", "manual"], "trade.exit_strategy": ["time_based", "tp_sl", "manual"],
"platform": ["pump_fun", "lets_bonk"], # Platform validation "platform": ["pump_fun", "lets_bonk"],
}
# Platform-specific listener compatibility
PLATFORM_LISTENER_COMPATIBILITY = {
Platform.PUMP_FUN: ["logs", "blocks", "geyser", "pumpportal"],
Platform.LETS_BONK: ["logs", "blocks", "geyser"], # PumpPortal is pump.fun only
} }
def load_bot_config(path: str) -> dict: def load_bot_config(path: str) -> dict:
""" """Load and validate a bot configuration from a YAML file."""
Load and validate a bot configuration from a YAML file.
Extended to support platform selection with backward compatibility.
Args:
path: Path to the YAML configuration file (relative or absolute)
Returns:
Validated configuration dictionary
Raises:
FileNotFoundError: If the configuration file doesn't exist
ValueError: If the configuration is invalid
"""
with open(path) as f: with open(path) as f:
config = yaml.safe_load(f) config = yaml.safe_load(f)
@@ -107,7 +60,6 @@ def load_bot_config(path: str) -> dict:
if os.path.exists(env_path): if os.path.exists(env_path):
load_dotenv(env_path, override=True) load_dotenv(env_path, override=True)
else: else:
# If not found relative to config, try relative to current working directory
load_dotenv(env_file, override=True) load_dotenv(env_file, override=True)
resolve_env_vars(config) resolve_env_vars(config)
@@ -117,18 +69,11 @@ def load_bot_config(path: str) -> dict:
config["platform"] = "pump_fun" config["platform"] = "pump_fun"
validate_config(config) validate_config(config)
return config return config
def resolve_env_vars(config: dict) -> None: def resolve_env_vars(config: dict) -> None:
""" """Recursively resolve environment variables in the configuration."""
Recursively resolve environment variables in the configuration.
Args:
config: Configuration dictionary to process
"""
def resolve_env(value): def resolve_env(value):
if isinstance(value, str) and value.startswith("${") and value.endswith("}"): if isinstance(value, str) and value.startswith("${") and value.endswith("}"):
env_var = value[2:-1] env_var = value[2:-1]
@@ -149,19 +94,7 @@ def resolve_env_vars(config: dict) -> None:
def get_nested_value(config: dict, path: str) -> Any: def get_nested_value(config: dict, path: str) -> Any:
""" """Get a nested value from the configuration using dot notation."""
Get a nested value from the configuration using dot notation.
Args:
config: Configuration dictionary
path: Path to the value using dot notation (e.g., "trade.buy_amount")
Returns:
The value at the specified path
Raises:
ValueError: If the path doesn't exist in the configuration
"""
keys = path.split(".") keys = path.split(".")
value = config value = config
for key in keys: for key in keys:
@@ -172,16 +105,7 @@ def get_nested_value(config: dict, path: str) -> Any:
def validate_config(config: dict) -> None: def validate_config(config: dict) -> None:
""" """Validate the configuration against defined rules with platform support."""
Validate the configuration against defined rules.
Extended to include platform-specific validation.
Args:
config: Configuration dictionary to validate
Raises:
ValueError: If the configuration is invalid
"""
# Validate required fields # Validate required fields
for field in REQUIRED_FIELDS: for field in REQUIRED_FIELDS:
get_nested_value(config, field) get_nested_value(config, field)
@@ -198,10 +122,8 @@ def validate_config(config: dict) -> None:
raise ValueError(f"Range error: {error_msg}") raise ValueError(f"Range error: {error_msg}")
except ValueError as e: except ValueError as e:
# Re-raise if it's our own error
if str(e).startswith(("Type error:", "Range error:")): if str(e).startswith(("Type error:", "Range error:")):
raise raise
# Otherwise, the field might be missing
continue continue
# Validate enum-like fields # Validate enum-like fields
@@ -211,7 +133,6 @@ def validate_config(config: dict) -> None:
if value not in valid_values: if value not in valid_values:
raise ValueError(f"{path} must be one of {valid_values}") raise ValueError(f"{path} must be one of {valid_values}")
except ValueError: except ValueError:
# Skip if the field is missing
continue continue
# Cannot enable both dynamic and fixed priority fees # Cannot enable both dynamic and fixed priority fees
@@ -219,11 +140,8 @@ def validate_config(config: dict) -> None:
dynamic = get_nested_value(config, "priority_fees.enable_dynamic") dynamic = get_nested_value(config, "priority_fees.enable_dynamic")
fixed = get_nested_value(config, "priority_fees.enable_fixed") fixed = get_nested_value(config, "priority_fees.enable_fixed")
if dynamic and fixed: if dynamic and fixed:
raise ValueError( raise ValueError("Cannot enable both dynamic and fixed priority fees simultaneously")
"Cannot enable both dynamic and fixed priority fees simultaneously"
)
except ValueError: except ValueError:
# Skip if one of the fields is missing
pass pass
# Platform-specific validation # Platform-specific validation
@@ -238,16 +156,30 @@ def validate_config(config: dict) -> None:
def validate_platform_config(config: dict, platform: Platform) -> None: def validate_platform_config(config: dict, platform: Platform) -> None:
""" """Validate platform-specific configuration requirements."""
Validate platform-specific configuration requirements. # Check if platform is supported
try:
from platforms import platform_factory
if not platform_factory.registry.is_platform_supported(platform):
raise ValueError(f"Platform {platform.value} is not supported. Available platforms: {[p.value for p in platform_factory.get_supported_platforms()]}")
except ImportError:
# If platform factory not available, just validate enum
pass
Args: # Validate listener compatibility with platform
config: Configuration dictionary try:
platform: Platform enum value listener_type = get_nested_value(config, "filters.listener_type")
compatible_listeners = PLATFORM_LISTENER_COMPATIBILITY.get(platform, [])
if listener_type not in compatible_listeners:
raise ValueError(
f"Listener type '{listener_type}' is not compatible with platform '{platform.value}'. "
f"Compatible listeners: {compatible_listeners}"
)
except ValueError:
pass
Raises: # Platform-specific configuration validation
ValueError: If platform-specific config is invalid
"""
if platform == Platform.PUMP_FUN: if platform == Platform.PUMP_FUN:
# pump.fun doesn't require additional config beyond base requirements # pump.fun doesn't require additional config beyond base requirements
pass pass
@@ -259,18 +191,7 @@ def validate_platform_config(config: dict, platform: Platform) -> None:
def get_platform_from_config(config: dict) -> Platform: def get_platform_from_config(config: dict) -> Platform:
""" """Extract platform enum from configuration."""
Extract platform enum from configuration.
Args:
config: Configuration dictionary
Returns:
Platform enum value
Raises:
ValueError: If platform is invalid
"""
platform_str = config.get("platform", "pump_fun") platform_str = config.get("platform", "pump_fun")
try: try:
return Platform(platform_str) return Platform(platform_str)
@@ -278,46 +199,156 @@ def get_platform_from_config(config: dict) -> Platform:
raise ValueError(f"Invalid platform '{platform_str}'. Must be one of: {[p.value for p in Platform]}") raise ValueError(f"Invalid platform '{platform_str}'. Must be one of: {[p.value for p in Platform]}")
def print_config_summary(config: dict) -> None: def validate_platform_listener_combination(platform: Platform, listener_type: str) -> bool:
""" """Check if a platform and listener type are compatible.
Print a summary of the loaded configuration.
Extended to show platform information.
Args: Args:
config: Configuration dictionary platform: Platform enum
listener_type: Listener type string
Returns:
True if combination is valid
""" """
compatible_listeners = PLATFORM_LISTENER_COMPATIBILITY.get(platform, [])
return listener_type in compatible_listeners
def get_supported_listeners_for_platform(platform: Platform) -> list[str]:
"""Get list of supported listener types for a platform.
Args:
platform: Platform enum
Returns:
List of supported listener types
"""
return PLATFORM_LISTENER_COMPATIBILITY.get(platform, [])
def get_platform_specific_required_config(platform: Platform) -> list[str]:
"""Get platform-specific required configuration paths.
Args:
platform: Platform enum
Returns:
List of additional required config paths for the platform
"""
if platform == Platform.PUMP_FUN:
return [] # No additional requirements
elif platform == Platform.LETS_BONK:
return [] # No additional requirements yet
else:
return []
def print_config_summary(config: dict) -> None:
"""Print a summary of the loaded configuration with platform info."""
platform_str = config.get("platform", "pump_fun") platform_str = config.get("platform", "pump_fun")
print(f"Bot name: {config.get('name', 'unnamed')}") print(f"Bot name: {config.get('name', 'unnamed')}")
print(f"Platform: {platform_str}") print(f"Platform: {platform_str}")
print( print(f"Listener type: {config.get('filters', {}).get('listener_type', 'not configured')}")
f"Listener type: {config.get('filters', {}).get('listener_type', 'not configured')}"
) # Validate platform-listener combination
try:
platform = Platform(platform_str)
listener_type = config.get('filters', {}).get('listener_type')
if listener_type and not validate_platform_listener_combination(platform, listener_type):
print(f"WARNING: Listener '{listener_type}' may not be compatible with platform '{platform_str}'")
except ValueError:
print(f"WARNING: Invalid platform '{platform_str}'")
trade = config.get("trade", {}) trade = config.get("trade", {})
print("Trade settings:") print("Trade settings:")
print(f" - Buy amount: {trade.get('buy_amount', 'not configured')} SOL") print(f" - Buy amount: {trade.get('buy_amount', 'not configured')} SOL")
print(f" - Buy slippage: {trade.get('buy_slippage', 'not configured') * 100}%") print(f" - Buy slippage: {trade.get('buy_slippage', 'not configured') * 100}%")
print( print(f" - Extreme fast mode: {'enabled' if trade.get('extreme_fast_mode') else 'disabled'}")
f" - Extreme fast mode: {'enabled' if trade.get('extreme_fast_mode') else 'disabled'}"
)
fees = config.get("priority_fees", {}) fees = config.get("priority_fees", {})
print("Priority fees:") print("Priority fees:")
if fees.get("enable_dynamic"): if fees.get("enable_dynamic"):
print(" - Dynamic fees enabled") print(" - Dynamic fees enabled")
elif fees.get("enable_fixed"): elif fees.get("enable_fixed"):
print( print(f" - Fixed fee: {fees.get('fixed_amount', 'not configured')} microlamports")
f" - Fixed fee: {fees.get('fixed_amount', 'not configured')} microlamports"
)
print("Configuration loaded successfully!") print("Configuration loaded successfully!")
if __name__ == "__main__": def validate_all_platform_configs(config_dir: str = "bots") -> dict[str, Any]:
# Example usage with platform configuration """Validate all bot configurations in a directory.
config = load_bot_config("bots/bot-sniper.yaml")
print_config_summary(config)
platform = get_platform_from_config(config) Args:
print(f"Detected platform: {platform}") config_dir: Directory containing bot config files
Returns:
Dictionary with validation results
"""
import glob
import os
results = {
"valid_configs": [],
"invalid_configs": [],
"platform_distribution": {},
"listener_distribution": {},
}
config_files = glob.glob(os.path.join(config_dir, "*.yaml"))
for config_file in config_files:
try:
config = load_bot_config(config_file)
platform = get_platform_from_config(config)
listener_type = config.get('filters', {}).get('listener_type', 'unknown')
results["valid_configs"].append({
"file": config_file,
"name": config.get("name"),
"platform": platform.value,
"listener": listener_type,
"enabled": config.get("enabled", True)
})
# Track distributions
platform_key = platform.value
results["platform_distribution"][platform_key] = results["platform_distribution"].get(platform_key, 0) + 1
results["listener_distribution"][listener_type] = results["listener_distribution"].get(listener_type, 0) + 1
except Exception as e:
results["invalid_configs"].append({
"file": config_file,
"error": str(e)
})
return results
if __name__ == "__main__":
# Example usage with platform configuration validation
import sys
if len(sys.argv) > 1:
config_path = sys.argv[1]
try:
config = load_bot_config(config_path)
print_config_summary(config)
platform = get_platform_from_config(config)
print(f"Detected platform: {platform}")
print(f"Supported listeners for this platform: {get_supported_listeners_for_platform(platform)}")
except Exception as e:
print(f"Configuration error: {e}")
else:
# Validate all configs in bots directory
results = validate_all_platform_configs()
print("Configuration validation results:")
print(f"Valid configs: {len(results['valid_configs'])}")
print(f"Invalid configs: {len(results['invalid_configs'])}")
print(f"Platform distribution: {results['platform_distribution']}")
print(f"Listener distribution: {results['listener_distribution']}")
if results['invalid_configs']:
print("\nInvalid configurations:")
for invalid in results['invalid_configs']:
print(f" {invalid['file']}: {invalid['error']}")
+165
View File
@@ -0,0 +1,165 @@
"""
Pump.fun bonding curve manager for price calculations and curve state management.
"""
import struct
from dataclasses import dataclass
from solders.pubkey import Pubkey
from core.client import SolanaClient
from core.pubkeys import LAMPORTS_PER_SOL, TOKEN_DECIMALS
from utils.logger import get_logger
logger = get_logger(__name__)
# Bonding curve discriminator
CURVE_DISCRIMINATOR = bytes([23, 203, 71, 8, 209, 70, 227, 3])
@dataclass
class BondingCurveState:
"""Represents the state of a pump.fun bonding curve."""
virtual_token_reserves: int
virtual_sol_reserves: int
real_token_reserves: int
real_sol_reserves: int
token_total_supply: int
complete: bool
creator: Pubkey
@property
def token_reserves(self) -> float:
"""Token reserves in decimal form."""
return self.virtual_token_reserves / 10**TOKEN_DECIMALS
@property
def sol_reserves(self) -> float:
"""SOL reserves in decimal form."""
return self.virtual_sol_reserves / LAMPORTS_PER_SOL
def calculate_price(self) -> float:
"""Calculate current token price in SOL."""
if self.virtual_token_reserves <= 0:
return 0.0
# Price = sol_reserves / token_reserves
price_lamports = self.virtual_sol_reserves / self.virtual_token_reserves
return price_lamports * (10**TOKEN_DECIMALS) / LAMPORTS_PER_SOL
class BondingCurveManager:
"""Manages pump.fun bonding curve operations."""
def __init__(self, client: SolanaClient):
"""Initialize bonding curve manager.
Args:
client: Solana RPC client
"""
self.client = client
async def get_curve_state(self, curve_address: Pubkey) -> BondingCurveState:
"""Get the current state of a bonding curve.
Args:
curve_address: Address of the bonding curve
Returns:
BondingCurveState with current curve data
Raises:
ValueError: If curve data is invalid or inaccessible
"""
try:
account = await self.client.get_account_info(curve_address)
if not account.data:
raise ValueError(f"No data in bonding curve account {curve_address}")
curve_state = self._decode_curve_state(account.data)
return curve_state
except Exception as e:
logger.error(f"Failed to get curve state: {e!s}")
raise ValueError(f"Invalid bonding curve state: {e!s}")
async def calculate_price(self, curve_address: Pubkey) -> float:
"""Calculate current token price from bonding curve.
Args:
curve_address: Address of the bonding curve
Returns:
Current token price in SOL
"""
curve_state = await self.get_curve_state(curve_address)
return curve_state.calculate_price()
def _decode_curve_state(self, data: bytes) -> BondingCurveState:
"""Decode bonding curve state from raw account data.
Args:
data: Raw account data
Returns:
Decoded BondingCurveState
Raises:
ValueError: If data format is invalid
"""
if len(data) < 8:
raise ValueError("Curve data too short")
# Check discriminator
if not data.startswith(CURVE_DISCRIMINATOR):
raise ValueError("Invalid curve discriminator")
offset = 8
try:
# Decode based on pump.fun BondingCurve structure:
# - virtual_token_reserves: u64 (8 bytes)
# - virtual_sol_reserves: u64 (8 bytes)
# - real_token_reserves: u64 (8 bytes)
# - real_sol_reserves: u64 (8 bytes)
# - token_total_supply: u64 (8 bytes)
# - complete: bool (1 byte)
# - padding: 7 bytes
# - creator: Pubkey (32 bytes)
virtual_token_reserves = struct.unpack_from("<Q", data, offset)[0]
offset += 8
virtual_sol_reserves = struct.unpack_from("<Q", data, offset)[0]
offset += 8
real_token_reserves = struct.unpack_from("<Q", data, offset)[0]
offset += 8
real_sol_reserves = struct.unpack_from("<Q", data, offset)[0]
offset += 8
token_total_supply = struct.unpack_from("<Q", data, offset)[0]
offset += 8
complete = bool(struct.unpack_from("<B", data, offset)[0])
offset += 1
# Skip padding
offset += 7
creator = Pubkey.from_bytes(data[offset:offset + 32])
return BondingCurveState(
virtual_token_reserves=virtual_token_reserves,
virtual_sol_reserves=virtual_sol_reserves,
real_token_reserves=real_token_reserves,
real_sol_reserves=real_sol_reserves,
token_total_supply=token_total_supply,
complete=complete,
creator=creator,
)
except Exception as e:
raise ValueError(f"Failed to decode curve state: {e}")
+19 -33
View File
@@ -6,12 +6,10 @@ import asyncio
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
import grpc import grpc
from solders.pubkey import Pubkey
from geyser.generated import geyser_pb2, geyser_pb2_grpc from geyser.generated import geyser_pb2, geyser_pb2_grpc
from interfaces.core import Platform, TokenInfo from interfaces.core import Platform, TokenInfo
from monitoring.base_listener import BaseTokenListener from monitoring.base_listener import BaseTokenListener
from platforms import get_platform_implementations
from utils.logger import get_logger from utils.logger import get_logger
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -27,14 +25,7 @@ class UniversalGeyserListener(BaseTokenListener):
geyser_auth_type: str, geyser_auth_type: str,
platforms: list[Platform] | None = None, platforms: list[Platform] | None = None,
): ):
"""Initialize universal Geyser listener. """Initialize universal Geyser listener."""
Args:
geyser_endpoint: Geyser gRPC endpoint URL
geyser_api_token: API token for authentication
geyser_auth_type: authentication type ('x-token' or 'basic')
platforms: List of platforms to monitor (if None, monitor all supported platforms)
"""
super().__init__() super().__init__()
self.geyser_endpoint = geyser_endpoint self.geyser_endpoint = geyser_endpoint
self.geyser_api_token = geyser_api_token self.geyser_api_token = geyser_api_token
@@ -51,7 +42,6 @@ class UniversalGeyserListener(BaseTokenListener):
from platforms import platform_factory from platforms import platform_factory
if platforms is None: if platforms is None:
# Monitor all supported platforms
self.platforms = platform_factory.get_supported_platforms() self.platforms = platform_factory.get_supported_platforms()
else: else:
self.platforms = platforms self.platforms = platforms
@@ -60,13 +50,13 @@ class UniversalGeyserListener(BaseTokenListener):
self.platform_parsers = {} self.platform_parsers = {}
self.platform_program_ids = set() self.platform_program_ids = set()
# Create a temporary client for getting parsers
from core.client import SolanaClient
temp_client = SolanaClient("http://temp")
for platform in self.platforms: for platform in self.platforms:
try: try:
# We'll need a dummy client for getting the parser - this is a design issue we should fix implementations = platform_factory.create_for_platform(platform, temp_client)
from core.client import SolanaClient
dummy_client = SolanaClient("http://localhost") # Won't be used for parsing
implementations = get_platform_implementations(platform, dummy_client)
parser = implementations.event_parser parser = implementations.event_parser
self.platform_parsers[platform] = parser self.platform_parsers[platform] = parser
self.platform_program_ids.add(parser.get_program_id()) self.platform_program_ids.add(parser.get_program_id())
@@ -78,6 +68,7 @@ class UniversalGeyserListener(BaseTokenListener):
async def _create_geyser_connection(self): async def _create_geyser_connection(self):
"""Establish a secure connection to the Geyser endpoint.""" """Establish a secure connection to the Geyser endpoint."""
if self.auth_type == "x-token": if self.auth_type == "x-token":
auth = grpc.metadata_call_credentials( auth = grpc.metadata_call_credentials(
lambda _, callback: callback( lambda _, callback: callback(
@@ -92,10 +83,12 @@ class UniversalGeyserListener(BaseTokenListener):
) )
creds = grpc.composite_channel_credentials(grpc.ssl_channel_credentials(), auth) creds = grpc.composite_channel_credentials(grpc.ssl_channel_credentials(), auth)
channel = grpc.aio.secure_channel(self.geyser_endpoint, creds) channel = grpc.aio.secure_channel(self.geyser_endpoint, creds)
return geyser_pb2_grpc.GeyserStub(channel), channel return geyser_pb2_grpc.GeyserStub(channel), channel
def _create_subscription_request(self): def _create_subscription_request(self):
"""Create a subscription request for all monitored platforms.""" """Create a subscription request for all monitored platforms."""
request = geyser_pb2.SubscribeRequest() request = geyser_pb2.SubscribeRequest()
# Add all platform program IDs to the filter # Add all platform program IDs to the filter
@@ -113,13 +106,7 @@ class UniversalGeyserListener(BaseTokenListener):
match_string: str | None = None, match_string: str | None = None,
creator_address: str | None = None, creator_address: str | None = None,
) -> None: ) -> None:
"""Listen for new token creations using Geyser subscription. """Listen for new token creations using Geyser subscription."""
Args:
token_callback: Callback function for new tokens
match_string: Optional string to match in token name/symbol
creator_address: Optional creator address to filter by
"""
if not self.platform_parsers: if not self.platform_parsers:
logger.error("No platform parsers available. Cannot listen for tokens.") logger.error("No platform parsers available. Cannot listen for tokens.")
return return
@@ -161,8 +148,12 @@ class UniversalGeyserListener(BaseTokenListener):
await token_callback(token_info) await token_callback(token_info)
except grpc.aio.AioRpcError as e: except Exception as e:
logger.error(f"gRPC error: {e.details()}") import grpc
if isinstance(e, grpc.aio.AioRpcError):
logger.error(f"gRPC error: {e.details()}")
else:
logger.error(f"Geyser error: {e}")
await asyncio.sleep(5) await asyncio.sleep(5)
finally: finally:
@@ -174,14 +165,7 @@ class UniversalGeyserListener(BaseTokenListener):
await asyncio.sleep(10) await asyncio.sleep(10)
async def _process_update(self, update) -> TokenInfo | None: async def _process_update(self, update) -> TokenInfo | None:
"""Process a Geyser update and extract token creation info. """Process a Geyser update and extract token creation info."""
Args:
update: Geyser update from the subscription
Returns:
TokenInfo if a token creation is found, None otherwise
"""
try: try:
if not update.HasField("transaction"): if not update.HasField("transaction"):
return None return None
@@ -191,6 +175,8 @@ class UniversalGeyserListener(BaseTokenListener):
if msg is None: if msg is None:
return None return None
from solders.pubkey import Pubkey
for ix in msg.instructions: for ix in msg.instructions:
# Check which platform this instruction belongs to # Check which platform this instruction belongs to
program_idx = ix.program_id_index program_idx = ix.program_id_index
+9 -23
View File
@@ -1,7 +1,6 @@
""" """
Universal logs listener that works with any platform through the interface system. Universal logs listener that works with any platform through the interface system.
""" """
import asyncio import asyncio
import json import json
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
@@ -10,7 +9,6 @@ import websockets
from interfaces.core import Platform, TokenInfo from interfaces.core import Platform, TokenInfo
from monitoring.base_listener import BaseTokenListener from monitoring.base_listener import BaseTokenListener
from platforms import get_platform_implementations
from utils.logger import get_logger from utils.logger import get_logger
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -47,13 +45,13 @@ class UniversalLogsListener(BaseTokenListener):
self.platform_parsers = {} self.platform_parsers = {}
self.platform_program_ids = [] self.platform_program_ids = []
# Create a temporary client for getting parsers (stateless parsers don't use it)
from core.client import SolanaClient
temp_client = SolanaClient("http://temp")
for platform in self.platforms: for platform in self.platforms:
try: try:
# We'll need a dummy client for getting the parser implementations = platform_factory.create_for_platform(platform, temp_client)
from core.client import SolanaClient
dummy_client = SolanaClient("http://localhost") # Won't be used for parsing
implementations = get_platform_implementations(platform, dummy_client)
parser = implementations.event_parser parser = implementations.event_parser
self.platform_parsers[platform] = parser self.platform_parsers[platform] = parser
self.platform_program_ids.append(str(parser.get_program_id())) self.platform_program_ids.append(str(parser.get_program_id()))
@@ -130,11 +128,11 @@ class UniversalLogsListener(BaseTokenListener):
websocket: Active WebSocket connection websocket: Active WebSocket connection
""" """
# Subscribe to logs for all monitored platforms # Subscribe to logs for all monitored platforms
for program_id in self.platform_program_ids: for i, program_id in enumerate(self.platform_program_ids):
subscription_message = json.dumps( subscription_message = json.dumps(
{ {
"jsonrpc": "2.0", "jsonrpc": "2.0",
"id": len(self.platform_program_ids), # Use different IDs "id": i + 1,
"method": "logsSubscribe", "method": "logsSubscribe",
"params": [ "params": [
{"mentions": [program_id]}, {"mentions": [program_id]},
@@ -155,11 +153,7 @@ class UniversalLogsListener(BaseTokenListener):
logger.warning(f"Unexpected subscription response: {response}") logger.warning(f"Unexpected subscription response: {response}")
async def _ping_loop(self, websocket) -> None: async def _ping_loop(self, websocket) -> None:
"""Keep connection alive with pings. """Keep connection alive with pings."""
Args:
websocket: Active WebSocket connection
"""
try: try:
while True: while True:
await asyncio.sleep(self.ping_interval) await asyncio.sleep(self.ping_interval)
@@ -168,7 +162,6 @@ class UniversalLogsListener(BaseTokenListener):
await asyncio.wait_for(pong_waiter, timeout=10) await asyncio.wait_for(pong_waiter, timeout=10)
except TimeoutError: except TimeoutError:
logger.warning("Ping timeout - server not responding") logger.warning("Ping timeout - server not responding")
# Force reconnection
await websocket.close() await websocket.close()
return return
except asyncio.CancelledError: except asyncio.CancelledError:
@@ -177,14 +170,7 @@ class UniversalLogsListener(BaseTokenListener):
logger.error(f"Ping error: {e!s}") logger.error(f"Ping error: {e!s}")
async def _wait_for_token_creation(self, websocket) -> TokenInfo | None: async def _wait_for_token_creation(self, websocket) -> TokenInfo | None:
"""Wait for token creation events from any platform. """Wait for token creation events from any platform."""
Args:
websocket: Active WebSocket connection
Returns:
TokenInfo if a token creation is found, None otherwise
"""
try: try:
response = await asyncio.wait_for(websocket.recv(), timeout=30) response = await asyncio.wait_for(websocket.recv(), timeout=30)
data = json.loads(response) data = json.loads(response)
+15 -9
View File
@@ -32,7 +32,7 @@ class PlatformRegistry:
def __init__(self): def __init__(self):
self._implementations: dict[Platform, dict[str, type]] = {} self._implementations: dict[Platform, dict[str, type]] = {}
self._instances: dict[Platform, PlatformImplementations] = {} self._instances: dict[tuple[Platform, str], PlatformImplementations] = {}
def register_platform( def register_platform(
self, self,
@@ -80,13 +80,16 @@ class PlatformRegistry:
if platform not in self._implementations: if platform not in self._implementations:
raise ValueError(f"Platform {platform} is not registered") raise ValueError(f"Platform {platform} is not registered")
# Check if we already have instances for this platform # Use client address as cache key to allow multiple clients
if platform in self._instances: cache_key = (platform, str(client.rpc_endpoint))
return self._instances[platform]
# Check if we already have instances for this platform + client combo
if cache_key in self._instances:
return self._instances[cache_key]
impl_classes = self._implementations[platform] impl_classes = self._implementations[platform]
# Create instances # Create instances - only curve_manager needs client
address_provider = impl_classes['address_provider']() address_provider = impl_classes['address_provider']()
instruction_builder = impl_classes['instruction_builder']() instruction_builder = impl_classes['instruction_builder']()
curve_manager = impl_classes['curve_manager'](client) curve_manager = impl_classes['curve_manager'](client)
@@ -100,20 +103,22 @@ class PlatformRegistry:
) )
# Cache the instances # Cache the instances
self._instances[platform] = implementations self._instances[cache_key] = implementations
return implementations return implementations
def get_platform_implementations(self, platform: Platform) -> PlatformImplementations | None: def get_platform_implementations(self, platform: Platform, client_endpoint: str) -> PlatformImplementations | None:
"""Get cached platform implementations. """Get cached platform implementations.
Args: Args:
platform: Platform to get implementations for platform: Platform to get implementations for
client_endpoint: Client endpoint for cache lookup
Returns: Returns:
PlatformImplementations if available, None otherwise PlatformImplementations if available, None otherwise
""" """
return self._instances.get(platform) cache_key = (platform, client_endpoint)
return self._instances.get(cache_key)
def get_supported_platforms(self) -> list[Platform]: def get_supported_platforms(self) -> list[Platform]:
"""Get list of supported platforms. """Get list of supported platforms.
@@ -144,7 +149,7 @@ class PlatformFactory:
def _setup_default_platforms(self) -> None: def _setup_default_platforms(self) -> None:
"""Setup default platform registrations.""" """Setup default platform registrations."""
# Import platform implementations dynamically to avoid circular imports # Import and register pump.fun platform
try: try:
from platforms.pumpfun import ( from platforms.pumpfun import (
PumpFunAddressProvider, PumpFunAddressProvider,
@@ -164,6 +169,7 @@ class PlatformFactory:
except ImportError as e: except ImportError as e:
print(f"Warning: Could not register pump.fun platform: {e}") print(f"Warning: Could not register pump.fun platform: {e}")
# Import and register LetsBonk platform
try: try:
from platforms.letsbonk import ( from platforms.letsbonk import (
LetsBonkAddressProvider, LetsBonkAddressProvider,
+105 -30
View File
@@ -5,11 +5,13 @@ This module parses LetsBonk-specific token creation events from various sources
by implementing the EventParser interface. by implementing the EventParser interface.
""" """
import base64
import struct import struct
from time import monotonic from time import monotonic
from typing import Any, Final from typing import Any, Final
from solders.pubkey import Pubkey from solders.pubkey import Pubkey
from solders.transaction import VersionedTransaction
from interfaces.core import EventParser, Platform, TokenInfo from interfaces.core import EventParser, Platform, TokenInfo
from platforms.letsbonk.address_provider import LetsBonkAddressProvider from platforms.letsbonk.address_provider import LetsBonkAddressProvider
@@ -184,6 +186,108 @@ class LetsBonkEventParser(EventParser):
""" """
return [self.INITIALIZE_DISCRIMINATOR] return [self.INITIALIZE_DISCRIMINATOR]
def parse_token_creation_from_block(self, block_data: dict) -> TokenInfo | None:
"""Parse token creation from block data (for block listener).
Args:
block_data: Block data from WebSocket
Returns:
TokenInfo if token creation found, None otherwise
"""
try:
if "transactions" not in block_data:
return None
for tx in block_data["transactions"]:
if not isinstance(tx, dict) or "transaction" not in tx:
continue
# Decode base64 transaction data if needed
tx_data = tx["transaction"]
if isinstance(tx_data, list) and len(tx_data) > 0:
try:
tx_data_encoded = tx_data[0]
tx_data_decoded = base64.b64decode(tx_data_encoded)
transaction = VersionedTransaction.from_bytes(tx_data_decoded)
for ix in transaction.message.instructions:
program_id = transaction.message.account_keys[ix.program_id_index]
# Check if instruction is from LetsBonk program
if str(program_id) != str(self.get_program_id()):
continue
ix_data = bytes(ix.data)
# Check for initialize discriminator
if len(ix_data) >= 8:
discriminator = struct.unpack("<Q", ix_data[:8])[0]
if discriminator == self.INITIALIZE_DISCRIMINATOR_INT:
# Token creation should have substantial data and many accounts
if len(ix_data) <= 8 or len(ix.accounts) < 10:
continue
# Parse the instruction
token_info = self.parse_token_creation_from_instruction(
ix_data, ix.accounts, transaction.message.account_keys
)
if token_info:
return token_info
except Exception:
continue
# Handle already decoded transaction data
elif isinstance(tx_data, dict) and "message" in tx_data:
try:
message = tx_data["message"]
if "instructions" not in message or "accountKeys" not in message:
continue
for ix in message["instructions"]:
if "programIdIndex" not in ix or "accounts" not in ix or "data" not in ix:
continue
program_idx = ix["programIdIndex"]
if program_idx >= len(message["accountKeys"]):
continue
program_id_str = message["accountKeys"][program_idx]
if program_id_str != str(self.get_program_id()):
continue
# Decode instruction data
ix_data = base64.b64decode(ix["data"])
if len(ix_data) >= 8:
discriminator = struct.unpack("<Q", ix_data[:8])[0]
if discriminator == self.INITIALIZE_DISCRIMINATOR_INT:
if len(ix_data) <= 8 or len(ix["accounts"]) < 10:
continue
# Convert account keys to bytes for parsing
account_keys_bytes = [
Pubkey.from_string(key).to_bytes()
for key in message["accountKeys"]
]
token_info = self.parse_token_creation_from_instruction(
ix_data, ix["accounts"], account_keys_bytes
)
if token_info:
return token_info
except Exception:
continue
return None
except Exception:
return None
def _parse_initialize_instruction_data(self, data: bytes) -> dict | None: def _parse_initialize_instruction_data(self, data: bytes) -> dict | None:
"""Parse the initialize instruction data from LetsBonk. """Parse the initialize instruction data from LetsBonk.
@@ -240,33 +344,4 @@ class LetsBonkEventParser(EventParser):
return parsed_data return parsed_data
except Exception: except Exception:
return None return None
def parse_token_creation_from_block(self, block_data: dict) -> list[TokenInfo]:
"""Parse token creations from block data (for block listener).
Args:
block_data: Block data from WebSocket
Returns:
List of TokenInfo for any token creations found
"""
tokens = []
try:
if "transactions" not in block_data:
return tokens
for tx in block_data["transactions"]:
if not isinstance(tx, dict) or "transaction" not in tx:
continue
# Process transaction (implementation would be similar to pump.fun)
# This is a simplified version - full implementation would decode
# the transaction and check for LetsBonk initialize instructions
pass
return tokens
except Exception:
return tokens
+29 -74
View File
@@ -1,8 +1,6 @@
""" """
Platform-aware trader implementations that use the interface system. Platform-aware trader implementations that use the interface system.
Final cleanup removing all platform-specific hardcoding.
This module provides new trader classes that work with any platform
through the interface system, while maintaining compatibility with existing code.
""" """
from solders.pubkey import Pubkey from solders.pubkey import Pubkey
@@ -33,18 +31,7 @@ class PlatformAwareBuyer(Trader):
extreme_fast_token_amount: int = 0, extreme_fast_token_amount: int = 0,
extreme_fast_mode: bool = False, extreme_fast_mode: bool = False,
): ):
"""Initialize platform-aware token buyer. """Initialize platform-aware token buyer."""
Args:
client: Solana client for RPC calls
wallet: Wallet for signing transactions
priority_fee_manager: Priority fee manager
amount: Amount of SOL to spend
slippage: Slippage tolerance (0.01 = 1%)
max_retries: Maximum number of retry attempts
extreme_fast_token_amount: Amount of token to buy if extreme fast mode is enabled
extreme_fast_mode: If enabled, avoid fetching pool state for price estimation
"""
self.client = client self.client = client
self.wallet = wallet self.wallet = wallet
self.priority_fee_manager = priority_fee_manager self.priority_fee_manager = priority_fee_manager
@@ -55,14 +42,7 @@ class PlatformAwareBuyer(Trader):
self.extreme_fast_token_amount = extreme_fast_token_amount self.extreme_fast_token_amount = extreme_fast_token_amount
async def execute(self, token_info: TokenInfo, *args, **kwargs) -> TradeResult: async def execute(self, token_info: TokenInfo, *args, **kwargs) -> TradeResult:
"""Execute buy operation using platform-specific implementations. """Execute buy operation using platform-specific implementations."""
Args:
token_info: Enhanced token information with platform
Returns:
TradeResult with buy outcome
"""
try: try:
# Get platform-specific implementations # Get platform-specific implementations
implementations = get_platform_implementations(token_info.platform, self.client) implementations = get_platform_implementations(token_info.platform, self.client)
@@ -78,7 +58,7 @@ class PlatformAwareBuyer(Trader):
token_amount = self.extreme_fast_token_amount token_amount = self.extreme_fast_token_amount
token_price_sol = self.amount / token_amount if token_amount > 0 else 0 token_price_sol = self.amount / token_amount if token_amount > 0 else 0
else: else:
# Get pool address based on platform # Get pool address based on platform using platform-agnostic method
pool_address = self._get_pool_address(token_info, address_provider) pool_address = self._get_pool_address(token_info, address_provider)
# Regular behavior with RPC call # Regular behavior with RPC call
@@ -92,7 +72,7 @@ class PlatformAwareBuyer(Trader):
# Calculate maximum SOL to spend with slippage # Calculate maximum SOL to spend with slippage
max_amount_lamports = int(amount_lamports * (1 + self.slippage)) max_amount_lamports = int(amount_lamports * (1 + self.slippage))
# Build buy instructions # Build buy instructions using platform-specific builder
instructions = await instruction_builder.build_buy_instruction( instructions = await instruction_builder.build_buy_instruction(
token_info, token_info,
self.wallet.pubkey, self.wallet.pubkey,
@@ -107,7 +87,7 @@ class PlatformAwareBuyer(Trader):
) )
logger.info( logger.info(
f"Buying {token_amount:.6f} tokens at {token_price_sol:.8f} SOL per token" f"Buying {token_amount:.6f} tokens at {token_price_sol:.8f} SOL per token on {token_info.platform.value}"
) )
logger.info( logger.info(
f"Total cost: {self.amount:.6f} SOL (max: {max_amount_lamports / LAMPORTS_PER_SOL:.6f} SOL)" f"Total cost: {self.amount:.6f} SOL (max: {max_amount_lamports / LAMPORTS_PER_SOL:.6f} SOL)"
@@ -151,22 +131,17 @@ class PlatformAwareBuyer(Trader):
) )
def _get_pool_address(self, token_info: TokenInfo, address_provider: AddressProvider) -> Pubkey: def _get_pool_address(self, token_info: TokenInfo, address_provider: AddressProvider) -> Pubkey:
"""Get the pool/curve address for price calculations. """Get the pool/curve address for price calculations using platform-agnostic method."""
# Try to get the address from token_info first, then derive if needed
Args:
token_info: Token information
address_provider: Platform address provider
Returns:
Pool/curve address
"""
if token_info.platform == Platform.PUMP_FUN: if token_info.platform == Platform.PUMP_FUN:
return token_info.bonding_curve or address_provider.derive_pool_address(token_info.mint) if hasattr(token_info, 'bonding_curve') and token_info.bonding_curve:
return token_info.bonding_curve
elif token_info.platform == Platform.LETS_BONK: elif token_info.platform == Platform.LETS_BONK:
return token_info.pool_state or address_provider.derive_pool_address(token_info.mint) if hasattr(token_info, 'pool_state') and token_info.pool_state:
else: return token_info.pool_state
# Fallback to deriving the address
return address_provider.derive_pool_address(token_info.mint) # Fallback to deriving the address using platform provider
return address_provider.derive_pool_address(token_info.mint)
class PlatformAwareSeller(Trader): class PlatformAwareSeller(Trader):
@@ -180,15 +155,7 @@ class PlatformAwareSeller(Trader):
slippage: float = 0.25, slippage: float = 0.25,
max_retries: int = 5, max_retries: int = 5,
): ):
"""Initialize platform-aware token seller. """Initialize platform-aware token seller."""
Args:
client: Solana client for RPC calls
wallet: Wallet for signing transactions
priority_fee_manager: Priority fee manager
slippage: Slippage tolerance (0.25 = 25%)
max_retries: Maximum number of retry attempts
"""
self.client = client self.client = client
self.wallet = wallet self.wallet = wallet
self.priority_fee_manager = priority_fee_manager self.priority_fee_manager = priority_fee_manager
@@ -196,14 +163,7 @@ class PlatformAwareSeller(Trader):
self.max_retries = max_retries self.max_retries = max_retries
async def execute(self, token_info: TokenInfo, *args, **kwargs) -> TradeResult: async def execute(self, token_info: TokenInfo, *args, **kwargs) -> TradeResult:
"""Execute sell operation using platform-specific implementations. """Execute sell operation using platform-specific implementations."""
Args:
token_info: Enhanced token information with platform
Returns:
TradeResult with sell outcome
"""
try: try:
# Get platform-specific implementations # Get platform-specific implementations
implementations = get_platform_implementations(token_info.platform, self.client) implementations = get_platform_implementations(token_info.platform, self.client)
@@ -229,7 +189,7 @@ class PlatformAwareSeller(Trader):
error_message="No tokens to sell" error_message="No tokens to sell"
) )
# Get pool address and current price # Get pool address and current price using platform-agnostic method
pool_address = self._get_pool_address(token_info, address_provider) pool_address = self._get_pool_address(token_info, address_provider)
token_price_sol = await curve_manager.calculate_price(pool_address) token_price_sol = await curve_manager.calculate_price(pool_address)
@@ -239,13 +199,13 @@ class PlatformAwareSeller(Trader):
expected_sol_output = float(token_balance_decimal) * float(token_price_sol) expected_sol_output = float(token_balance_decimal) * float(token_price_sol)
min_sol_output = int((expected_sol_output * (1 - self.slippage)) * LAMPORTS_PER_SOL) min_sol_output = int((expected_sol_output * (1 - self.slippage)) * LAMPORTS_PER_SOL)
logger.info(f"Selling {token_balance_decimal} tokens") logger.info(f"Selling {token_balance_decimal} tokens on {token_info.platform.value}")
logger.info(f"Expected SOL output: {expected_sol_output:.8f} SOL") logger.info(f"Expected SOL output: {expected_sol_output:.8f} SOL")
logger.info( logger.info(
f"Minimum SOL output (with {self.slippage * 100}% slippage): {min_sol_output / LAMPORTS_PER_SOL:.8f} SOL" f"Minimum SOL output (with {self.slippage * 100}% slippage): {min_sol_output / LAMPORTS_PER_SOL:.8f} SOL"
) )
# Build sell instructions # Build sell instructions using platform-specific builder
instructions = await instruction_builder.build_sell_instruction( instructions = await instruction_builder.build_sell_instruction(
token_info, token_info,
self.wallet.pubkey, self.wallet.pubkey,
@@ -297,19 +257,14 @@ class PlatformAwareSeller(Trader):
) )
def _get_pool_address(self, token_info: TokenInfo, address_provider: AddressProvider) -> Pubkey: def _get_pool_address(self, token_info: TokenInfo, address_provider: AddressProvider) -> Pubkey:
"""Get the pool/curve address for price calculations. """Get the pool/curve address for price calculations using platform-agnostic method."""
# Try to get the address from token_info first, then derive if needed
Args:
token_info: Token information
address_provider: Platform address provider
Returns:
Pool/curve address
"""
if token_info.platform == Platform.PUMP_FUN: if token_info.platform == Platform.PUMP_FUN:
return token_info.bonding_curve or address_provider.derive_pool_address(token_info.mint) if hasattr(token_info, 'bonding_curve') and token_info.bonding_curve:
return token_info.bonding_curve
elif token_info.platform == Platform.LETS_BONK: elif token_info.platform == Platform.LETS_BONK:
return token_info.pool_state or address_provider.derive_pool_address(token_info.mint) if hasattr(token_info, 'pool_state') and token_info.pool_state:
else: return token_info.pool_state
# Fallback to deriving the address
return address_provider.derive_pool_address(token_info.mint) # Fallback to deriving the address using platform provider
return address_provider.derive_pool_address(token_info.mint)
+46 -95
View File
@@ -1,6 +1,6 @@
""" """
Universal trading coordinator that works with any platform. Universal trading coordinator that works with any platform.
Replaces PumpTrader with platform-agnostic implementation. Cleaned up to remove all platform-specific hardcoding.
""" """
import asyncio import asyncio
@@ -84,18 +84,7 @@ class UniversalTrader:
marry_mode: bool = False, marry_mode: bool = False,
yolo_mode: bool = False, yolo_mode: bool = False,
): ):
"""Initialize the universal trader. """Initialize the universal trader."""
Args:
rpc_endpoint: RPC endpoint URL
wss_endpoint: WebSocket endpoint URL
private_key: Wallet private key
buy_amount: Amount of SOL to spend on buys
buy_slippage: Slippage tolerance for buys
sell_slippage: Slippage tolerance for sells
platform: Platform to trade on (Platform enum or string)
... (other args same as PumpTrader)
"""
# Core components # Core components
self.solana_client = SolanaClient(rpc_endpoint) self.solana_client = SolanaClient(rpc_endpoint)
self.wallet = Wallet(private_key) self.wallet = Wallet(private_key)
@@ -116,6 +105,15 @@ class UniversalTrader:
logger.info(f"Initialized Universal Trader for platform: {self.platform.value}") logger.info(f"Initialized Universal Trader for platform: {self.platform.value}")
# Validate platform support
try:
from platforms import platform_factory
if not platform_factory.registry.is_platform_supported(self.platform):
raise ValueError(f"Platform {self.platform.value} is not supported")
except Exception as e:
logger.error(f"Platform validation failed: {e}")
raise
# Get platform-specific implementations # Get platform-specific implementations
self.platform_implementations = get_platform_implementations( self.platform_implementations = get_platform_implementations(
self.platform, self.solana_client self.platform, self.solana_client
@@ -141,7 +139,7 @@ class UniversalTrader:
max_retries, max_retries,
) )
# Initialize the appropriate listener # Initialize the appropriate listener with platform filtering
self.token_listener = ListenerFactory.create_listener( self.token_listener = ListenerFactory.create_listener(
listener_type=listener_type, listener_type=listener_type,
wss_endpoint=wss_endpoint, wss_endpoint=wss_endpoint,
@@ -253,11 +251,7 @@ class UniversalTrader:
logger.info("Universal Trader has shut down") logger.info("Universal Trader has shut down")
async def _wait_for_token(self) -> TokenInfo | None: async def _wait_for_token(self) -> TokenInfo | None:
"""Wait for a single token to be detected. """Wait for a single token to be detected."""
Returns:
TokenInfo or None if timeout occurs
"""
# Create a one-time event to signal when a token is found # Create a one-time event to signal when a token is found
token_found = asyncio.Event() token_found = asyncio.Event()
found_token = None found_token = None
@@ -322,11 +316,7 @@ class UniversalTrader:
await self.solana_client.close() await self.solana_client.close()
async def _queue_token(self, token_info: TokenInfo) -> None: async def _queue_token(self, token_info: TokenInfo) -> None:
"""Queue a token for processing if not already processed. """Queue a token for processing if not already processed."""
Args:
token_info: Token information to queue
"""
token_key = str(token_info.mint) token_key = str(token_info.mint)
if token_key in self.processed_tokens: if token_key in self.processed_tokens:
@@ -360,7 +350,6 @@ class UniversalTrader:
await self._handle_token(token_info) await self._handle_token(token_info)
except asyncio.CancelledError: except asyncio.CancelledError:
# Handle cancellation gracefully
logger.info("Token queue processor was cancelled") logger.info("Token queue processor was cancelled")
break break
except Exception as e: except Exception as e:
@@ -369,11 +358,7 @@ class UniversalTrader:
self.token_queue.task_done() self.token_queue.task_done()
async def _handle_token(self, token_info: TokenInfo) -> None: async def _handle_token(self, token_info: TokenInfo) -> None:
"""Handle a new token creation event. """Handle a new token creation event."""
Args:
token_info: Token information
"""
try: try:
# Validate that token is for our platform # Validate that token is for our platform
if token_info.platform != self.platform: if token_info.platform != self.platform:
@@ -404,12 +389,7 @@ class UniversalTrader:
logger.error(f"Error handling token {token_info.symbol}: {e!s}") logger.error(f"Error handling token {token_info.symbol}: {e!s}")
async def _handle_successful_buy(self, token_info: TokenInfo, buy_result: TradeResult) -> None: async def _handle_successful_buy(self, token_info: TokenInfo, buy_result: TradeResult) -> None:
"""Handle successful token purchase. """Handle successful token purchase."""
Args:
token_info: Token information
buy_result: The result of the buy operation
"""
logger.info(f"Successfully bought {token_info.symbol} on {token_info.platform.value}") logger.info(f"Successfully bought {token_info.symbol} on {token_info.platform.value}")
self._log_trade("buy", token_info, buy_result.price, buy_result.amount, buy_result.tx_signature) self._log_trade("buy", token_info, buy_result.price, buy_result.amount, buy_result.tx_signature)
self.traded_mints.add(token_info.mint) self.traded_mints.add(token_info.mint)
@@ -426,12 +406,7 @@ class UniversalTrader:
logger.info("Marry mode enabled. Skipping sell operation.") logger.info("Marry mode enabled. Skipping sell operation.")
async def _handle_failed_buy(self, token_info: TokenInfo, buy_result: TradeResult) -> None: async def _handle_failed_buy(self, token_info: TokenInfo, buy_result: TradeResult) -> None:
"""Handle failed token purchase. """Handle failed token purchase."""
Args:
token_info: Token information
buy_result: The result of the buy operation
"""
logger.error(f"Failed to buy {token_info.symbol}: {buy_result.error_message}") logger.error(f"Failed to buy {token_info.symbol}: {buy_result.error_message}")
# Close ATA if enabled # Close ATA if enabled
await handle_cleanup_after_failure( await handle_cleanup_after_failure(
@@ -445,12 +420,7 @@ class UniversalTrader:
) )
async def _handle_tp_sl_exit(self, token_info: TokenInfo, buy_result: TradeResult) -> None: async def _handle_tp_sl_exit(self, token_info: TokenInfo, buy_result: TradeResult) -> None:
"""Handle take profit/stop loss exit strategy. """Handle take profit/stop loss exit strategy."""
Args:
token_info: Token information
buy_result: Result from the buy operation
"""
# Create position # Create position
position = Position.create_from_buy_result( position = Position.create_from_buy_result(
mint=token_info.mint, mint=token_info.mint,
@@ -472,11 +442,7 @@ class UniversalTrader:
await self._monitor_position_until_exit(token_info, position) await self._monitor_position_until_exit(token_info, position)
async def _handle_time_based_exit(self, token_info: TokenInfo) -> None: async def _handle_time_based_exit(self, token_info: TokenInfo) -> None:
"""Handle legacy time-based exit strategy. """Handle legacy time-based exit strategy."""
Args:
token_info: Token information
"""
logger.info(f"Waiting for {self.wait_time_after_buy} seconds before selling...") logger.info(f"Waiting for {self.wait_time_after_buy} seconds before selling...")
await asyncio.sleep(self.wait_time_after_buy) await asyncio.sleep(self.wait_time_after_buy)
@@ -500,15 +466,10 @@ class UniversalTrader:
logger.error(f"Failed to sell {token_info.symbol}: {sell_result.error_message}") logger.error(f"Failed to sell {token_info.symbol}: {sell_result.error_message}")
async def _monitor_position_until_exit(self, token_info: TokenInfo, position: Position) -> None: async def _monitor_position_until_exit(self, token_info: TokenInfo, position: Position) -> None:
"""Monitor a position until exit conditions are met. """Monitor a position until exit conditions are met."""
Args:
token_info: Token information
position: Position to monitor
"""
logger.info(f"Starting position monitoring (check interval: {self.price_check_interval}s)") logger.info(f"Starting position monitoring (check interval: {self.price_check_interval}s)")
# Get pool address for price monitoring # Get pool address for price monitoring using platform-agnostic method
pool_address = self._get_pool_address(token_info) pool_address = self._get_pool_address(token_info)
curve_manager = self.platform_implementations.curve_manager curve_manager = self.platform_implementations.curve_manager
@@ -570,35 +531,25 @@ class UniversalTrader:
await asyncio.sleep(self.price_check_interval) # Continue monitoring despite errors await asyncio.sleep(self.price_check_interval) # Continue monitoring despite errors
def _get_pool_address(self, token_info: TokenInfo) -> Pubkey: def _get_pool_address(self, token_info: TokenInfo) -> Pubkey:
"""Get the pool/curve address for price monitoring. """Get the pool/curve address for price monitoring using platform-agnostic method."""
Args:
token_info: Token information
Returns:
Pool/curve address
"""
address_provider = self.platform_implementations.address_provider address_provider = self.platform_implementations.address_provider
if token_info.platform == Platform.PUMP_FUN: # Use platform-specific logic to get the appropriate address
return token_info.bonding_curve or address_provider.derive_pool_address(token_info.mint) if hasattr(token_info, 'bonding_curve') and token_info.bonding_curve:
elif token_info.platform == Platform.LETS_BONK: return token_info.bonding_curve
return token_info.pool_state or address_provider.derive_pool_address(token_info.mint) elif hasattr(token_info, 'pool_state') and token_info.pool_state:
return token_info.pool_state
else: else:
# Fallback to deriving the address # Fallback to deriving the address using platform provider
return address_provider.derive_pool_address(token_info.mint) return address_provider.derive_pool_address(token_info.mint)
async def _save_token_info(self, token_info: TokenInfo) -> None: async def _save_token_info(self, token_info: TokenInfo) -> None:
"""Save token information to a file. """Save token information to a file."""
Args:
token_info: Token information
"""
try: try:
os.makedirs("trades", exist_ok=True) os.makedirs("trades", exist_ok=True)
file_name = os.path.join("trades", f"{token_info.mint}.txt") file_name = os.path.join("trades", f"{token_info.mint}.txt")
# Convert to dictionary for saving # Convert to dictionary for saving - platform-agnostic
token_dict = { token_dict = {
"name": token_info.name, "name": token_info.name,
"symbol": token_info.symbol, "symbol": token_info.symbol,
@@ -607,14 +558,22 @@ class UniversalTrader:
"platform": token_info.platform.value, "platform": token_info.platform.value,
"user": str(token_info.user) if token_info.user else None, "user": str(token_info.user) if token_info.user else None,
"creator": str(token_info.creator) if token_info.creator else None, "creator": str(token_info.creator) if token_info.creator else None,
# Platform-specific fields "creation_timestamp": token_info.creation_timestamp,
"bonding_curve": str(token_info.bonding_curve) if token_info.bonding_curve else None,
"associated_bonding_curve": str(token_info.associated_bonding_curve) if token_info.associated_bonding_curve else None,
"creator_vault": str(token_info.creator_vault) if token_info.creator_vault else None,
"pool_state": str(token_info.pool_state) if token_info.pool_state else None,
"base_vault": str(token_info.base_vault) if token_info.base_vault else None,
"quote_vault": str(token_info.quote_vault) if token_info.quote_vault else None,
} }
# Add platform-specific fields only if they exist
platform_fields = {
"bonding_curve": token_info.bonding_curve,
"associated_bonding_curve": token_info.associated_bonding_curve,
"creator_vault": token_info.creator_vault,
"pool_state": token_info.pool_state,
"base_vault": token_info.base_vault,
"quote_vault": token_info.quote_vault,
}
for field_name, field_value in platform_fields.items():
if field_value is not None:
token_dict[field_name] = str(field_value)
with open(file_name, "w") as file: with open(file_name, "w") as file:
file.write(json.dumps(token_dict, indent=2)) file.write(json.dumps(token_dict, indent=2))
@@ -624,15 +583,7 @@ class UniversalTrader:
logger.error(f"Failed to save token information: {e!s}") logger.error(f"Failed to save token information: {e!s}")
def _log_trade(self, action: str, token_info: TokenInfo, price: float, amount: float, tx_hash: str | None) -> None: def _log_trade(self, action: str, token_info: TokenInfo, price: float, amount: float, tx_hash: str | None) -> None:
"""Log trade information. """Log trade information."""
Args:
action: Trade action (buy/sell)
token_info: Token information
price: Token price in SOL
amount: Trade amount in SOL
tx_hash: Transaction hash
"""
try: try:
os.makedirs("trades", exist_ok=True) os.makedirs("trades", exist_ok=True)