Team Ai
Modelpublic

ParallelLLC/algorithmic_trading

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
27likes32downloads
websocket_server.py461 linesDownload Raw Back to ui
1"""2WebSocket Server for Real-time Trading Data3 4Provides real-time updates for:5- Market data streaming6- Trading signals7- Portfolio updates8- System alerts9"""10 11import asyncio12import websockets13import json14import logging15import threading16import time17from datetime import datetime, timedelta18from typing import Dict, Any, List, Optional19import pandas as pd20import os21import sys22 23# Add project root to path24sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))25 26from agentic_ai_system.main import load_config27from agentic_ai_system.data_ingestion import load_data, add_technical_indicators28from agentic_ai_system.alpaca_broker import AlpacaBroker29from agentic_ai_system.finrl_agent import FinRLAgent, FinRLConfig30 31class TradingWebSocketServer:32    def __init__(self, host="localhost", port=8765):33        self.host = host34        self.port = port35        self.clients = set()36        self.config = None37        self.alpaca_broker = None38        self.finrl_agent = None39        self.trading_active = False40        self.market_data = None41        self.portfolio_data = {}42        43        # Setup logging44        logging.basicConfig(level=logging.INFO)45        self.logger = logging.getLogger(__name__)46    47    async def register(self, websocket):48        """Register a new client"""49        self.clients.add(websocket)50        self.logger.info(f"Client connected. Total clients: {len(self.clients)}")51        52        # Send initial data53        await self.send_initial_data(websocket)54    55    async def unregister(self, websocket):56        """Unregister a client"""57        self.clients.remove(websocket)58        self.logger.info(f"Client disconnected. Total clients: {len(self.clients)}")59    60    async def send_initial_data(self, websocket):61        """Send initial data to new client"""62        initial_data = {63            "type": "initial_data",64            "timestamp": datetime.now().isoformat(),65            "config": self.config,66            "portfolio": self.portfolio_data,67            "trading_status": self.trading_active68        }69        await websocket.send(json.dumps(initial_data))70    71    async def broadcast(self, message):72        """Broadcast message to all connected clients"""73        if self.clients:74            message_str = json.dumps(message)75            await asyncio.gather(76                *[client.send(message_str) for client in self.clients],77                return_exceptions=True78            )79    80    async def handle_market_data(self):81        """Handle real-time market data updates"""82        while True:83            try:84                if self.config and self.alpaca_broker:85                    # Get real-time market data86                    symbol = self.config['trading']['symbol']87                    88                    # Get current price89                    current_price = await self.get_current_price(symbol)90                    91                    if current_price:92                        market_update = {93                            "type": "market_data",94                            "timestamp": datetime.now().isoformat(),95                            "symbol": symbol,96                            "price": current_price,97                            "volume": await self.get_current_volume(symbol)98                        }99                        100                        await self.broadcast(market_update)101                        self.logger.info(f"Broadcasted market data for {symbol}: ${current_price}")102                103                await asyncio.sleep(1)  # Update every second104                105            except Exception as e:106                self.logger.error(f"Error in market data handler: {e}")107                await asyncio.sleep(5)  # Wait before retrying108    109    async def handle_portfolio_updates(self):110        """Handle portfolio updates"""111        while True:112            try:113                if self.alpaca_broker:114                    # Get portfolio information115                    account_info = self.alpaca_broker.get_account_info()116                    positions = self.alpaca_broker.get_positions()117                    118                    if account_info:119                        portfolio_update = {120                            "type": "portfolio_update",121                            "timestamp": datetime.now().isoformat(),122                            "account": {123                                "buying_power": float(account_info['buying_power']),124                                "portfolio_value": float(account_info['portfolio_value']),125                                "equity": float(account_info['equity']),126                                "cash": float(account_info['cash'])127                            },128                            "positions": positions if positions else []129                        }130                        131                        await self.broadcast(portfolio_update)132                        self.portfolio_data = portfolio_update133                134                await asyncio.sleep(5)  # Update every 5 seconds135                136            except Exception as e:137                self.logger.error(f"Error in portfolio updates: {e}")138                await asyncio.sleep(10)  # Wait before retrying139    140    async def handle_trading_signals(self):141        """Handle trading signals from FinRL agent"""142        while True:143            try:144                if self.trading_active and self.finrl_agent and self.market_data is not None:145                    # Generate trading signals146                    signal = await self.generate_trading_signal()147                    148                    if signal:149                        signal_update = {150                            "type": "trading_signal",151                            "timestamp": datetime.now().isoformat(),152                            "signal": signal153                        }154                        155                        await self.broadcast(signal_update)156                        self.logger.info(f"Broadcasted trading signal: {signal}")157                158                await asyncio.sleep(10)  # Generate signals every 10 seconds159                160            except Exception as e:161                self.logger.error(f"Error in trading signals: {e}")162                await asyncio.sleep(30)  # Wait before retrying163    164    async def get_current_price(self, symbol):165        """Get current price for symbol"""166        try:167            if self.alpaca_broker:168                # Get latest price from Alpaca169                latest_trade = self.alpaca_broker.get_latest_trade(symbol)170                if latest_trade:171                    return float(latest_trade['p'])172            return None173        except Exception as e:174            self.logger.error(f"Error getting current price: {e}")175            return None176    177    async def get_current_volume(self, symbol):178        """Get current volume for symbol"""179        try:180            if self.alpaca_broker:181                # Get latest trade volume182                latest_trade = self.alpaca_broker.get_latest_trade(symbol)183                if latest_trade:184                    return int(latest_trade['s'])185            return None186        except Exception as e:187            self.logger.error(f"Error getting current volume: {e}")188            return None189    190    async def generate_trading_signal(self):191        """Generate trading signal using FinRL agent"""192        try:193            if self.finrl_agent and self.market_data is not None:194                # Use recent data for prediction195                recent_data = self.market_data.tail(100)196                197                prediction_result = self.finrl_agent.predict(198                    data=recent_data,199                    config=self.config,200                    use_real_broker=False201                )202                203                if prediction_result['success']:204                    # Generate signal based on prediction205                    current_price = await self.get_current_price(self.config['trading']['symbol'])206                    207                    if current_price:208                        signal = {209                            "action": "HOLD",  # Default action210                            "confidence": 0.5,211                            "price": current_price,212                            "reasoning": "Model prediction"213                        }214                        215                        # Determine action based on prediction216                        if prediction_result['total_return'] > 0.02:  # 2% positive return217                            signal["action"] = "BUY"218                            signal["confidence"] = min(0.9, 0.5 + abs(prediction_result['total_return']))219                        elif prediction_result['total_return'] < -0.02:  # 2% negative return220                            signal["action"] = "SELL"221                            signal["confidence"] = min(0.9, 0.5 + abs(prediction_result['total_return']))222                        223                        return signal224            225            return None226        except Exception as e:227            self.logger.error(f"Error generating trading signal: {e}")228            return None229    230    async def handle_client_message(self, websocket, message):231        """Handle incoming client messages"""232        try:233            data = json.loads(message)234            message_type = data.get("type")235            236            if message_type == "load_config":237                # Load configuration238                config_file = data.get("config_file", "config.yaml")239                self.config = load_config(config_file)240                241                response = {242                    "type": "config_loaded",243                    "success": True,244                    "config": self.config245                }246                await websocket.send(json.dumps(response))247            248            elif message_type == "connect_alpaca":249                # Connect to Alpaca250                api_key = data.get("api_key")251                secret_key = data.get("secret_key")252                253                if api_key and secret_key:254                    self.config['alpaca']['api_key'] = api_key255                    self.config['alpaca']['secret_key'] = secret_key256                    self.config['execution']['broker_api'] = 'alpaca_paper'257                    258                    self.alpaca_broker = AlpacaBroker(self.config)259                    260                    response = {261                        "type": "alpaca_connected",262                        "success": True263                    }264                    await websocket.send(json.dumps(response))265                else:266                    response = {267                        "type": "alpaca_connected",268                        "success": False,269                        "error": "Missing API credentials"270                    }271                    await websocket.send(json.dumps(response))272            273            elif message_type == "start_trading":274                # Start trading275                self.trading_active = True276                277                response = {278                    "type": "trading_started",279                    "success": True280                }281                await websocket.send(json.dumps(response))282                283                # Broadcast to all clients284                await self.broadcast({285                    "type": "trading_status",286                    "active": True,287                    "timestamp": datetime.now().isoformat()288                })289            290            elif message_type == "stop_trading":291                # Stop trading292                self.trading_active = False293                294                response = {295                    "type": "trading_stopped",296                    "success": True297                }298                await websocket.send(json.dumps(response))299                300                # Broadcast to all clients301                await self.broadcast({302                    "type": "trading_status",303                    "active": False,304                    "timestamp": datetime.now().isoformat()305                })306            307            elif message_type == "load_data":308                # Load market data309                if self.config:310                    self.market_data = load_data(self.config)311                    if self.market_data is not None:312                        self.market_data = add_technical_indicators(self.market_data)313                        314                        response = {315                            "type": "data_loaded",316                            "success": True,317                            "data_points": len(self.market_data)318                        }319                    else:320                        response = {321                            "type": "data_loaded",322                            "success": False,323                            "error": "Failed to load data"324                        }325                else:326                    response = {327                        "type": "data_loaded",328                        "success": False,329                        "error": "Configuration not loaded"330                    }331                332                await websocket.send(json.dumps(response))333            334            elif message_type == "train_model":335                # Train FinRL model336                if self.market_data is not None:337                    algorithm = data.get("algorithm", "PPO")338                    learning_rate = data.get("learning_rate", 0.0003)339                    training_steps = data.get("training_steps", 100000)340                    341                    finrl_config = FinRLConfig(342                        algorithm=algorithm,343                        learning_rate=learning_rate,344                        batch_size=64,345                        buffer_size=1000000,346                        learning_starts=100,347                        gamma=0.99,348                        tau=0.005,349                        train_freq=1,350                        gradient_steps=1,351                        verbose=1,352                        tensorboard_log='logs/finrl_tensorboard'353                    )354                    355                    self.finrl_agent = FinRLAgent(finrl_config)356                    357                    # Train in background thread358                    def train_model():359                        try:360                            result = self.finrl_agent.train(361                                data=self.market_data,362                                config=self.config,363                                total_timesteps=training_steps,364                                use_real_broker=False365                            )366                            367                            # Broadcast training completion368                            asyncio.create_task(self.broadcast({369                                "type": "training_completed",370                                "success": result['success'],371                                "result": result372                            }))373                        except Exception as e:374                            asyncio.create_task(self.broadcast({375                                "type": "training_completed",376                                "success": False,377                                "error": str(e)378                            }))379                    380                    training_thread = threading.Thread(target=train_model)381                    training_thread.daemon = True382                    training_thread.start()383                    384                    response = {385                        "type": "training_started",386                        "success": True387                    }388                else:389                    response = {390                        "type": "training_started",391                        "success": False,392                        "error": "Market data not loaded"393                    }394                395                await websocket.send(json.dumps(response))396            397            else:398                # Unknown message type399                response = {400                    "type": "error",401                    "message": f"Unknown message type: {message_type}"402                }403                await websocket.send(json.dumps(response))404                405        except json.JSONDecodeError:406            response = {407                "type": "error",408                "message": "Invalid JSON message"409            }410            await websocket.send(json.dumps(response))411        except Exception as e:412            response = {413                "type": "error",414                "message": f"Server error: {str(e)}"415            }416            await websocket.send(json.dumps(response))417    418    async def websocket_handler(self, websocket, path):419        """Main WebSocket handler"""420        await self.register(websocket)421        try:422            async for message in websocket:423                await self.handle_client_message(websocket, message)424        except websockets.exceptions.ConnectionClosed:425            pass426        finally:427            await self.unregister(websocket)428    429    async def start_server(self):430        """Start the WebSocket server"""431        # Start background tasks432        asyncio.create_task(self.handle_market_data())433        asyncio.create_task(self.handle_portfolio_updates())434        asyncio.create_task(self.handle_trading_signals())435        436        # Start WebSocket server437        server = await websockets.serve(438            self.websocket_handler,439            self.host,440            self.port441        )442        443        self.logger.info(f"WebSocket server started on ws://{self.host}:{self.port}")444        445        # Keep server running446        await server.wait_closed()447    448    def run_server(self):449        """Run the server in a separate thread"""450        def run():451            asyncio.run(self.start_server())452        453        server_thread = threading.Thread(target=run)454        server_thread.daemon = True455        server_thread.start()456        457        return server_thread458 459def create_websocket_server(host="localhost", port=8765):460    """Create and return a WebSocket server instance"""461    return TradingWebSocketServer(host=host, port=port)