diff --git a/bin/fxt/build.rs b/bin/fxt/build.rs index 8982929af..74c0694ce 100644 --- a/bin/fxt/build.rs +++ b/bin/fxt/build.rs @@ -27,7 +27,7 @@ fn main() -> Result<(), Box> { .server_mod_attribute(".", "#[allow(unused_qualifications)]") .client_mod_attribute(".", "#[allow(unused_qualifications)]") .protoc_arg("--experimental_allow_proto3_optional") - .compile_protos(&protos, &[proto_root.to_string()])?; + .compile_protos(&protos, &[proto_root.to_owned()])?; for proto in &[ "trading", diff --git a/bin/fxt/src/auth/encryption.rs b/bin/fxt/src/auth/encryption.rs index 98cbcecaf..f336c72c7 100644 --- a/bin/fxt/src/auth/encryption.rs +++ b/bin/fxt/src/auth/encryption.rs @@ -230,11 +230,16 @@ pub fn decrypt_token(encrypted: &str, key: &[u8]) -> Result )); } - // Extract nonce (first 12 bytes) - let nonce = Nonce::from_slice(&combined[0..12]); + // Extract nonce (first 12 bytes) -- length validated above (>= 28). + let nonce_bytes = combined.get(..12).ok_or_else(|| { + CommonError::service(ErrorCategory::Security, "Nonce extraction failed".to_owned()) + })?; + let nonce = Nonce::from_slice(nonce_bytes); // Extract ciphertext + tag (remaining bytes, tag is last 16 bytes included in ciphertext) - let ciphertext = &combined[12..]; + let ciphertext = combined.get(12..).ok_or_else(|| { + CommonError::service(ErrorCategory::Security, "Ciphertext extraction failed".to_owned()) + })?; // Create AES-256-GCM cipher let cipher = Aes256Gcm::new_from_slice(key).map_err(|e| { diff --git a/bin/fxt/src/auth/token_manager.rs b/bin/fxt/src/auth/token_manager.rs index 25929fa55..2bbacd0b0 100644 --- a/bin/fxt/src/auth/token_manager.rs +++ b/bin/fxt/src/auth/token_manager.rs @@ -682,20 +682,17 @@ impl AuthTokenManager { /// Check if the token needs refresh (expired or near expiration) pub async fn needs_refresh(&self) -> bool { // Read tokens directly from storage (don't use get_current_token which filters expired tokens) - let access_token = match self.storage.get_access_token().await { - Ok(Some(token)) => token, - _ => return false, // No token = no refresh needed + let Ok(Some(access_token)) = self.storage.get_access_token().await else { + return false; // No token = no refresh needed }; - let refresh_token = match self.storage.get_refresh_token().await { - Ok(Some(token)) => token, - _ => return false, // No refresh token = can't refresh + let Ok(Some(refresh_token)) = self.storage.get_refresh_token().await else { + return false; // No refresh token = can't refresh }; // Extract expiry from access token - let expires_at = match extract_token_expiry(&access_token) { - Ok(exp) => exp, - Err(_) => return false, // Can't parse = assume no refresh needed + let Ok(expires_at) = extract_token_expiry(&access_token) else { + return false; // Can't parse = assume no refresh needed }; let token_info = TokenInfo { diff --git a/bin/fxt/src/client/backtesting_client.rs b/bin/fxt/src/client/backtesting_client.rs deleted file mode 100644 index 87e1f05f6..000000000 --- a/bin/fxt/src/client/backtesting_client.rs +++ /dev/null @@ -1,198 +0,0 @@ -//! Backtesting service gRPC client -//! -//! Provides a gRPC client for communicating with the backtesting service, -//! including strategy testing, performance analysis, and historical simulation. - -use anyhow::{Context, Result as AnyhowResult}; -use serde::{Deserialize, Serialize}; -use tonic::transport::Channel; - -/// Configuration for the backtesting service gRPC client -/// -/// Contains connection parameters and client behavior settings for -/// communicating with the backtesting service. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BacktestingClientConfig { - /// gRPC endpoint URL for the backtesting service (MUST be https://) - pub endpoint: String, - /// Request timeout in milliseconds - pub timeout_ms: u64, -} - -impl Default for BacktestingClientConfig { - /// Create default backtesting client configuration - /// - /// Returns configuration with API Gateway HTTPS endpoint and 60-second timeout - /// (longer than trading client due to potentially long-running backtests). - /// - /// # Security - /// - /// Defaults to HTTPS (not HTTP) to enforce encrypted connections. - /// - /// # Wave 71 Update - /// - /// Endpoint changed to API Gateway (port 50050) instead of direct backtesting service. - /// - /// All requests now route through the API Gateway for centralized authentication. - fn default() -> Self { - Self { - endpoint: "https://localhost:50050".to_owned(), // API Gateway endpoint - timeout_ms: 60_000, - } - } -} - -/// gRPC client for the backtesting service -/// -/// Manages connection to the backtesting service and provides methods for -/// running strategy backtests, retrieving performance metrics, and managing -/// historical simulations. -#[derive(Debug)] -pub struct BacktestingClient { - /// Client configuration - config: BacktestingClientConfig, - /// Active gRPC channel (None if disconnected) - channel: Option, -} - -impl BacktestingClient { - /// Create a new backtesting client with the specified configuration - /// - /// # Arguments - /// * `config` - Client configuration including endpoint and timeout settings - /// - /// # Returns - /// - /// New `BacktestingClient` instance (not yet connected) - pub const fn new(config: BacktestingClientConfig) -> Self { - Self { - config, - channel: None, - } - } - - /// Validate URL scheme is HTTPS (rejects HTTP) - /// - /// # Security - /// - /// This enforces fail-closed behavior - only HTTPS connections are allowed. - /// - /// Insecure HTTP connections are rejected with a clear error message. - fn validate_endpoint_security(endpoint: &str) -> AnyhowResult<()> { - if endpoint.starts_with("http://") { - anyhow::bail!( - "SECURITY ERROR: Insecure HTTP endpoint rejected: {}. TLS (https://) is required for all gRPC connections to protect sensitive backtesting data.", - endpoint - ); - } - - if !endpoint.starts_with("https://") { - anyhow::bail!( - "Invalid endpoint scheme in URL: {}. Only HTTPS is supported (example: https://backtesting-service:50053).", - endpoint - ); - } - - Ok(()) - } - - /// Establish connection to the backtesting service - /// - /// Creates a gRPC channel to the configured backtesting service endpoint. - /// - /// Must be called before making any service requests. - /// - /// # Returns - /// `Result<(), anyhow::Error>` - Ok if connection successful - /// - /// # Errors - /// - /// Returns error if: - /// - Endpoint uses insecure HTTP scheme (security validation failure) - /// - /// - Unable to parse endpoint URL - /// - Unable to connect to the service endpoint - /// - /// # Security - /// - /// This method enforces TLS by rejecting any HTTP endpoints. - /// - /// Use HTTPS endpoints only (e.g., ). - pub async fn connect(&mut self) -> AnyhowResult<()> { - // SECURITY: Validate endpoint uses HTTPS before attempting connection - Self::validate_endpoint_security(&self.config.endpoint) - .context("Backtesting client endpoint security validation failed")?; - - // Parse endpoint with proper error handling (no unwrap/panic) - let channel = Channel::from_shared(self.config.endpoint.clone()) - .context("Failed to parse backtesting service endpoint URL - check URL format")? - .connect() - .await - .context("Failed to establish connection to backtesting service - check network and TLS configuration")?; - - self.channel = Some(channel); - - tracing::info!( - "\u{2705} Backtesting client connected securely via TLS to {}", - self.config.endpoint - ); - - Ok(()) - } - - /// Check if the client is currently connected to the backtesting service - /// - /// # Returns - /// `true` if connected, `false` if disconnected - pub const fn is_connected(&self) -> bool { - self.channel.is_some() - } - - /// Shutdown the client and close the connection - /// - /// Cleanly closes the gRPC channel and releases associated resources. - /// - /// The client can be reconnected after shutdown by calling `connect()`. - pub async fn shutdown(&mut self) { - if self.channel.is_some() { - tracing::info!( - "Shutting down backtesting client connection to {}", - self.config.endpoint - ); - } - self.channel = None; - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_default_uses_https() { - let config = BacktestingClientConfig::default(); - assert!( - config.endpoint.starts_with("https://"), - "Default config must use HTTPS" - ); - } - - #[test] - fn test_http_validation_rejects_insecure() { - let result = BacktestingClient::validate_endpoint_security("http://localhost:50053"); - assert!(result.is_err(), "HTTP endpoints should be rejected"); - assert!(result.unwrap_err().to_string().contains("Insecure HTTP")); - } - - #[test] - fn test_https_validation_accepts_secure() { - let result = BacktestingClient::validate_endpoint_security("https://localhost:50053"); - assert!(result.is_ok(), "HTTPS endpoints should be accepted"); - } - - #[test] - fn test_invalid_scheme_rejected() { - let result = BacktestingClient::validate_endpoint_security("grpc://localhost:50053"); - assert!(result.is_err(), "Non-HTTPS schemes should be rejected"); - } -} diff --git a/bin/fxt/src/client/connection_manager.rs b/bin/fxt/src/client/connection_manager.rs deleted file mode 100644 index 53f031e77..000000000 --- a/bin/fxt/src/client/connection_manager.rs +++ /dev/null @@ -1,167 +0,0 @@ -//! Connection management for TLI gRPC clients -//! -//! This module provides connection pooling, health monitoring, and statistics -//! tracking for all gRPC client connections to backend services. - -use serde::{Deserialize, Serialize}; -use std::sync::Arc; -use tokio::sync::RwLock; - -/// Configuration parameters for gRPC client connections -/// -/// Contains all settings needed to establish and maintain connections to backend services, -/// including authentication, timeouts, and retry policies. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ConnectionConfig { - /// Server URL for the gRPC service (MUST be https://, e.g., "") - pub server_url: String, - /// Optional authentication token for secure connections - pub auth_token: Option, - /// Connection timeout in milliseconds - pub timeout_ms: u64, - /// Maximum number of retry attempts for failed requests - pub max_retries: u32, -} - -impl Default for ConnectionConfig { - /// Create default configuration for local development - /// - /// # Wave 71 Update - /// - /// Default endpoint changed to API Gateway (port 50050) instead of direct service connections. - /// - /// All TLI traffic now routes through the API Gateway for centralized authentication and routing. - fn default() -> Self { - // SECURITY: Default to HTTPS, not HTTP - // Wave 71: Connect to API Gateway instead of direct service endpoints - Self { - server_url: "https://localhost:50050".to_owned(), // API Gateway endpoint - auth_token: None, - timeout_ms: 10000, - max_retries: 3, - } - } -} - -/// Connection statistics and metrics -/// -/// Tracks performance metrics and error counts for monitoring -/// connection health and diagnosing issues. -#[derive(Debug, Clone)] -pub struct ConnectionStats { - /// Total number of messages sent to the server - pub messages_sent: u64, - /// Total number of messages received from the server - pub messages_received: u64, - /// Total bytes sent to the server - pub bytes_sent: u64, - /// Total bytes received from the server - pub bytes_received: u64, - /// Total number of connection errors encountered - pub connection_errors: u64, - /// Most recent error message, if any - pub last_error: Option, -} - -/// Connection manager for gRPC client connections -/// -/// Manages connection pools, health monitoring, and statistics tracking -/// for all backend service connections. Provides automatic reconnection -/// and load balancing capabilities. -#[derive(Debug)] -pub struct ConnectionManager { - /// Connection configuration - config: ConnectionConfig, - /// Thread-safe connection statistics - stats: Arc>, -} - -impl ConnectionManager { - /// Create a new connection manager with the given configuration - /// - /// # Arguments - /// * `config` - Connection configuration parameters - /// - /// # Returns - /// - /// A new `ConnectionManager` instance ready to manage connections - pub fn new(config: ConnectionConfig) -> Self { - Self { - config, - stats: Arc::new(RwLock::new(ConnectionStats { - messages_sent: 0, - messages_received: 0, - bytes_sent: 0, - bytes_received: 0, - connection_errors: 0, - last_error: None, - })), - } - } - - /// Get a reference to the connection configuration - pub fn config(&self) -> &ConnectionConfig { - &self.config - } - - /// Establish a connection to the configured server - /// - /// # Returns - /// - /// Ok(()) if connection succeeds, Err with error message if it fails - pub async fn connect(&self) -> Result<(), String> { - Ok(()) - } - - /// Disconnect from the server and cleanup resources - /// - /// # Returns - /// - /// Ok(()) if disconnection succeeds, Err with error message if it fails - pub async fn disconnect(&self) -> Result<(), String> { - Ok(()) - } - - /// Get current connection statistics - /// - /// # Returns - /// - /// A clone of the current connection statistics - pub async fn get_stats(&self) -> ConnectionStats { - self.stats.read().await.clone() - } - - /// Add a new service to the connection pool - /// - /// # Arguments - /// * `service_name` - Unique identifier for the service - /// - /// * `config` - Connection configuration for the service - /// - /// # Returns - /// - /// Ok(()) if service added successfully, Err with error message if it fails - pub async fn add_service( - &self, - _service_name: String, - _config: ConnectionConfig, - ) -> Result<(), String> { - Ok(()) - } - - /// Get statistics for all connections in the pool - /// - /// # Returns - /// `HashMap` mapping service names to their connection statistics - pub async fn get_pool_stats(&self) -> std::collections::HashMap> { - std::collections::HashMap::new() - } - - /// Gracefully shutdown all connections and cleanup resources - /// - /// This method ensures all active connections are properly closed - /// and resources are released before the manager is destroyed. - pub async fn shutdown(&self) { - // Shutdown implementation - } -} diff --git a/bin/fxt/src/client/ml_training_client.rs b/bin/fxt/src/client/ml_training_client.rs deleted file mode 100644 index 7b381d2c0..000000000 --- a/bin/fxt/src/client/ml_training_client.rs +++ /dev/null @@ -1,247 +0,0 @@ -//! ML training service gRPC client -//! -//! Provides a gRPC client for communicating with the ML training service, -//! including model training, resource monitoring, and training job management. - -use anyhow::{Context, Result as AnyhowResult}; -use serde::{Deserialize, Serialize}; -use tonic::transport::Channel; - -/// Configuration for the ML training service gRPC client -/// -/// Contains connection parameters and client behavior settings for -/// communicating with the ML training service. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct MLTrainingClientConfig { - /// gRPC endpoint URL for the ML training service (MUST be https://) - pub endpoint: String, - /// Request timeout in milliseconds - pub timeout_ms: u64, -} - -impl Default for MLTrainingClientConfig { - /// Create default ML training client configuration - /// - /// Returns configuration with API Gateway HTTPS endpoint and 120-second timeout - /// (longer timeout due to potentially long-running ML operations). - /// - /// # Security - /// - /// Defaults to HTTPS (not HTTP) to enforce encrypted connections. - /// - /// # Wave 71 Update - /// - /// Endpoint changed to API Gateway (port 50050) instead of direct ML training service. - /// - /// All requests now route through the API Gateway for centralized authentication. - fn default() -> Self { - Self { - endpoint: "https://localhost:50050".to_owned(), // API Gateway endpoint - timeout_ms: 120_000, - } - } -} - -/// gRPC client for the ML training service -/// -/// Manages connection to the ML training service and provides methods for -/// starting training jobs, monitoring progress, and managing model lifecycles. -#[derive(Debug)] -pub struct MLTrainingClient { - /// Client configuration - config: MLTrainingClientConfig, - /// Active gRPC channel (None if disconnected) - channel: Option, -} - -impl MLTrainingClient { - /// Create a new ML training client with the specified configuration - /// - /// # Arguments - /// * `config` - Client configuration including endpoint and timeout settings - /// - /// # Returns - /// - /// New `MLTrainingClient` instance (not yet connected) - pub const fn new(config: MLTrainingClientConfig) -> Self { - Self { - config, - channel: None, - } - } - - /// Validate URL scheme is HTTPS (rejects HTTP) - /// - /// # Security - /// - /// This enforces fail-closed behavior - only HTTPS connections are allowed. - /// - /// Insecure HTTP connections are rejected with a clear error message. - fn validate_endpoint_security(endpoint: &str) -> AnyhowResult<()> { - if endpoint.starts_with("http://") { - anyhow::bail!( - "SECURITY ERROR: Insecure HTTP endpoint rejected: {}. TLS (https://) is required for all gRPC connections to protect sensitive ML model data.", - endpoint - ); - } - - if !endpoint.starts_with("https://") { - anyhow::bail!( - "Invalid endpoint scheme in URL: {}. Only HTTPS is supported (example: https://ml-training-service:50054).", - endpoint - ); - } - - Ok(()) - } - - /// Establish connection to the ML training service - /// - /// Creates a gRPC channel to the configured ML training service endpoint. - /// - /// Must be called before making any service requests. - /// - /// # Returns - /// `Result<(), anyhow::Error>` - Ok if connection successful - /// - /// # Errors - /// - /// Returns error if: - /// - Endpoint uses insecure HTTP scheme (security validation failure) - /// - /// - Unable to parse endpoint URL - /// - Unable to connect to the service endpoint - /// - /// # Security - /// - /// This method enforces TLS by rejecting any HTTP endpoints. - /// - /// Use HTTPS endpoints only (e.g., ). - pub async fn connect(&mut self) -> AnyhowResult<()> { - // SECURITY: Validate endpoint uses HTTPS before attempting connection - Self::validate_endpoint_security(&self.config.endpoint) - .context("ML training client endpoint security validation failed")?; - - // Parse endpoint with proper error handling (no unwrap/panic) - let channel = Channel::from_shared(self.config.endpoint.clone()) - .context("Failed to parse ML training service endpoint URL - check URL format")? - .connect() - .await - .context("Failed to establish connection to ML training service - check network and TLS configuration")?; - - self.channel = Some(channel); - - tracing::info!( - "\u{2705} ML training client connected securely via TLS to {}", - self.config.endpoint - ); - - Ok(()) - } - - /// Check if the client is currently connected to the ML training service - /// - /// # Returns - /// `true` if connected, `false` if disconnected - pub const fn is_connected(&self) -> bool { - self.channel.is_some() - } - - /// Shutdown the client and close the connection - /// - /// Cleanly closes the gRPC channel and releases associated resources. - /// - /// The client can be reconnected after shutdown by calling `connect()`. - pub async fn shutdown(&mut self) { - if self.channel.is_some() { - tracing::info!( - "Shutting down ML training client connection to {}", - self.config.endpoint - ); - } - self.channel = None; - } -} - -/// Resource monitoring event for ML training operations -/// -/// Contains system resource utilization metrics during model training, -/// including CPU, memory, and optional GPU usage statistics. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ResourceMonitoringEvent { - /// CPU utilization as a percentage (0.0 to 100.0) - pub cpu_usage: f64, - /// Memory utilization as a percentage (0.0 to 100.0) - pub memory_usage: f64, - /// GPU utilization as a percentage (None if no GPU available) - pub gpu_usage: Option, - /// Unix timestamp when the metrics were recorded (nanoseconds) - pub timestamp: i64, -} - -/// Training job context and status information -/// -/// Provides metadata about an active or completed ML training job, -/// including identification, model information, and progress tracking. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TrainingJobContext { - /// Unique identifier for the training job - pub job_id: String, - /// Name of the model being trained - pub model_name: String, - /// Current status of the training job (e.g., "running", "completed", "failed") - pub status: String, - /// Training progress as a percentage (0.0 to 100.0) - pub progress: f64, -} - -/// Training progress event with performance metrics -/// -/// Contains detailed progress information from ML model training, -/// including loss metrics, accuracy measurements, and timing data. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TrainingProgressEvent { - /// Training job identifier this event belongs to - pub job_id: String, - /// Current training epoch number - pub epoch: u32, - /// Current loss value for the epoch - pub loss: f64, - /// Optional accuracy metric (if available for the model type) - pub accuracy: Option, - /// Unix timestamp when the progress was recorded (nanoseconds) - pub timestamp: i64, -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_default_uses_https() { - let config = MLTrainingClientConfig::default(); - assert!( - config.endpoint.starts_with("https://"), - "Default config must use HTTPS" - ); - } - - #[test] - fn test_http_validation_rejects_insecure() { - let result = MLTrainingClient::validate_endpoint_security("http://localhost:50054"); - assert!(result.is_err(), "HTTP endpoints should be rejected"); - assert!(result.unwrap_err().to_string().contains("Insecure HTTP")); - } - - #[test] - fn test_https_validation_accepts_secure() { - let result = MLTrainingClient::validate_endpoint_security("https://localhost:50054"); - assert!(result.is_ok(), "HTTPS endpoints should be accepted"); - } - - #[test] - fn test_invalid_scheme_rejected() { - let result = MLTrainingClient::validate_endpoint_security("ws://localhost:50054"); - assert!(result.is_err(), "Non-HTTPS schemes should be rejected"); - } -} diff --git a/bin/fxt/src/client/mod.rs b/bin/fxt/src/client/mod.rs deleted file mode 100644 index d95553e1d..000000000 --- a/bin/fxt/src/client/mod.rs +++ /dev/null @@ -1,319 +0,0 @@ -//! TLI gRPC client modules -//! -//! This module contains comprehensive gRPC client implementations for all -//! core trading system services with advanced features including: -//! - Connection pooling and health monitoring -//! - Real-time streaming support -//! - Automatic reconnection and circuit breakers -//! - Comprehensive error handling -//! - Metrics collection and alerting - -pub mod backtesting_client; -pub mod connection_manager; -pub mod ml_training_client; -pub mod trading_client; - -use anyhow::{Context, Result}; -use tonic::transport::Channel; - -/// Create a gRPC channel that auto-configures TLS for `https://` URLs. -/// -/// Connects eagerly (awaits the TCP handshake). -pub async fn connect_channel(url: &str) -> Result { - let mut endpoint = Channel::from_shared(url.to_owned()).context("Invalid gRPC endpoint URL")?; - if url.starts_with("https://") { - endpoint = endpoint - .tls_config(tonic::transport::ClientTlsConfig::new().with_enabled_roots()) - .context("Failed to configure TLS")?; - } - endpoint - .connect() - .await - .context("Failed to connect to API Gateway") -} - -/// Create a lazily-connected gRPC channel with TLS auto-detection. -/// -/// Does not attempt the TCP connection until the first RPC call. -pub fn connect_channel_lazy(url: &str) -> Result { - let mut endpoint = Channel::from_shared(url.to_owned()).context("Invalid gRPC endpoint URL")?; - if url.starts_with("https://") { - endpoint = endpoint - .tls_config(tonic::transport::ClientTlsConfig::new().with_enabled_roots()) - .context("Failed to configure TLS")?; - } - Ok(endpoint.connect_lazy()) -} - -/// Service endpoints configuration -#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)] -pub struct ServiceEndpoints { - /// Trading engine endpoint - pub trading_engine: String, - /// Market data endpoint - pub market_data: String, - /// Backtesting service endpoint - pub backtesting_service: String, - /// ML training service endpoint - pub ml_training_service: String, -} - -impl ServiceEndpoints { - /// Create default endpoints for local development - /// - /// # Security - /// - /// All endpoints use HTTPS (not HTTP) to enforce encrypted connections. - /// - /// # Wave 71 Update - /// - /// All endpoints now point to API Gateway (port 50050) instead of direct services. - /// - /// The API Gateway handles routing to backend services based on gRPC service names. - pub fn localhost() -> Self { - Self { - trading_engine: "https://localhost:50050".to_owned(), // API Gateway - market_data: "https://localhost:50050".to_owned(), // API Gateway - backtesting_service: "https://localhost:50050".to_owned(), // API Gateway - ml_training_service: "https://localhost:50050".to_owned(), // API Gateway - } - } -} - -/// Client factory for creating and managing all service clients -#[derive(Debug)] -pub struct ClientFactory { - /// Connection manager shared across all clients - connection_manager: std::sync::Arc, -} - -impl ClientFactory { - /// Create a new client factory - pub fn new(connection_config: connection_manager::ConnectionConfig) -> Self { - let connection_manager = std::sync::Arc::new(connection_manager::ConnectionManager::new( - connection_config, - )); - - Self { - connection_manager, - } - } - - /// Create a trading client - pub const fn create_trading_client( - &self, - config: trading_client::TradingClientConfig, - ) -> trading_client::TradingClient { - trading_client::TradingClient::new(config) - } - - /// Create a backtesting client - pub const fn create_backtesting_client( - &self, - config: backtesting_client::BacktestingClientConfig, - ) -> backtesting_client::BacktestingClient { - backtesting_client::BacktestingClient::new(config) - } - - /// Create an ML training client - pub const fn create_ml_training_client( - &self, - config: ml_training_client::MLTrainingClientConfig, - ) -> ml_training_client::MLTrainingClient { - ml_training_client::MLTrainingClient::new(config) - } - - /// Add a service connection to the pool - pub async fn add_service( - &self, - service_name: String, - config: connection_manager::ConnectionConfig, - ) -> crate::error::TliResult<()> { - self.connection_manager - .add_service(service_name, config) - .await - .map_err(crate::error::TliError::Connection) - } - - /// Get connection statistics for all services - pub async fn get_connection_stats( - &self, - ) -> std::collections::HashMap> { - self.connection_manager.get_pool_stats().await - } - - /// Shutdown all connections and clients - pub async fn shutdown(&self) { - self.connection_manager.shutdown().await; - } -} - -/// Convenience builder for creating a complete TLI client setup -#[derive(Debug)] -pub struct TliClientBuilder { - /// Connection configuration - connection_config: connection_manager::ConnectionConfig, - /// Service endpoints - service_endpoints: std::collections::HashMap, - /// Client configurations - trading_config: Option, - backtesting_config: Option, - ml_training_config: Option, -} - -impl Default for TliClientBuilder { - fn default() -> Self { - Self::new() - } -} - -impl TliClientBuilder { - /// Create a new builder - pub fn new() -> Self { - Self { - connection_config: connection_manager::ConnectionConfig::default(), - service_endpoints: std::collections::HashMap::new(), - trading_config: None, - backtesting_config: None, - ml_training_config: None, - } - } - - /// Set connection configuration - pub fn with_connection_config(mut self, config: connection_manager::ConnectionConfig) -> Self { - self.connection_config = config; - self - } - - /// Add a service endpoint - pub fn with_service_endpoint(mut self, service_name: String, endpoint: String) -> Self { - self.service_endpoints.insert(service_name, endpoint); - self - } - - /// Set trading client configuration - pub fn with_trading_config(mut self, config: trading_client::TradingClientConfig) -> Self { - self.trading_config = Some(config); - self - } - - /// Set backtesting client configuration - pub fn with_backtesting_config( - mut self, - config: backtesting_client::BacktestingClientConfig, - ) -> Self { - self.backtesting_config = Some(config); - self - } - - /// Set ML training client configuration - pub fn with_ml_training_config( - mut self, - config: ml_training_client::MLTrainingClientConfig, - ) -> Self { - self.ml_training_config = Some(config); - self - } - - /// Build the complete TLI client setup - pub async fn build(self) -> crate::error::TliResult { - let factory = ClientFactory::new(self.connection_config.clone()); - - // Add service connections - for (service_name, endpoint) in self.service_endpoints { - let mut service_config = self.connection_config.clone(); - service_config.server_url = endpoint; - factory.add_service(service_name, service_config).await?; - } - - // Create clients - let trading_client = self - .trading_config - .map(|config| factory.create_trading_client(config)); - - let backtesting_client = self - .backtesting_config - .map(|config| factory.create_backtesting_client(config)); - - let ml_training_client = self - .ml_training_config - .map(|config| factory.create_ml_training_client(config)); - - Ok(TliClientSuite { - factory, - trading_client, - backtesting_client, - ml_training_client, - }) - } -} - -/// Complete TLI client suite with all service clients -#[derive(Debug)] -pub struct TliClientSuite { - /// Client factory - pub factory: ClientFactory, - /// Trading client (includes all operations: trading, risk, monitoring, config, system status) - pub trading_client: Option, - /// Backtesting client - pub backtesting_client: Option, - /// ML training client - pub ml_training_client: Option, -} - -impl TliClientSuite { - /// Get connection statistics for all services - pub async fn get_connection_stats( - &self, - ) -> std::collections::HashMap> { - self.factory.get_connection_stats().await - } - - /// Shutdown all clients and connections - pub async fn shutdown(self) { - if let Some(mut client) = self.trading_client { - client.shutdown().await; - } - if let Some(mut client) = self.backtesting_client { - client.shutdown().await; - } - if let Some(mut client) = self.ml_training_client { - client.shutdown().await; - } - - self.factory.shutdown().await; - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_client_factory_creation() { - let config = connection_manager::ConnectionConfig::default(); - let factory = ClientFactory::new(config); - - // Test that factory can create clients - let trading_config = trading_client::TradingClientConfig::default(); - let _trading_client = factory.create_trading_client(trading_config); - } - - #[test] - fn test_builder_pattern() { - let builder = TliClientBuilder::new() - .with_service_endpoint( - "trading_service".to_owned(), - "http://localhost:50051".to_owned(), - ) - .with_trading_config(trading_client::TradingClientConfig::default()) - .with_backtesting_config(backtesting_client::BacktestingClientConfig::default()) - .with_ml_training_config(ml_training_client::MLTrainingClientConfig::default()); - - // Builder should have the configuration set - assert!(builder.trading_config.is_some()); - assert!(builder.backtesting_config.is_some()); - assert!(builder.ml_training_config.is_some()); - assert!(builder.service_endpoints.contains_key("trading_service")); - } -} diff --git a/bin/fxt/src/client/trading_client.rs b/bin/fxt/src/client/trading_client.rs deleted file mode 100644 index b948239a7..000000000 --- a/bin/fxt/src/client/trading_client.rs +++ /dev/null @@ -1,196 +0,0 @@ -//! Trading service gRPC client -//! -//! Provides a gRPC client for communicating with the trading service, -//! including order management, position tracking, and system monitoring. - -use anyhow::{Context, Result as AnyhowResult}; -use serde::{Deserialize, Serialize}; -use tonic::transport::Channel; - -/// Configuration for the trading service gRPC client -/// -/// Contains connection parameters and client behavior settings for -/// communicating with the trading service. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TradingClientConfig { - /// gRPC endpoint URL for the trading service (MUST be https://) - pub endpoint: String, - /// Request timeout in milliseconds - pub timeout_ms: u64, -} - -impl Default for TradingClientConfig { - /// Create default trading client configuration - /// - /// Returns configuration with API Gateway HTTPS endpoint and 30-second timeout. - /// - /// # Security - /// - /// Defaults to HTTPS (not HTTP) to enforce encrypted connections. - /// - /// # Wave 71 Update - /// - /// Endpoint changed to API Gateway (port 50050) instead of direct trading service. - /// - /// All requests now route through the API Gateway for centralized authentication. - fn default() -> Self { - Self { - endpoint: "https://localhost:50050".to_owned(), // API Gateway endpoint - timeout_ms: 30_000, - } - } -} - -/// gRPC client for the trading service -/// -/// Manages connection to the trading service and provides methods for -/// order submission, position queries, and system status monitoring. -#[derive(Debug)] -pub struct TradingClient { - /// Client configuration - config: TradingClientConfig, - /// Active gRPC channel (None if disconnected) - channel: Option, -} - -impl TradingClient { - /// Create a new trading client with the specified configuration - /// - /// # Arguments - /// * `config` - Client configuration including endpoint and timeout settings - /// - /// # Returns - /// - /// New `TradingClient` instance (not yet connected) - pub const fn new(config: TradingClientConfig) -> Self { - Self { - config, - channel: None, - } - } - - /// Validate URL scheme is HTTPS (rejects HTTP) - /// - /// # Security - /// - /// This enforces fail-closed behavior - only HTTPS connections are allowed. - /// - /// Insecure HTTP connections are rejected with a clear error message. - fn validate_endpoint_security(endpoint: &str) -> AnyhowResult<()> { - if endpoint.starts_with("http://") { - anyhow::bail!( - "SECURITY ERROR: Insecure HTTP endpoint rejected: {}. TLS (https://) is required for all gRPC connections to protect sensitive trading data.", - endpoint - ); - } - - if !endpoint.starts_with("https://") { - anyhow::bail!( - "Invalid endpoint scheme in URL: {}. Only HTTPS is supported (example: https://trading-service:50051).", - endpoint - ); - } - - Ok(()) - } - - /// Establish connection to the trading service - /// - /// Creates a gRPC channel to the configured trading service endpoint. - /// - /// Must be called before making any service requests. - /// - /// # Returns - /// `Result<(), anyhow::Error>` - Ok if connection successful - /// - /// # Errors - /// - /// Returns error if: - /// - Endpoint uses insecure HTTP scheme (security validation failure) - /// - /// - Unable to parse endpoint URL - /// - Unable to connect to the service endpoint - /// - /// # Security - /// - /// This method enforces TLS by rejecting any HTTP endpoints. - /// - /// Use HTTPS endpoints only (e.g., ). - pub async fn connect(&mut self) -> AnyhowResult<()> { - // SECURITY: Validate endpoint uses HTTPS before attempting connection - Self::validate_endpoint_security(&self.config.endpoint) - .context("Trading client endpoint security validation failed")?; - - // Parse endpoint with proper error handling (no unwrap/panic) - let channel = Channel::from_shared(self.config.endpoint.clone()) - .context("Failed to parse trading service endpoint URL - check URL format")? - .connect() - .await - .context("Failed to establish connection to trading service - check network and TLS configuration")?; - - self.channel = Some(channel); - - tracing::info!( - "\u{2705} Trading client connected securely via TLS to {}", - self.config.endpoint - ); - - Ok(()) - } - - /// Check if the client is currently connected to the trading service - /// - /// # Returns - /// `true` if connected, `false` if disconnected - pub const fn is_connected(&self) -> bool { - self.channel.is_some() - } - - /// Shutdown the client and close the connection - /// - /// Cleanly closes the gRPC channel and releases associated resources. - /// - /// The client can be reconnected after shutdown by calling `connect()`. - pub async fn shutdown(&mut self) { - if self.channel.is_some() { - tracing::info!( - "Shutting down trading client connection to {}", - self.config.endpoint - ); - } - self.channel = None; - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_default_uses_https() { - let config = TradingClientConfig::default(); - assert!( - config.endpoint.starts_with("https://"), - "Default config must use HTTPS" - ); - } - - #[test] - fn test_http_validation_rejects_insecure() { - let result = TradingClient::validate_endpoint_security("http://localhost:50051"); - assert!(result.is_err(), "HTTP endpoints should be rejected"); - assert!(result.unwrap_err().to_string().contains("Insecure HTTP")); - } - - #[test] - fn test_https_validation_accepts_secure() { - let result = TradingClient::validate_endpoint_security("https://localhost:50051"); - assert!(result.is_ok(), "HTTPS endpoints should be accepted"); - } - - #[test] - fn test_invalid_scheme_rejected() { - let result = TradingClient::validate_endpoint_security("ftp://localhost:50051"); - assert!(result.is_err(), "Non-HTTPS schemes should be rejected"); - } -} diff --git a/bin/fxt/src/commands/agent.rs b/bin/fxt/src/commands/agent.rs index 6ec60870c..0696648e7 100644 --- a/bin/fxt/src/commands/agent.rs +++ b/bin/fxt/src/commands/agent.rs @@ -1,480 +1,36 @@ -//! TLI Agent Commands -//! -//! Command-line interface for Trading Agent operations. -//! Connects to Trading Agent Service via API Gateway for portfolio allocation, -//! asset selection, and strategy coordination. -//! -//! # Commands -//! - `allocate-portfolio` - Allocate capital across selected assets -//! -//! # Architecture -//! - Pure client implementation (connects ONLY to API Gateway at port 50051) -//! - gRPC communication with `TradingAgentService` via API Gateway proxy -//! - No direct service dependencies (proper microservice architecture) +//! `fxt agent` -- trading agent control. -use anyhow::{Context, Result}; -use clap::{Args, Subcommand}; -use colored::Colorize; -// Channel creation via crate::client::connect_channel -use tonic::Request; +use anyhow::Result; +use clap::{Parser, Subcommand}; -use tonic::metadata::MetadataValue; +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; -use crate::proto::trading_agent::{ - trading_agent_service_client::TradingAgentServiceClient, AllocatePortfolioRequest, - AllocationStrategy, AllocationType, GetSelectedAssetsRequest, RiskConstraints, -}; - -/// Agent command arguments -#[derive(Args, Debug)] -pub struct AgentArgs { +#[derive(Parser, Debug)] +pub struct AgentCommand { #[command(subcommand)] - command: AgentCommand, + action: AgentAction, } -/// Agent subcommands -#[derive(Subcommand, Debug, Clone)] -pub enum AgentCommand { - /// Allocate portfolio capital across assets - #[clap( - long_about = "Allocate capital across selected assets using various strategies.\n\n\ - Strategies:\n\ - - equal-weight: 1/N allocation across all assets\n\ - - risk-parity: Equal risk contribution per asset\n\ - - ml-optimized: ML-based portfolio optimization\n\ - - mean-variance: Mean-variance optimization (Markowitz)\n\ - - kelly: Kelly criterion allocation\n\n\ - Examples:\n\ - tli agent allocate-portfolio --selection-id abc-123 --total-capital 100000\n\ - tli agent allocate-portfolio --selection-id abc-123 --total-capital 100000 --strategy risk-parity" - )] - AllocatePortfolio(AllocatePortfolioArgs), +#[derive(Subcommand, Debug)] +enum AgentAction { + /// Start the trading agent + Start, + /// Stop the trading agent + Stop, + /// Show agent status (strategy, positions, signals) + Status, + /// Show/update agent configuration + Config { + /// Configuration key to get/set + key: Option, + /// Value to set (omit to read current value) + value: Option, + }, } -/// Portfolio allocation arguments (public for testing) -#[derive(Debug, Args, Clone)] -pub struct AllocatePortfolioArgs { - /// Asset selection ID from previous `SelectAssets` call - #[arg(long, required = true)] - pub selection_id: String, - - /// Total capital to allocate (USD) - #[arg(long, required = true)] - pub total_capital: f64, - - /// Allocation strategy - #[arg(long, default_value = "ml-optimized")] - pub strategy: String, - - /// Maximum position size (percentage of portfolio, 0.0-1.0) - #[arg(long, default_value = "0.20")] - pub max_position_size: f64, - - /// Minimum position size (percentage of portfolio, 0.0-1.0) - #[arg(long, default_value = "0.05")] - pub min_position_size: f64, -} - -impl AgentArgs { - /// Execute agent command - /// - /// Routes to appropriate subcommand handler. - /// All commands connect to API Gateway (). - pub async fn execute(&self, api_gateway_url: &str, jwt_token: &str) -> Result<()> { - match &self.command { - AgentCommand::AllocatePortfolio(args) => { - handle_allocate_portfolio(args.clone(), api_gateway_url, jwt_token).await - }, - } - } -} - -/// Parse allocation strategy string to `AllocationType` enum -fn parse_allocation_strategy(strategy: &str) -> Result { - match strategy.to_lowercase().as_str() { - "equal-weight" | "equalweight" => Ok(AllocationType::EqualWeight), - "risk-parity" | "riskparity" => Ok(AllocationType::RiskParity), - "ml-optimized" | "mloptimized" => Ok(AllocationType::MlOptimized), - "mean-variance" | "meanvariance" => Ok(AllocationType::MeanVariance), - "kelly" => Ok(AllocationType::Kelly), - _ => anyhow::bail!( - "Unknown allocation strategy: {}. Valid options: equal-weight, risk-parity, ml-optimized, mean-variance, kelly", - strategy - ), - } -} - -/// Validate portfolio allocation constraints -fn validate_constraints(args: &AllocatePortfolioArgs) -> Result<()> { - // Validate total capital is positive - if args.total_capital <= 0.0 { - anyhow::bail!( - "Total capital must be positive, got: {}", - args.total_capital - ); - } - - // Validate position size constraints (0 < min < max < 1.0) - if args.min_position_size <= 0.0 || args.min_position_size >= 1.0 { - anyhow::bail!( - "Minimum position size must be between 0.0 and 1.0, got: {}", - args.min_position_size - ); - } - - if args.max_position_size <= 0.0 || args.max_position_size > 1.0 { - anyhow::bail!( - "Maximum position size must be between 0.0 and 1.0, got: {}", - args.max_position_size - ); - } - - if args.min_position_size >= args.max_position_size { - anyhow::bail!( - "Minimum position size ({}) must be less than maximum position size ({})", - args.min_position_size, - args.max_position_size - ); - } - - Ok(()) -} - -/// Handle allocate-portfolio command (public for testing) -/// -/// # Arguments -/// * `args` - Portfolio allocation arguments -/// * `api_gateway_url` - API Gateway URL -/// * `jwt_token` - JWT authentication token -/// -/// # Production Implementation -/// Connects to Trading Agent Service via API Gateway and requests portfolio allocation. -pub async fn handle_allocate_portfolio( - args: AllocatePortfolioArgs, - api_gateway_url: &str, - jwt_token: &str, -) -> Result<()> { - // Validate constraints - validate_constraints(&args).context("Invalid portfolio allocation constraints")?; - - // Parse allocation strategy - let allocation_type = - parse_allocation_strategy(&args.strategy).context("Failed to parse allocation strategy")?; - - println!("{}", "\u{1f4ca} Allocating Portfolio...".bold()); - println!( - "Strategy: {} | Total Capital: ${:.2}", - args.strategy.bright_magenta(), - args.total_capital - ); - println!( - "Position Size Range: {:.1}% - {:.1}%", - args.min_position_size * 100.0, - args.max_position_size * 100.0 - ); - println!(); - - // Connect to API Gateway - let channel = crate::client::connect_channel(api_gateway_url).await?; - - let mut client = TradingAgentServiceClient::new(channel); - - // Fetch real assets from selection - let mut get_assets_request = Request::new(GetSelectedAssetsRequest { - universe_id: Some(args.selection_id.clone()), - }); - let get_token = MetadataValue::try_from(format!("Bearer {}", jwt_token)) - .context("Invalid JWT token format")?; - get_assets_request - .metadata_mut() - .insert("authorization", get_token); - - let assets_response = client - .get_selected_assets(get_assets_request) - .await - .context("Failed to fetch selected assets")? - .into_inner(); - - let assets = assets_response.assets; - - // Create allocation request - let request = AllocatePortfolioRequest { - assets, - strategy: Some(AllocationStrategy { - allocation_type: allocation_type as i32, - parameters: std::collections::HashMap::new(), - }), - risk_constraints: Some(RiskConstraints { - max_position_size_pct: args.max_position_size, - max_sector_exposure_pct: 0.40, - max_volatility: 0.25, - max_var_95: 0.05, - max_leverage: 1.0, - }), - total_capital: args.total_capital, - }; - - // Add JWT token to metadata - let mut grpc_request = Request::new(request); - if let Ok(auth_value) = format!("Bearer {}", jwt_token).parse() { - grpc_request.metadata_mut().insert("authorization", auth_value); - } - - // Call AllocatePortfolio RPC - let response = client - .allocate_portfolio(grpc_request) - .await - .context("Failed to allocate portfolio")? - .into_inner(); - - // Display allocation results - println!( - "{}", - format!("Portfolio Allocation (ID: {})", response.allocation_id) - .green() - .bold() - ); - println!( - "Strategy: {} | Total Capital: ${}", - args.strategy, args.total_capital - ); - println!(); - - // Display allocation table - println!("\u{250c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{252c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{252c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{252c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2510}"); - println!( - "\u{2502} {:<10} \u{2502} {:<8} \u{2502} {:<12} \u{2502} {:<15} \u{2502}", - "Symbol".bold(), - "Weight".bold(), - "Capital".bold(), - "Position Size".bold() - ); - println!("\u{251c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{253c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{253c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{253c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2524}"); - - for allocation in &response.allocations { - println!( - "\u{2502} {:<10} \u{2502} {:>7.1}% \u{2502} ${:>10.2} \u{2502} {:>12.0} contracts\u{2502}", - allocation.symbol, - allocation.target_weight * 100.0, - allocation.target_capital, - allocation.target_quantity - ); - } - - println!("\u{2514}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2534}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2534}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2534}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2518}"); - println!(); - - // Display risk metrics - if let Some(metrics) = response.metrics { - println!("{}", "Risk Metrics:".bold()); - println!( - " Portfolio Volatility: {:.1}%", - metrics.portfolio_volatility * 100.0 - ); - println!(" Sharpe Ratio: {:.2}", metrics.portfolio_sharpe); - println!( - " Max Drawdown: {:.1}%", - metrics.max_drawdown_estimate * 100.0 - ); - println!(" VaR (95%): {:.1}%", metrics.var_95 * 100.0); - } - - Ok(()) -} - -/// Execute agent command (public interface for main.rs) -/// -/// # Arguments -/// * `args` - Agent command arguments -/// * `api_gateway_url` - API Gateway URL -/// * `jwt_token` - JWT authentication token -pub async fn execute_agent_command( - args: AgentArgs, - api_gateway_url: &str, - jwt_token: &str, -) -> Result<()> { - args.execute(api_gateway_url, jwt_token).await -} - -#[cfg(test)] -#[allow(clippy::unwrap_used, clippy::expect_used)] -mod tests { - use super::*; - - #[test] - fn test_parse_allocation_strategy_equal_weight() { - let result = parse_allocation_strategy("equal-weight"); - assert!(result.is_ok()); - assert_eq!(result.unwrap() as i32, AllocationType::EqualWeight as i32); - } - - #[test] - fn test_parse_allocation_strategy_risk_parity() { - let result = parse_allocation_strategy("risk-parity"); - assert!(result.is_ok()); - assert_eq!(result.unwrap() as i32, AllocationType::RiskParity as i32); - } - - #[test] - fn test_parse_allocation_strategy_ml_optimized() { - let result = parse_allocation_strategy("ml-optimized"); - assert!(result.is_ok()); - assert_eq!(result.unwrap() as i32, AllocationType::MlOptimized as i32); - } - - #[test] - fn test_parse_allocation_strategy_mean_variance() { - let result = parse_allocation_strategy("mean-variance"); - assert!(result.is_ok()); - assert_eq!(result.unwrap() as i32, AllocationType::MeanVariance as i32); - } - - #[test] - fn test_parse_allocation_strategy_kelly() { - let result = parse_allocation_strategy("kelly"); - assert!(result.is_ok()); - assert_eq!(result.unwrap() as i32, AllocationType::Kelly as i32); - } - - #[test] - fn test_parse_allocation_strategy_case_insensitive() { - parse_allocation_strategy("EQUAL-WEIGHT").unwrap(); - parse_allocation_strategy("Risk-Parity").unwrap(); - parse_allocation_strategy("ML-OPTIMIZED").unwrap(); - } - - #[test] - fn test_parse_allocation_strategy_invalid() { - let result = parse_allocation_strategy("invalid-strategy"); - assert!(result.is_err()); - assert!(result - .unwrap_err() - .to_string() - .contains("Unknown allocation strategy")); - } - - #[test] - fn test_validate_constraints_valid() { - let args = AllocatePortfolioArgs { - selection_id: "test-123".to_owned(), - total_capital: 100000.0, - strategy: "ml-optimized".to_owned(), - max_position_size: 0.20, - min_position_size: 0.05, - }; - - let result = validate_constraints(&args); - result.unwrap(); - } - - #[test] - fn test_validate_constraints_negative_capital() { - let args = AllocatePortfolioArgs { - selection_id: "test-123".to_owned(), - total_capital: -1000.0, - strategy: "ml-optimized".to_owned(), - max_position_size: 0.20, - min_position_size: 0.05, - }; - - let result = validate_constraints(&args); - assert!(result.is_err()); - assert!(result.unwrap_err().to_string().contains("positive")); - } - - #[test] - fn test_validate_constraints_zero_capital() { - let args = AllocatePortfolioArgs { - selection_id: "test-123".to_owned(), - total_capital: 0.0, - strategy: "ml-optimized".to_owned(), - max_position_size: 0.20, - min_position_size: 0.05, - }; - - let result = validate_constraints(&args); - assert!(result.is_err()); - } - - #[test] - fn test_validate_constraints_min_size_too_small() { - let args = AllocatePortfolioArgs { - selection_id: "test-123".to_owned(), - total_capital: 100000.0, - strategy: "ml-optimized".to_owned(), - max_position_size: 0.20, - min_position_size: 0.0, - }; - - let result = validate_constraints(&args); - assert!(result.is_err()); - assert!(result - .unwrap_err() - .to_string() - .contains("Minimum position size")); - } - - #[test] - fn test_validate_constraints_min_size_too_large() { - let args = AllocatePortfolioArgs { - selection_id: "test-123".to_owned(), - total_capital: 100000.0, - strategy: "ml-optimized".to_owned(), - max_position_size: 0.20, - min_position_size: 1.0, - }; - - let result = validate_constraints(&args); - assert!(result.is_err()); - } - - #[test] - fn test_validate_constraints_max_size_too_large() { - let args = AllocatePortfolioArgs { - selection_id: "test-123".to_owned(), - total_capital: 100000.0, - strategy: "ml-optimized".to_owned(), - max_position_size: 1.5, - min_position_size: 0.05, - }; - - let result = validate_constraints(&args); - assert!(result.is_err()); - assert!(result - .unwrap_err() - .to_string() - .contains("Maximum position size")); - } - - #[test] - fn test_validate_constraints_min_greater_than_max() { - let args = AllocatePortfolioArgs { - selection_id: "test-123".to_owned(), - total_capital: 100000.0, - strategy: "ml-optimized".to_owned(), - max_position_size: 0.10, - min_position_size: 0.20, - }; - - let result = validate_constraints(&args); - assert!(result.is_err()); - assert!(result - .unwrap_err() - .to_string() - .contains("must be less than maximum")); - } - - #[test] - fn test_validate_constraints_min_equals_max() { - let args = AllocatePortfolioArgs { - selection_id: "test-123".to_owned(), - total_capital: 100000.0, - strategy: "ml-optimized".to_owned(), - max_position_size: 0.15, - min_position_size: 0.15, - }; - - let result = validate_constraints(&args); - assert!(result.is_err()); +impl AgentCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("agent command not yet implemented") } } diff --git a/bin/fxt/src/commands/auth.rs b/bin/fxt/src/commands/auth.rs index e222d453d..80b4a39d3 100644 --- a/bin/fxt/src/commands/auth.rs +++ b/bin/fxt/src/commands/auth.rs @@ -1,322 +1,33 @@ -//! Authentication Commands -//! -//! CLI commands for user authentication operations: -//! - Login with username/password (interactive password prompt) -//! - Logout (clear stored credentials) -//! - Status (show authentication status) -//! - Refresh (manually refresh access token) +//! `fxt auth` -- authentication and session management. -use anyhow::{Context, Result}; -use clap::Subcommand; -use colored::Colorize; -use std::io::Write; -use std::time::{SystemTime, UNIX_EPOCH}; +use anyhow::Result; +use clap::{Parser, Subcommand}; -use crate::auth::{ - token_manager::{AuthTokenManager, FileTokenStorage, TokenStorage}, - LoginClient, -}; +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; -/// Authentication subcommands -#[derive(Subcommand, Debug, Clone)] -pub enum AuthCommand { - /// Login to Foxhunt system with username/password +#[derive(Parser, Debug)] +pub struct AuthCommand { + #[command(subcommand)] + action: AuthAction, +} + +#[derive(Subcommand, Debug)] +enum AuthAction { + /// Login with username and password Login { - /// Username for authentication - #[clap(short, long)] + /// Username + #[arg(long)] username: Option, - - /// Password (if not provided, will prompt securely) - #[clap(short, long)] - password: Option, - - /// API Gateway URL - #[clap( - long, - env = "API_GATEWAY_URL", - default_value = "https://api.fxhnt.ai" - )] - api_gateway_url: String, }, - - /// Logout and clear stored credentials + /// Logout and clear stored tokens Logout, - /// Show current authentication status Status, - - /// Refresh access token using refresh token - Refresh { - /// API Gateway URL - #[clap( - long, - env = "API_GATEWAY_URL", - default_value = "https://api.fxhnt.ai" - )] - api_gateway_url: String, - }, } -/// Execute authentication command -pub async fn execute_auth_command(command: AuthCommand) -> Result<()> { - match command { - AuthCommand::Login { - username, - password, - api_gateway_url, - } => execute_login(username, password, &api_gateway_url).await, - AuthCommand::Logout => execute_logout().await, - AuthCommand::Status => execute_status().await, - AuthCommand::Refresh { api_gateway_url } => execute_refresh(&api_gateway_url).await, +impl AuthCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("auth command not yet implemented") } } - -/// Execute login command -async fn execute_login( - username: Option, - password: Option, - api_gateway_url: &str, -) -> Result<()> { - println!("\n{}", "=== Foxhunt TLI Authentication ===".cyan().bold()); - println!(); - - // If username or password not provided, use interactive flow - if username.is_none() || password.is_none() { - return execute_interactive_login(api_gateway_url).await; - } - - // Non-interactive login - both are Some() since we checked is_none() above - let resolved_username = match username { - Some(u) => u, - None => return execute_interactive_login(api_gateway_url).await, - }; - let resolved_password = match password { - Some(p) => p, - None => return execute_interactive_login(api_gateway_url).await, - }; - - println!("{}", "Connecting to API Gateway...".cyan()); - - // Connect to API Gateway - let channel = crate::client::connect_channel(api_gateway_url).await?; - - // Create auth components - let storage = FileTokenStorage::new().context("Failed to initialize token storage")?; - let auth_manager = AuthTokenManager::new(storage); - let login_client = LoginClient::new(channel); - - println!("{}", "Authenticating...".cyan()); - - // Use LoginClient for non-interactive login with provided credentials - login_client - .login_with_credentials(&resolved_username, &resolved_password, &auth_manager) - .await - .context("Authentication failed")?; - - println!(); - println!("{}", "\u{2713} Login successful!".green().bold()); - println!("{}", format!(" User: {}", resolved_username).green()); - println!(); - - Ok(()) -} - -/// Execute interactive login (prompts for credentials) -async fn execute_interactive_login(api_gateway_url: &str) -> Result<()> { - use std::io; - - // Prompt for username - print!("{}", "Username: ".cyan().bold()); - io::stdout().flush()?; - let mut raw_username = String::new(); - io::stdin() - .read_line(&mut raw_username) - .context("Failed to read username")?; - let username = raw_username.trim().to_owned(); - - if username.is_empty() { - anyhow::bail!("Username cannot be empty"); - } - - // Prompt for password (hidden input) - print!("{}", "Password: ".cyan().bold()); - io::stdout().flush()?; - let password = rpassword::read_password().context("Failed to read password")?; - - if password.is_empty() { - anyhow::bail!("Password cannot be empty"); - } - - // Connect to API Gateway - println!(); - println!("{}", "Connecting to API Gateway...".cyan()); - - let channel = crate::client::connect_channel(api_gateway_url).await?; - - // Create auth components - let storage = FileTokenStorage::new().context("Failed to initialize token storage")?; - let auth_manager = AuthTokenManager::new(storage); - let login_client = LoginClient::new(channel); - - println!("{}", "Authenticating...".cyan()); - - // Use LoginClient's interactive login (handles MFA if needed) - login_client - .interactive_login(&auth_manager) - .await - .context("Authentication failed")?; - - println!("{}", format!(" User: {}", username).green()); - println!(); - - Ok(()) -} - -/// Execute logout command -async fn execute_logout() -> Result<()> { - let storage = FileTokenStorage::new().context("Failed to initialize token storage")?; - - // Clear both access and refresh tokens from storage - storage - .clear_access_token() - .await - .context("Failed to clear access token")?; - - storage - .remove_refresh_token() - .await - .context("Failed to clear refresh token")?; - - println!("{}", "\u{2713} Logged out successfully".green().bold()); - println!(" All tokens cleared from storage"); - println!(" Run: {} to login again", "tli auth login".bright_cyan()); - - Ok(()) -} - -/// JWT claims structure for token parsing -#[derive(Debug, serde::Deserialize)] -struct JwtClaims { - exp: u64, - sub: String, -} - -/// Parse JWT claims without signature verification -fn parse_jwt_claims(token: &str) -> Result { - use jsonwebtoken::{decode, Algorithm, DecodingKey, Validation}; - - let mut validation = Validation::new(Algorithm::HS256); - validation.insecure_disable_signature_validation(); - validation.validate_exp = false; - - let token_data = decode::(token, &DecodingKey::from_secret(b"dummy"), &validation)?; - - Ok(token_data.claims) -} - -/// Execute status command -async fn execute_status() -> Result<()> { - println!("{}", "=== Authentication Status ===".cyan().bold()); - println!(); - - // Read tokens directly from file storage - let storage = FileTokenStorage::new().context("Failed to initialize token storage")?; - - // Check for access token in storage - if let Some(token) = storage.get_access_token().await? { - println!("{}", "\u{2713} Authenticated".green().bold()); - println!( - "{}", - format!(" Token: {}...", token.get(..token.len().min(30)).unwrap_or(&token)).green() - ); - - // Try to parse token and show expiry - if let Ok(claims) = parse_jwt_claims(&token) { - // Display username from JWT subject - println!("{}", format!(" User: {}", claims.sub).green()); - - let now = SystemTime::now() - .duration_since(UNIX_EPOCH) - .context("Failed to get current time")? - .as_secs(); - - if claims.exp > now { - let remaining = claims.exp - now; - let minutes = remaining / 60; - let seconds = remaining % 60; - - if remaining > 60 { - println!( - "{}", - format!(" Expires in: {} minutes, {} seconds", minutes, seconds).cyan() - ); - } else { - println!( - "{}", - format!(" Expires in: {} seconds (refresh recommended)", seconds).yellow() - ); - } - } else { - println!("{}", " Status: EXPIRED".red()); - } - } - - // Check for refresh token availability - match storage.get_refresh_token().await { - Ok(Some(_)) => println!("{}", " Refresh token: Available".green()), - Ok(None) => println!("{}", " Refresh token: Not available".yellow()), - Err(_) => println!("{}", " Refresh token: Not available".yellow()), - } - } else { - println!("{}", "\u{2717} Not authenticated".red().bold()); - println!("{}", " Run 'tli auth login' to authenticate".yellow()); - } - - println!(); - Ok(()) -} - -/// Execute refresh command -async fn execute_refresh(api_gateway_url: &str) -> Result<()> { - println!("{}", "Refreshing tokens...".cyan()); - - // Connect to API Gateway - let channel = crate::client::connect_channel(api_gateway_url).await?; - - // Create auth components - let storage = FileTokenStorage::new().context("Failed to initialize token storage")?; - let auth_manager = AuthTokenManager::new(storage); - let login_client = LoginClient::new(channel); - - // Check if refresh token exists - if auth_manager.get_refresh_token().await?.is_none() { - anyhow::bail!("No refresh token available. Please login first with 'tli auth login'"); - } - - // Attempt refresh - login_client - .refresh_tokens(&auth_manager) - .await - .context("Token refresh failed")?; - - println!( - "{}", - "\u{2713} Tokens refreshed successfully".green().bold() - ); - - // Show new expiry - if let Some(time_remaining) = auth_manager.time_until_expiry().await { - let minutes = time_remaining.as_secs() / 60; - let seconds = time_remaining.as_secs() % 60; - println!( - "{}", - format!( - " New token expires in: {} minutes, {} seconds", - minutes, seconds - ) - .cyan() - ); - } - - Ok(()) -} diff --git a/bin/fxt/src/commands/backtest.rs b/bin/fxt/src/commands/backtest.rs new file mode 100644 index 000000000..8ec3f8674 --- /dev/null +++ b/bin/fxt/src/commands/backtest.rs @@ -0,0 +1,48 @@ +//! `fxt backtest` -- backtesting operations. + +use anyhow::Result; +use clap::{Parser, Subcommand}; + +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; + +#[derive(Parser, Debug)] +pub struct BacktestCommand { + #[command(subcommand)] + action: BacktestAction, +} + +#[derive(Subcommand, Debug)] +enum BacktestAction { + /// Run a backtest + Run { + /// Strategy name or model ID + #[arg(long)] + strategy: String, + /// Symbol to backtest + #[arg(long)] + symbol: String, + /// Start date (YYYY-MM-DD) + #[arg(long)] + from: String, + /// End date (YYYY-MM-DD) + #[arg(long)] + to: String, + }, + /// Check backtest status + Status { + /// Backtest job ID + job_id: String, + }, + /// Show backtest results + Results { + /// Backtest job ID + job_id: String, + }, +} + +impl BacktestCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("backtest command not yet implemented") + } +} diff --git a/bin/fxt/src/commands/backtest_ml.rs b/bin/fxt/src/commands/backtest_ml.rs deleted file mode 100644 index dd0cf46f1..000000000 --- a/bin/fxt/src/commands/backtest_ml.rs +++ /dev/null @@ -1,428 +0,0 @@ -//! ML Backtesting Commands for TLI -//! -//! Command-line interface for ML-powered backtesting operations. -//! Connects to `BacktestingService` gRPC endpoint. - -use anyhow::{Context, Result}; -use clap::{Args, Subcommand}; -use colored::Colorize; -use tonic::Request; -use tracing::{debug, error}; - -use crate::proto::trading::{ - backtesting_service_client::BacktestingServiceClient, BacktestStatus, - GetBacktestResultsRequest, GetBacktestStatusRequest, StartBacktestRequest, -}; - -/// ML Backtesting command arguments -#[derive(Args, Debug)] -pub struct BacktestMlArgs { - /// Subcommand to execute - #[clap(subcommand)] - pub command: BacktestMlCommand, - - /// API Gateway URL (override config) - #[clap(long, env = "API_GATEWAY_URL")] - pub api_gateway_url: Option, -} - -/// ML Backtesting subcommands -#[derive(Subcommand, Debug)] -pub enum BacktestMlCommand { - /// Run ML ensemble backtest - Run { - /// Symbol to backtest (e.g., ES.FUT, NQ.FUT) - #[arg(short, long)] - symbol: String, - - /// Start date (YYYY-MM-DD) - #[arg(long)] - start: String, - - /// End date (YYYY-MM-DD) - #[arg(long)] - end: String, - - /// Initial capital - #[arg(short, long, default_value = "100000.0")] - capital: f64, - - /// Confidence threshold (0.0-1.0) - #[arg(short = 't', long, default_value = "0.6")] - threshold: f64, - - /// Use ensemble (all models) or single model - #[arg(long, default_value = "true")] - ensemble: bool, - - /// Specific model name if not using ensemble (DQN, PPO, MAMBA2, TFT) - #[arg(long)] - model: Option, - - /// Compare with rule-based strategy - #[arg(long)] - compare: bool, - - /// Description for this backtest run - #[arg(short, long)] - description: Option, - }, - - /// Get status of running backtest - Status { - /// Backtest ID to check - #[arg(short, long)] - id: String, - }, - - /// Get results of completed backtest - Results { - /// Backtest ID to fetch results for - #[arg(short, long)] - id: String, - - /// Include individual trades in output - #[arg(long)] - trades: bool, - }, -} - -/// Execute ML backtesting command -pub async fn execute_backtest_ml_command(args: BacktestMlArgs) -> Result<()> { - let gateway_url = args - .api_gateway_url - .unwrap_or_else(|| "http://localhost:50051".to_owned()); - - debug!("Connecting to API Gateway at: {}", gateway_url); - - let mut client = BacktestingServiceClient::connect(gateway_url.clone()) - .await - .context("Failed to connect to Backtesting Service")?; - - match args.command { - BacktestMlCommand::Run { - symbol, - start, - end, - capital, - threshold, - ensemble, - model, - compare, - description, - } => { - run_ml_backtest( - &mut client, - symbol, - start, - end, - capital, - threshold, - ensemble, - model, - compare, - description, - ) - .await - }, - BacktestMlCommand::Status { id } => get_backtest_status(&mut client, id).await, - BacktestMlCommand::Results { id, trades } => { - get_backtest_results(&mut client, id, trades).await - }, - } -} - -/// Helper to convert date string to Unix nanos -fn date_to_unix_nanos(date_str: &str) -> Result { - let date = chrono::NaiveDate::parse_from_str(date_str, "%Y-%m-%d") - .context("Invalid date format, use YYYY-MM-DD")? - .and_hms_opt(0, 0, 0) - .context("Failed to create datetime")?; - date.and_utc() - .timestamp_nanos_opt() - .context("Date out of range for nanosecond timestamp") -} - -/// Run ML backtest -async fn run_ml_backtest( - client: &mut BacktestingServiceClient, - symbol: String, - start: String, - end: String, - capital: f64, - threshold: f64, - ensemble: bool, - model: Option, - compare: bool, - description: Option, -) -> Result<()> { - println!("{}", "\u{1f680} Starting ML Backtest".bold().green()); - println!("\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}"); - - let start_nanos = date_to_unix_nanos(&start)?; - let end_nanos = date_to_unix_nanos(&end)?; - - // Build parameters - let mut parameters = vec![ - ("confidence_threshold".to_owned(), threshold.to_string()), - ("use_ensemble".to_owned(), ensemble.to_string()), - ]; - - if let Some(ref model_name) = model { - parameters.push(("model_name".to_owned(), model_name.clone())); - } - - // Run ML backtest - let ml_request = Request::new(StartBacktestRequest { - strategy_name: "MLEnsemble".to_owned(), - symbols: vec![symbol.clone()], - start_date_unix_nanos: start_nanos, - end_date_unix_nanos: end_nanos, - initial_capital: capital, - parameters: parameters.into_iter().collect(), - save_results: true, - description: description - .clone() - .unwrap_or_else(|| "ML backtest via TLI".to_owned()), - }); - - let ml_response = client - .start_backtest(ml_request) - .await - .context("Failed to start ML backtest")?; - let ml_result = ml_response.into_inner(); - - if !ml_result.success { - error!("ML backtest failed to start: {}", ml_result.message); - return Err(anyhow::anyhow!( - "Failed to start backtest: {}", - ml_result.message - )); - } - - let ml_id = ml_result.backtest_id.clone(); - println!("\u{2705} ML Backtest started: {}", ml_id.bright_cyan()); - println!(" Symbol: {}", symbol.bright_yellow()); - println!(" Period: {} to {}", start, end); - println!(" Capital: ${:.2}", capital); - println!(" Threshold: {:.1}%", threshold * 100.0); - println!( - " Mode: {}", - if ensemble { - "Ensemble (All Models)".bright_green() - } else { - format!( - "Single Model ({})", - model.unwrap_or_else(|| "DQN".to_owned()) - ) - .bright_blue() - } - ); - - // If compare flag is set, also run rule-based backtest - if compare { - println!( - "\n{}", - "\u{1f4ca} Running comparison backtest...".bold().cyan() - ); - - let rule_request = Request::new(StartBacktestRequest { - strategy_name: "MovingAverageCrossover".to_owned(), - symbols: vec![symbol.clone()], - start_date_unix_nanos: start_nanos, - end_date_unix_nanos: end_nanos, - initial_capital: capital, - parameters: vec![ - ("fast_period".to_owned(), "10".to_owned()), - ("slow_period".to_owned(), "20".to_owned()), - ] - .into_iter() - .collect(), - save_results: true, - description: "Rule-based comparison backtest".to_owned(), - }); - - let rule_response = client - .start_backtest(rule_request) - .await - .context("Failed to start comparison backtest")?; - let rule_result = rule_response.into_inner(); - - if rule_result.success { - println!( - "\u{2705} Comparison backtest started: {}", - rule_result.backtest_id.bright_cyan() - ); - } - } - - println!( - "\n\u{1f4a1} Use {} to check status", - format!("tli backtest ml status --id {}", ml_id).bright_yellow() - ); - println!( - "\u{1f4a1} Use {} to get results", - format!("tli backtest ml results --id {}", ml_id).bright_yellow() - ); - - Ok(()) -} - -/// Get backtest status -async fn get_backtest_status( - client: &mut BacktestingServiceClient, - id: String, -) -> Result<()> { - let request = Request::new(GetBacktestStatusRequest { - backtest_id: id.clone(), - }); - - let response = client - .get_backtest_status(request) - .await - .context("Failed to get backtest status")?; - let status = response.into_inner(); - - println!("{}", "\u{1f4ca} Backtest Status".bold().green()); - println!("\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}"); - println!("ID: {}", status.backtest_id.bright_cyan()); - println!("Status: {}", format_backtest_status(status.status())); - println!("Progress: {:.1}%", status.progress_percentage); - println!("Current Date: {}", status.current_date); - println!("Trades Executed: {}", status.trades_executed); - println!("Current P&L: ${:.2}", status.current_pnl); - - if let Some(error) = status.error_message { - println!("{}: {}", "Error".bright_red(), error); - } - - Ok(()) -} - -/// Get backtest results -async fn get_backtest_results( - client: &mut BacktestingServiceClient, - id: String, - include_trades: bool, -) -> Result<()> { - let request = Request::new(GetBacktestResultsRequest { - backtest_id: id.clone(), - include_trades, - include_metrics: true, - }); - - let response = client - .get_backtest_results(request) - .await - .context("Failed to get backtest results")?; - let results = response.into_inner(); - - println!("{}", "\u{1f4c8} ML Backtest Results".bold().green()); - println!("\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}"); - - if let Some(metrics) = results.metrics { - println!("\n{}", "Performance Metrics:".bold()); - println!(" Total Return: {:.2}%", metrics.total_return * 100.0); - println!( - " Annualized Return: {:.2}%", - metrics.annualized_return * 100.0 - ); - println!(" Sharpe Ratio: {:.2}", metrics.sharpe_ratio); - println!(" Sortino Ratio: {:.2}", metrics.sortino_ratio); - println!(" Max Drawdown: {:.2}%", metrics.max_drawdown * 100.0); - println!(" Calmar Ratio: {:.2}", metrics.calmar_ratio); - - println!("\n{}", "Trade Statistics:".bold()); - println!(" Total Trades: {}", metrics.total_trades); - println!( - " Winning Trades: {} ({:.1}%)", - metrics.winning_trades, - metrics.win_rate * 100.0 - ); - println!(" Losing Trades: {}", metrics.losing_trades); - println!(" Profit Factor: {:.2}", metrics.profit_factor); - println!(" Average Win: ${:.2}", metrics.avg_win); - println!(" Average Loss: ${:.2}", metrics.avg_loss); - println!(" Largest Win: ${:.2}", metrics.largest_win); - println!(" Largest Loss: ${:.2}", metrics.largest_loss); - - // Highlight target achievements - println!("\n{}", "Target Metrics:".bold()); - if metrics.sharpe_ratio > 1.5 { - println!(" \u{2705} Sharpe Ratio > 1.5 (ACHIEVED)"); - } else { - println!( - " \u{26a0}\u{fe0f} Sharpe Ratio: {:.2} (target: >1.5)", - metrics.sharpe_ratio - ); - } - - if metrics.win_rate > 0.55 { - println!(" \u{2705} Win Rate > 55% (ACHIEVED)"); - } else { - println!( - " \u{26a0}\u{fe0f} Win Rate: {:.1}% (target: >55%)", - metrics.win_rate * 100.0 - ); - } - - if metrics.max_drawdown < 0.20 { - println!(" \u{2705} Max Drawdown < 20% (ACHIEVED)"); - } else { - println!( - " \u{26a0}\u{fe0f} Max Drawdown: {:.1}% (target: <20%)", - metrics.max_drawdown * 100.0 - ); - } - } else { - println!("{}", "No metrics available".bright_red()); - } - - if include_trades && !results.trades.is_empty() { - println!( - "\n{}", - format!("Recent Trades ({} total):", results.trades.len()).bold() - ); - for (i, trade) in results.trades.iter().take(10).enumerate() { - println!( - " {}. {} {} @ ${:.2} \u{2192} ${:.2} = {}", - i + 1, - trade.symbol, - format_order_side(trade.side), - trade.entry_price, - trade.exit_price, - if trade.pnl >= 0.0 { - format!("+${:.2}", trade.pnl).bright_green() - } else { - format!("-${:.2}", trade.pnl.abs()).bright_red() - } - ); - } - if results.trades.len() > 10 { - println!(" ... and {} more trades", results.trades.len() - 10); - } - } - - Ok(()) -} - -/// Format backtest status for display -fn format_backtest_status(status: BacktestStatus) -> colored::ColoredString { - match status { - BacktestStatus::Queued => "QUEUED".bright_yellow(), - BacktestStatus::Running => "RUNNING".bright_cyan(), - BacktestStatus::Completed => "COMPLETED".bright_green(), - BacktestStatus::Failed => "FAILED".bright_red(), - BacktestStatus::Cancelled => "CANCELLED".bright_magenta(), - BacktestStatus::Unspecified | BacktestStatus::Paused => "UNKNOWN".bright_red(), - } -} - -/// Format order side for display -const fn format_order_side(side: i32) -> &'static str { - match side { - 1 => "BUY", - 2 => "SELL", - _ => "UNKNOWN", - } -} diff --git a/bin/fxt/src/commands/broker.rs b/bin/fxt/src/commands/broker.rs index 16a773677..5a03178fc 100644 --- a/bin/fxt/src/commands/broker.rs +++ b/bin/fxt/src/commands/broker.rs @@ -1,517 +1,40 @@ -//! Broker connectivity check command -//! -//! Validates IBKR TWS/Gateway connectivity via: -//! 1. Direct ibapi handshake (feature-gated behind `broker-check`) -//! 2. Broker Gateway gRPC health check +//! `fxt broker` -- broker connectivity. use anyhow::Result; -use clap::{Args, Subcommand}; -use colored::Colorize; -use std::time::Instant; +use clap::{Parser, Subcommand}; -use crate::config::TliConfig; +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; -/// Resolved broker check settings (CLI > env > config > defaults). -#[derive(Debug)] -pub struct ResolvedBrokerConfig { - pub ibkr_host: String, - pub ibkr_port: u16, - pub ibkr_client_id: i32, - pub gateway_url: String, - pub skip_ibkr: bool, - pub skip_gateway: bool, -} - -/// Broker command arguments -#[derive(Args, Debug)] -pub struct BrokerArgs { +#[derive(Parser, Debug)] +pub struct BrokerCommand { #[command(subcommand)] - command: BrokerCommand, + action: BrokerAction, } -/// Broker subcommands -#[derive(Subcommand, Debug, Clone)] -pub enum BrokerCommand { - /// Check broker connectivity (IB Gateway + Broker Gateway gRPC) - #[clap( - long_about = "Validate IBKR TWS/Gateway connectivity.\n\n\ - Runs two checks:\n\ - 1. Direct IB Gateway: TCP connect + ibapi handshake\n\ - 2. Broker Gateway gRPC: HealthCheck + GetSessionStatus RPCs\n\n\ - Config precedence: CLI flags > env vars > ~/.foxhunt/config.toml > defaults\n\n\ - Examples:\n\ - fxt broker check\n\ - fxt broker check --host 10.0.0.5 --port 4002\n\ - fxt broker check --skip-ibkr\n\ - fxt broker check --skip-gateway" - )] - Check(CheckArgs), +#[derive(Subcommand, Debug)] +enum BrokerAction { + /// Show broker connection status + Status, + /// Connect to broker gateway + Connect { + /// Broker host override + #[arg(long)] + host: Option, + /// Broker port override + #[arg(long)] + port: Option, + }, + /// List recent executions + Executions { + /// Maximum number of executions to show + #[arg(long, default_value = "20")] + limit: u32, + }, } -/// Arguments for `fxt broker check` -#[derive(Debug, Args, Clone)] -pub struct CheckArgs { - /// IB Gateway host - #[arg(long, env = "IBKR_HOST")] - pub host: Option, - - /// IB Gateway socat port (4004 = paper, 4003 = live) - #[arg(long, env = "IBKR_PORT")] - pub port: Option, - - /// TWS client ID (must be unique per connection) - #[arg(long, env = "IBKR_CLIENT_ID")] - pub client_id: Option, - - /// Broker Gateway gRPC URL - #[arg(long, env = "BROKER_GATEWAY_URL")] - pub gateway_url: Option, - - /// Skip direct IB Gateway check - #[arg(long, default_value_t = false)] - pub skip_ibkr: bool, - - /// Skip Broker Gateway gRPC check - #[arg(long, default_value_t = false)] - pub skip_gateway: bool, -} - -/// Resolve config: CLI flags > env vars (handled by clap) > config file > defaults. -pub fn resolve_config(args: &CheckArgs, config: &TliConfig) -> ResolvedBrokerConfig { - ResolvedBrokerConfig { - ibkr_host: args - .host - .clone() - .unwrap_or_else(|| config.broker.ibkr.host.clone()), - ibkr_port: args.port.unwrap_or(config.broker.ibkr.port), - ibkr_client_id: args.client_id.unwrap_or(config.broker.ibkr.client_id), - gateway_url: args - .gateway_url - .clone() - .unwrap_or_else(|| config.broker.gateway_url.clone()), - skip_ibkr: args.skip_ibkr, - skip_gateway: args.skip_gateway, - } -} - -/// Single check result -#[derive(Debug)] -pub struct CheckResult { - pub name: String, - pub passed: bool, - pub detail: String, - pub duration_ms: u128, -} - -/// Print a check result line -fn print_check(result: &CheckResult) { - let status = if result.passed { - "\u{2713}".green().bold() - } else { - "\u{2717}".red().bold() - }; - let dots = ".".repeat(24_usize.saturating_sub(result.name.len())); - println!( - " {} {} {} {}", - result.name, dots, status, result.detail - ); -} - -// ── IB Gateway direct check ────────────────────────────────────────── - -/// TCP connect probe (fail-fast before ibapi handshake). -async fn check_tcp_connect(host: &str, port: u16) -> CheckResult { - let addr = format!("{host}:{port}"); - let start = Instant::now(); - match tokio::time::timeout( - std::time::Duration::from_secs(5), - tokio::net::TcpStream::connect(&addr), - ) - .await - { - Ok(Ok(_)) => CheckResult { - name: "TCP connect".to_owned(), - passed: true, - detail: format!("{addr} ({}ms)", start.elapsed().as_millis()), - duration_ms: start.elapsed().as_millis(), - }, - Ok(Err(e)) => CheckResult { - name: "TCP connect".to_owned(), - passed: false, - detail: format!("{e} ({addr})"), - duration_ms: start.elapsed().as_millis(), - }, - Err(_) => CheckResult { - name: "TCP connect".to_owned(), - passed: false, - detail: format!("Timeout after 5s ({addr})"), - duration_ms: start.elapsed().as_millis(), - }, - } -} - -/// Full ibapi handshake (blocking, runs on spawn_blocking). -/// Retries once after 1s if the first attempt fails (IB Gateway may need time -/// to release a stale client_id slot after a previous disconnection). -#[cfg(feature = "broker-check")] -async fn check_ibapi_handshake(host: &str, port: u16, client_id: i32) -> CheckResult { - let addr = format!("{host}:{port}"); - let start = Instant::now(); - - let a1 = addr.clone(); - let first = tokio::task::spawn_blocking(move || ibapi::Client::connect(&a1, client_id)).await; - - match first { - Ok(Ok(client)) => { - let sv = client.server_version(); - return CheckResult { - name: "ibapi handshake".to_owned(), - passed: true, - detail: format!( - "client_id={client_id}, server v{sv} ({}ms)", - start.elapsed().as_millis() - ), - duration_ms: start.elapsed().as_millis(), - }; - } - Ok(Err(_)) => { - // Retry once — IB Gateway may need time to release stale client_id slot - tokio::time::sleep(std::time::Duration::from_secs(1)).await; - } - Err(e) => { - return CheckResult { - name: "ibapi handshake".to_owned(), - passed: false, - detail: format!("Task join error: {e}"), - duration_ms: start.elapsed().as_millis(), - }; - } - } - - let a2 = addr; - match tokio::task::spawn_blocking(move || ibapi::Client::connect(&a2, client_id)).await { - Ok(Ok(client)) => { - let sv = client.server_version(); - CheckResult { - name: "ibapi handshake".to_owned(), - passed: true, - detail: format!( - "client_id={client_id}, server v{sv} ({}ms, retry)", - start.elapsed().as_millis() - ), - duration_ms: start.elapsed().as_millis(), - } - } - Ok(Err(e)) => CheckResult { - name: "ibapi handshake".to_owned(), - passed: false, - detail: format!("{e}"), - duration_ms: start.elapsed().as_millis(), - }, - Err(e) => CheckResult { - name: "ibapi handshake".to_owned(), - passed: false, - detail: format!("Task join error: {e}"), - duration_ms: start.elapsed().as_millis(), - }, - } -} - -#[cfg(not(feature = "broker-check"))] -async fn check_ibapi_handshake(_host: &str, _port: u16, _client_id: i32) -> CheckResult { - CheckResult { - name: "ibapi handshake".to_owned(), - passed: false, - detail: "Skipped (compile with --features broker-check)".to_owned(), - duration_ms: 0, - } -} - -// ── Broker Gateway gRPC check ──────────────────────────────────────── - -/// gRPC HealthCheck against broker_gateway_service. -async fn check_grpc_health(gateway_url: &str) -> CheckResult { - use crate::proto::broker_gateway::broker_gateway_service_client::BrokerGatewayServiceClient; - use crate::proto::broker_gateway::HealthCheckRequest; - - let start = Instant::now(); - let channel = match tonic::transport::Channel::from_shared(gateway_url.to_owned()) { - Ok(endpoint) => { - match tokio::time::timeout(std::time::Duration::from_secs(5), endpoint.connect()).await - { - Ok(Ok(ch)) => ch, - Ok(Err(e)) => { - return CheckResult { - name: "Health check".to_owned(), - passed: false, - detail: format!("Connection failed: {e}"), - duration_ms: start.elapsed().as_millis(), - }; - } - Err(_) => { - return CheckResult { - name: "Health check".to_owned(), - passed: false, - detail: format!("Timeout after 5s ({gateway_url})"), - duration_ms: start.elapsed().as_millis(), - }; - } - } - } - Err(e) => { - return CheckResult { - name: "Health check".to_owned(), - passed: false, - detail: format!("Invalid URL: {e}"), - duration_ms: start.elapsed().as_millis(), - }; - } - }; - - let mut client = BrokerGatewayServiceClient::new(channel); - match client.health_check(HealthCheckRequest {}).await { - Ok(resp) => { - let inner = resp.into_inner(); - CheckResult { - name: "Health check".to_owned(), - passed: inner.healthy, - detail: format!( - "{} ({}ms)", - inner.message, - start.elapsed().as_millis() - ), - duration_ms: start.elapsed().as_millis(), - } - } - Err(e) => CheckResult { - name: "Health check".to_owned(), - passed: false, - detail: format!("{} ({})", e.code(), e.message()), - duration_ms: start.elapsed().as_millis(), - }, - } -} - -/// gRPC GetSessionStatus against broker_gateway_service. -async fn check_grpc_session(gateway_url: &str) -> CheckResult { - use crate::proto::broker_gateway::broker_gateway_service_client::BrokerGatewayServiceClient; - use crate::proto::broker_gateway::GetSessionStatusRequest; - - let start = Instant::now(); - let channel = match tonic::transport::Channel::from_shared(gateway_url.to_owned()) { - Ok(endpoint) => match endpoint.connect().await { - Ok(ch) => ch, - Err(e) => { - return CheckResult { - name: "Session status".to_owned(), - passed: false, - detail: format!("Connection failed: {e}"), - duration_ms: start.elapsed().as_millis(), - }; - } - }, - Err(e) => { - return CheckResult { - name: "Session status".to_owned(), - passed: false, - detail: format!("Invalid URL: {e}"), - duration_ms: start.elapsed().as_millis(), - }; - } - }; - - let mut client = BrokerGatewayServiceClient::new(channel); - match client - .get_session_status(GetSessionStatusRequest { session_id: None }) - .await - { - Ok(resp) => { - let inner = resp.into_inner(); - let state_name = match inner.state { - 0 => "DISCONNECTED", - 1 => "CONNECTED", - 2 => "LOGGING_IN", - 3 => "ACTIVE", - 4 => "LOGGING_OUT", - _ => "UNKNOWN", - }; - CheckResult { - name: "Session status".to_owned(), - passed: inner.state >= 1, // CONNECTED or better - detail: format!( - "{state_name} (seq: {}/{}, RTT: {:.1}ms)", - inner.sender_seq_num, inner.target_seq_num, inner.heartbeat_rtt_ms - ), - duration_ms: start.elapsed().as_millis(), - } - } - Err(e) => CheckResult { - name: "Session status".to_owned(), - passed: false, - detail: format!("{} ({})", e.code(), e.message()), - duration_ms: start.elapsed().as_millis(), - }, - } -} - -// ── Orchestrator ───────────────────────────────────────────────────── - -/// Run all broker checks and return overall pass/fail. -pub async fn run_broker_check(resolved: &ResolvedBrokerConfig) -> Result { - println!("{}", "Broker Connectivity Check".bold()); - println!("{}", "\u{2550}".repeat(40)); - - let mut all_passed = true; - - // ── IB Gateway (direct) ── - if !resolved.skip_ibkr { - println!("{}", "IB Gateway (direct)".bold()); - let tcp = check_tcp_connect(&resolved.ibkr_host, resolved.ibkr_port).await; - print_check(&tcp); - - if tcp.passed { - let handshake = check_ibapi_handshake( - &resolved.ibkr_host, - resolved.ibkr_port, - resolved.ibkr_client_id, - ) - .await; - print_check(&handshake); - if !handshake.passed { - all_passed = false; - } - } else { - all_passed = false; - } - println!(); - } - - // ── Broker Gateway (gRPC) ── - if !resolved.skip_gateway { - println!("{}", "Broker Gateway (gRPC)".bold()); - let health = check_grpc_health(&resolved.gateway_url).await; - print_check(&health); - - if health.passed { - let session = check_grpc_session(&resolved.gateway_url).await; - print_check(&session); - if !session.passed { - all_passed = false; - } - } else { - all_passed = false; - } - println!(); - } - - // ── Summary ── - if all_passed { - println!("{}", "Result: All checks passed \u{2713}".green().bold()); - } else { - println!("{}", "Result: One or more checks failed \u{2717}".red().bold()); - } - - Ok(all_passed) -} - -/// Execute broker command (public interface for main.rs). -pub async fn execute_broker_command(args: BrokerArgs, config: &TliConfig) -> Result { - match args.command { - BrokerCommand::Check(check_args) => { - let resolved = resolve_config(&check_args, config); - run_broker_check(&resolved).await - } - } -} - -#[cfg(test)] -#[allow(clippy::unwrap_used, clippy::expect_used)] -mod tests { - use super::*; - use crate::config::TliConfig; - - #[test] - fn test_resolve_config_uses_defaults() { - let args = CheckArgs { - host: None, - port: None, - client_id: None, - gateway_url: None, - skip_ibkr: false, - skip_gateway: false, - }; - let config = TliConfig::default(); - let resolved = resolve_config(&args, &config); - - assert_eq!(resolved.ibkr_host, "127.0.0.1"); - assert_eq!(resolved.ibkr_port, 4004); - assert_eq!(resolved.ibkr_client_id, 99); - assert_eq!(resolved.gateway_url, "http://localhost:50056"); - assert!(!resolved.skip_ibkr); - assert!(!resolved.skip_gateway); - } - - #[test] - fn test_resolve_config_cli_overrides() { - let args = CheckArgs { - host: Some("10.0.0.5".to_owned()), - port: Some(4001), - client_id: Some(7), - gateway_url: Some("http://custom:9090".to_owned()), - skip_ibkr: true, - skip_gateway: false, - }; - let config = TliConfig::default(); - let resolved = resolve_config(&args, &config); - - assert_eq!(resolved.ibkr_host, "10.0.0.5"); - assert_eq!(resolved.ibkr_port, 4001); - assert_eq!(resolved.ibkr_client_id, 7); - assert_eq!(resolved.gateway_url, "http://custom:9090"); - assert!(resolved.skip_ibkr); - } - - #[test] - fn test_resolve_config_file_overrides_defaults() { - let args = CheckArgs { - host: None, - port: None, - client_id: None, - gateway_url: None, - skip_ibkr: false, - skip_gateway: false, - }; - let config: TliConfig = toml::from_str( - r#" - [broker] - gateway_url = "http://broker-gw:50060" - [broker.ibkr] - host = "192.168.1.100" - port = 4001 - client_id = 55 - "#, - ) - .unwrap(); - let resolved = resolve_config(&args, &config); - - assert_eq!(resolved.ibkr_host, "192.168.1.100"); - assert_eq!(resolved.ibkr_port, 4001); - assert_eq!(resolved.ibkr_client_id, 55); - assert_eq!(resolved.gateway_url, "http://broker-gw:50060"); - } - - #[tokio::test] - async fn test_tcp_connect_refused() { - // Connect to a port that's almost certainly not listening - let result = check_tcp_connect("127.0.0.1", 19999).await; - assert!(!result.passed); - assert_eq!(result.name, "TCP connect"); - } - - #[tokio::test] - async fn test_grpc_health_unreachable() { - let result = check_grpc_health("http://127.0.0.1:19998").await; - assert!(!result.passed); - assert_eq!(result.name, "Health check"); +impl BrokerCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("broker command not yet implemented") } } diff --git a/bin/fxt/src/commands/cluster.rs b/bin/fxt/src/commands/cluster.rs new file mode 100644 index 000000000..9cf1ed323 --- /dev/null +++ b/bin/fxt/src/commands/cluster.rs @@ -0,0 +1,33 @@ +//! `fxt cluster` -- cluster operations. + +use anyhow::Result; +use clap::{Parser, Subcommand}; + +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; + +#[derive(Parser, Debug)] +pub struct ClusterCommand { + #[command(subcommand)] + action: ClusterAction, +} + +#[derive(Subcommand, Debug)] +enum ClusterAction { + /// Show cluster resource utilization + Status, + /// Show node and pod resource details + Resources, + /// Show recent cluster events + Events { + /// Maximum number of events to show + #[arg(long, default_value = "50")] + limit: u32, + }, +} + +impl ClusterCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("cluster command not yet implemented") + } +} diff --git a/bin/fxt/src/commands/config_cmd.rs b/bin/fxt/src/commands/config_cmd.rs new file mode 100644 index 000000000..12b704a2b --- /dev/null +++ b/bin/fxt/src/commands/config_cmd.rs @@ -0,0 +1,39 @@ +//! `fxt config` -- system configuration. + +use anyhow::Result; +use clap::{Parser, Subcommand}; + +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; + +#[derive(Parser, Debug)] +pub struct ConfigCommand { + #[command(subcommand)] + action: ConfigAction, +} + +#[derive(Subcommand, Debug)] +enum ConfigAction { + /// Get a configuration value + Get { + /// Configuration key + key: String, + }, + /// Set a configuration value + Set { + /// Configuration key + key: String, + /// Configuration value + value: String, + }, + /// Export full configuration as TOML + Export, + /// Show resolved environment (merged CLI + env + config file) + Env, +} + +impl ConfigCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("config command not yet implemented") + } +} diff --git a/bin/fxt/src/commands/data.rs b/bin/fxt/src/commands/data.rs new file mode 100644 index 000000000..e9f513a15 --- /dev/null +++ b/bin/fxt/src/commands/data.rs @@ -0,0 +1,44 @@ +//! `fxt data` -- data pipeline management. + +use anyhow::Result; +use clap::{Parser, Subcommand}; + +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; + +#[derive(Parser, Debug)] +pub struct DataCommand { + #[command(subcommand)] + action: DataAction, +} + +#[derive(Subcommand, Debug)] +enum DataAction { + /// Download market data + Download { + /// Symbol to download + #[arg(long)] + symbol: String, + /// Start date (YYYY-MM-DD) + #[arg(long)] + from: String, + /// End date (YYYY-MM-DD) + #[arg(long)] + to: String, + /// Schema (ohlcv-1m, ohlcv-1s, trades, mbp-1) + #[arg(long, default_value = "ohlcv-1m")] + schema: String, + }, + /// Show data pipeline status + Status, + /// Show data cache inventory + Cache, + /// Show live data feed status + Feeds, +} + +impl DataCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("data command not yet implemented") + } +} diff --git a/bin/fxt/src/commands/mcp.rs b/bin/fxt/src/commands/mcp.rs new file mode 100644 index 000000000..6179ddb0e --- /dev/null +++ b/bin/fxt/src/commands/mcp.rs @@ -0,0 +1,18 @@ +//! `fxt mcp` -- MCP (Model Context Protocol) server mode. + +use anyhow::Result; +use clap::Parser; + +#[derive(Parser, Debug)] +pub struct McpCommand; + +impl McpCommand { + /// Start an MCP server that exposes Foxhunt operations as tools + /// for LLM agents (e.g. Claude, Cursor). + /// + /// This command does not need a pre-built gRPC client -- + /// the server manages its own connections. + pub async fn execute(&self) -> Result<()> { + anyhow::bail!("mcp server not yet implemented") + } +} diff --git a/bin/fxt/src/commands/mod.rs b/bin/fxt/src/commands/mod.rs index 2475dbc9b..a993c32e9 100644 --- a/bin/fxt/src/commands/mod.rs +++ b/bin/fxt/src/commands/mod.rs @@ -1,38 +1,17 @@ -//! TLI Command Modules -//! -//! Command-line interface subcommands for TLI operations. -//! Each module implements a specific command category with rich terminal output. -//! -//! # Available Commands -//! - `tune` - Hyperparameter tuning job management (start, status, best, stop) -//! - `train` - Training job management (list, status, details) -//! - `model` - Model promotion management (list, approve, reject) -//! - `agent` - Trading agent operations (status, performance) -//! -//! # Future Commands (Planned) -//! - `backtest` - Backtesting operations -//! - `trading` - Trading operations (orders, positions) -//! - `risk` - Risk management queries -//! - `config` - Configuration management +//! CLI command modules -- one file per top-level subcommand. pub mod agent; pub mod auth; -pub mod backtest_ml; +pub mod backtest; pub mod broker; +pub mod cluster; +pub mod config_cmd; +pub mod data; +pub mod mcp; pub mod model; +pub mod risk; +pub mod service; pub mod trade; -pub mod trade_ml; pub mod train; pub mod tune; pub mod watch; -pub mod tune_stream; - -pub use agent::{execute_agent_command, AgentArgs}; -pub use auth::{execute_auth_command, AuthCommand}; -pub use broker::{execute_broker_command, BrokerArgs}; -pub use backtest_ml::{execute_backtest_ml_command, BacktestMlArgs, BacktestMlCommand}; -pub use model::{execute_model_command, ModelCommand}; -pub use trade::{execute_trade_command, TradeArgs}; -pub use trade_ml::{execute_trade_ml_command, TradeMlArgs}; -pub use train::{execute_train_command, TrainCommand}; -pub use tune::{execute_tune_command, TuneCommand}; diff --git a/bin/fxt/src/commands/model.rs b/bin/fxt/src/commands/model.rs new file mode 100644 index 000000000..7a2fbe23c --- /dev/null +++ b/bin/fxt/src/commands/model.rs @@ -0,0 +1,55 @@ +//! `fxt model` -- model management and promotion. + +use anyhow::Result; +use clap::{Parser, Subcommand}; + +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; + +#[derive(Parser, Debug)] +pub struct ModelCommand { + #[command(subcommand)] + action: ModelAction, +} + +#[derive(Subcommand, Debug)] +enum ModelAction { + /// List registered models + List { + /// Filter by model type + #[arg(long)] + model: Option, + /// Filter by status (staging, production, archived) + #[arg(long)] + status: Option, + }, + /// Show detailed model status + Status { + /// Model ID + model_id: String, + }, + /// Run inference on a model + Predict { + /// Model ID + model_id: String, + /// Symbol to predict + #[arg(long)] + symbol: String, + }, + /// Show ensemble composition and weights + Ensemble, + /// Promote a model to production + Promote { + /// Model ID + model_id: String, + /// Reason for promotion + #[arg(long)] + reason: Option, + }, +} + +impl ModelCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("model command not yet implemented") + } +} diff --git a/bin/fxt/src/commands/model/approve.rs b/bin/fxt/src/commands/model/approve.rs deleted file mode 100644 index bc4b461bd..000000000 --- a/bin/fxt/src/commands/model/approve.rs +++ /dev/null @@ -1,49 +0,0 @@ -//! FXT Model Approve Command - Approve a Pending Model Promotion -//! -//! Approves a model for promotion to production. -//! -//! # Usage -//! -//! ```bash -//! fxt model approve -//! ``` - -use crate::proto::ml_training::{ - ml_training_service_client::MlTrainingServiceClient, ApproveModelRequest, -}; -use anyhow::{Context, Result}; -use tonic::{metadata::MetadataValue, Request}; - -/// Approve a pending model promotion. -/// -/// Calls `ApproveModel` on the ML training service via the API gateway -/// to promote the model to production. -pub async fn run(api_gateway_url: &str, jwt_token: &str, model_id: &str) -> Result<()> { - let channel = crate::client::connect_channel(api_gateway_url).await?; - - let mut client = MlTrainingServiceClient::new(channel); - - let mut request = Request::new(ApproveModelRequest { - model_id: model_id.to_owned(), - promoted_to: "production".to_owned(), - }); - let token_value = MetadataValue::try_from(format!("Bearer {}", jwt_token)) - .context("Invalid JWT token format")?; - request - .metadata_mut() - .insert("authorization", token_value); - - let response = client - .approve_model(request) - .await - .context("Failed to approve model")? - .into_inner(); - - if response.success { - println!("Model '{}' approved for production.", model_id); - } else { - println!("Failed to approve model '{}': {}", model_id, response.message); - } - - Ok(()) -} diff --git a/bin/fxt/src/commands/model/list.rs b/bin/fxt/src/commands/model/list.rs deleted file mode 100644 index 15ca9518c..000000000 --- a/bin/fxt/src/commands/model/list.rs +++ /dev/null @@ -1,72 +0,0 @@ -//! FXT Model List Command - Display Available Models for Training -//! -//! Lists all available model types that can be trained. -//! -//! # Usage -//! -//! ```bash -//! fxt model list -//! ``` - -use crate::proto::ml_training::{ - ml_training_service_client::MlTrainingServiceClient, ListAvailableModelsRequest, -}; -use anyhow::{Context, Result}; -use tonic::{metadata::MetadataValue, Request}; - -/// List available models for training. -/// -/// Calls `ListAvailableModels` on the ML training service via the API gateway -/// and displays a formatted table of available model types. -pub async fn run(api_gateway_url: &str, jwt_token: &str) -> Result<()> { - let channel = crate::client::connect_channel(api_gateway_url).await?; - - let mut client = MlTrainingServiceClient::new(channel); - - let mut request = Request::new(ListAvailableModelsRequest {}); - let token_value = MetadataValue::try_from(format!("Bearer {}", jwt_token)) - .context("Invalid JWT token format")?; - request - .metadata_mut() - .insert("authorization", token_value); - - let response = client - .list_available_models(request) - .await - .context("Failed to list available models")? - .into_inner(); - - let models = response.models; - - if models.is_empty() { - println!("No available models found."); - return Ok(()); - } - - println!("Available Models for Training"); - println!("{:-<80}", ""); - println!( - "{:<15} {:<35} {:<8} {:<10}", - "Model Type", "Description", "GPU", "Est. Time" - ); - println!("{:-<80}", ""); - - for model in &models { - let gpu_label = if model.requires_gpu { "Yes" } else { "No" }; - let time_label = if model.estimated_training_time_minutes > 0 { - format!("{}m", model.estimated_training_time_minutes) - } else { - "N/A".to_owned() - }; - - println!( - "{:<15} {:<35} {:<8} {:<10}", - model.model_type, model.description, gpu_label, time_label - ); - } - - println!("{:-<80}", ""); - println!("Total: {} model(s)", models.len()); - - Ok(()) -} diff --git a/bin/fxt/src/commands/model/mod.rs b/bin/fxt/src/commands/model/mod.rs deleted file mode 100644 index 9bba47984..000000000 --- a/bin/fxt/src/commands/model/mod.rs +++ /dev/null @@ -1,52 +0,0 @@ -//! FXT Model Command Module -//! -//! Provides CLI interface for managing ML model promotions. -//! -//! # Subcommands -//! - `list` - List models with pending promotion status -//! - `approve` - Approve a pending model promotion -//! - `reject` - Reject a pending model promotion - -pub mod approve; -pub mod list; -pub mod reject; - -use anyhow::Result; -use clap::Subcommand; - -/// Model management subcommands -#[derive(Subcommand, Debug)] -pub enum ModelCommand { - /// List models with active/pending promotion status - List, - /// Approve a pending model promotion - Approve { - /// Model ID to approve - model_id: String, - }, - /// Reject a pending model promotion - Reject { - /// Model ID to reject - model_id: String, - /// Reason for rejection - #[arg(long, default_value = "operator rejected")] - reason: String, - }, -} - -/// Execute model command -pub async fn execute_model_command( - command: ModelCommand, - api_gateway_url: &str, - jwt_token: &str, -) -> Result<()> { - match command { - ModelCommand::List => list::run(api_gateway_url, jwt_token).await, - ModelCommand::Approve { model_id } => { - approve::run(api_gateway_url, jwt_token, &model_id).await - } - ModelCommand::Reject { model_id, reason } => { - reject::run(api_gateway_url, jwt_token, &model_id, &reason).await - } - } -} diff --git a/bin/fxt/src/commands/model/reject.rs b/bin/fxt/src/commands/model/reject.rs deleted file mode 100644 index d1a6167e5..000000000 --- a/bin/fxt/src/commands/model/reject.rs +++ /dev/null @@ -1,55 +0,0 @@ -//! FXT Model Reject Command - Reject a Pending Model Promotion -//! -//! Rejects a model promotion and records the reason. -//! -//! # Usage -//! -//! ```bash -//! fxt model reject -//! fxt model reject --reason "metrics degraded vs baseline" -//! ``` - -use crate::proto::ml_training::{ - ml_training_service_client::MlTrainingServiceClient, RejectModelRequest, -}; -use anyhow::{Context, Result}; -use tonic::{metadata::MetadataValue, Request}; - -/// Reject a pending model promotion. -/// -/// Calls `RejectModel` on the ML training service via the API gateway -/// to reject the model promotion with an operator-supplied reason. -pub async fn run( - api_gateway_url: &str, - jwt_token: &str, - model_id: &str, - reason: &str, -) -> Result<()> { - let channel = crate::client::connect_channel(api_gateway_url).await?; - - let mut client = MlTrainingServiceClient::new(channel); - - let mut request = Request::new(RejectModelRequest { - model_id: model_id.to_owned(), - reason: reason.to_owned(), - }); - let token_value = MetadataValue::try_from(format!("Bearer {}", jwt_token)) - .context("Invalid JWT token format")?; - request - .metadata_mut() - .insert("authorization", token_value); - - let response = client - .reject_model(request) - .await - .context("Failed to reject model")? - .into_inner(); - - if response.success { - println!("Model '{}' rejected. Reason: {}", model_id, reason); - } else { - println!("Failed to reject model '{}': {}", model_id, response.message); - } - - Ok(()) -} diff --git a/bin/fxt/src/commands/risk.rs b/bin/fxt/src/commands/risk.rs new file mode 100644 index 000000000..8a7af1ef8 --- /dev/null +++ b/bin/fxt/src/commands/risk.rs @@ -0,0 +1,46 @@ +//! `fxt risk` -- risk management. + +use anyhow::Result; +use clap::{Parser, Subcommand}; + +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; + +#[derive(Parser, Debug)] +pub struct RiskCommand { + #[command(subcommand)] + action: RiskAction, +} + +#[derive(Subcommand, Debug)] +enum RiskAction { + /// Show current risk status (VaR, exposure, limits) + Status, + /// Show/update risk limits + Limits { + /// Limit key to get/set (e.g. max_position_size) + key: Option, + /// Value to set (omit to read current value) + value: Option, + }, + /// Show drawdown analysis + Drawdown, + /// Run stress test scenarios + Stress { + /// Scenario name (e.g. flash-crash, vol-spike) + #[arg(long)] + scenario: Option, + }, + /// Emergency: kill switch + Emergency { + /// Confirm emergency shutdown + #[arg(long)] + confirm: bool, + }, +} + +impl RiskCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("risk command not yet implemented") + } +} diff --git a/bin/fxt/src/commands/service.rs b/bin/fxt/src/commands/service.rs new file mode 100644 index 000000000..e5d49264b --- /dev/null +++ b/bin/fxt/src/commands/service.rs @@ -0,0 +1,50 @@ +//! `fxt service` -- service operations. + +use anyhow::Result; +use clap::{Parser, Subcommand}; + +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; + +#[derive(Parser, Debug)] +pub struct ServiceCommand { + #[command(subcommand)] + action: ServiceAction, +} + +#[derive(Subcommand, Debug)] +enum ServiceAction { + /// List all services with health status + List, + /// Detailed service status + Status { + /// Service name + service: String, + }, + /// Stream service logs + Logs { + /// Service name + service: String, + /// Follow log output + #[arg(long, short)] + follow: bool, + }, + /// Restart a service + Restart { + /// Service name + service: String, + }, + /// Deploy binary update + Deploy { + /// Service name + service: String, + }, + /// Health check all services + Health, +} + +impl ServiceCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("service command not yet implemented") + } +} diff --git a/bin/fxt/src/commands/trade.rs b/bin/fxt/src/commands/trade.rs index 8a71f7fcd..5c6cfb319 100644 --- a/bin/fxt/src/commands/trade.rs +++ b/bin/fxt/src/commands/trade.rs @@ -1,129 +1,49 @@ -//! TLI Trade Commands -//! -//! Trading operations with ML-powered decision making. -//! -//! # Architecture -//! This module acts as a routing layer for trade-related commands: -//! - `ml` - ML-powered trading operations (ensemble voting, predictions, performance) -//! -//! # Command Flow -//! User → main.rs → trade.rs → `trade_ml.rs` → API Gateway → Trading Service -//! -//! # Future Extensions -//! - `manual` - Manual order submission -//! - `modify` - Order modification -//! - `cancel` - Order cancellation +//! `fxt trade` -- order management. use anyhow::Result; -use clap::{Args, Subcommand}; +use clap::{Parser, Subcommand}; -use crate::commands::trade_ml::{execute_trade_ml_command, TradeMlArgs}; +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; -/// Trade command arguments -#[derive(Debug, Args)] -pub struct TradeArgs { +#[derive(Parser, Debug)] +pub struct TradeCommand { #[command(subcommand)] - pub command: TradeCommand, + action: TradeAction, } -/// Trade subcommands -#[derive(Debug, Subcommand)] -pub enum TradeCommand { - /// ML-powered trading commands - #[command(name = "ml")] - Ml(TradeMlArgs), +#[derive(Subcommand, Debug)] +enum TradeAction { + /// Submit a new order + Submit { + /// Trading symbol (e.g. ES.FUT) + #[arg(long)] + symbol: String, + /// Order side: buy or sell + #[arg(long)] + side: String, + /// Quantity + #[arg(long)] + qty: f64, + /// Limit price (omit for market orders) + #[arg(long)] + price: Option, + }, + /// Cancel an open order + Cancel { + /// Order ID to cancel + order_id: String, + }, + /// List open positions + Positions, + /// List open orders + Orders, + /// Account summary (balance, margin, P&L) + Account, } -/// Execute trade command -/// -/// # Arguments -/// * `args` - Trade command arguments (contains subcommand) -/// * `api_gateway_url` - API Gateway URL for gRPC connection -/// * `jwt_token` - JWT authentication token -/// -/// # Returns -/// - `Ok(())` - Command executed successfully -/// - `Err(anyhow::Error)` - Command execution failed -/// -/// # Routing -/// This function routes to the appropriate subcommand handler: -/// - `TradeCommand::Ml` → `execute_trade_ml_command()` -/// -/// # Example -/// ```no_run -/// use fxt::commands::trade::{TradeArgs, TradeCommand, execute_trade_command}; -/// use fxt::commands::trade_ml::TradeMlArgs; -/// -/// # async fn example() -> anyhow::Result<()> { -/// let args = TradeArgs { -/// command: TradeCommand::Ml(TradeMlArgs { /* ... */ }), -/// }; -/// -/// execute_trade_command(args, "http://localhost:50051", "jwt-token").await?; -/// # Ok(()) -/// # } -/// ``` -pub async fn execute_trade_command( - args: TradeArgs, - api_gateway_url: &str, - jwt_token: &str, -) -> Result<()> { - match args.command { - TradeCommand::Ml(ml_args) => { - execute_trade_ml_command(ml_args, api_gateway_url, jwt_token).await - }, - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_trade_args_structure() { - // Verify TradeArgs struct is correctly defined - // This ensures the command structure is valid for clap parsing - use clap::Parser; - - #[derive(Parser)] - struct TestCli { - #[command(flatten)] - trade_args: TradeArgs, - } - - // Test that the structure compiles and can be parsed - // (Actual parsing is tested in main.rs integration tests) - } - - #[test] - fn test_trade_command_variants() { - // Verify TradeCommand enum has expected variants - use crate::commands::trade_ml::TradeMlArgs; - - let _ml_variant = TradeCommand::Ml(TradeMlArgs { - command: crate::commands::trade_ml::TradeMlCommand::Performance { model: None }, - }); - - // Test compiles = variants are correctly defined - } - - #[tokio::test] - async fn test_execute_trade_command_routing() { - use crate::commands::trade_ml::{TradeMlArgs, TradeMlCommand}; - - // Create a test TradeArgs with ML subcommand - let args = TradeArgs { - command: TradeCommand::Ml(TradeMlArgs { - command: TradeMlCommand::Performance { - model: Some("DQN".to_owned()), - }, - }), - }; - - // Execute command (will fail due to no actual API Gateway, but tests routing) - let result = execute_trade_command(args, "http://localhost:50051", "mock-token").await; - - // Should attempt to execute (may fail due to connection, but routing works) - assert!(result.is_ok() || result.is_err()); +impl TradeCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("trade command not yet implemented") } } diff --git a/bin/fxt/src/commands/trade_ml.rs b/bin/fxt/src/commands/trade_ml.rs deleted file mode 100644 index bb740788a..000000000 --- a/bin/fxt/src/commands/trade_ml.rs +++ /dev/null @@ -1,1312 +0,0 @@ -//! TLI ML Trading Commands -//! -//! Command-line interface for ML-powered trading operations. -//! Connects to API Gateway for ML order submission, prediction viewing, and performance metrics. -//! -//! # Commands -//! - `submit` - Execute ML-based trade (ensemble or single model) -//! - `predictions` - View ML prediction history with outcomes -//! - `performance` - View ML model performance metrics -//! -//! # Architecture -//! - Pure client implementation (connects ONLY to API Gateway at port 50051) -//! - gRPC communication with `TradingService` via API Gateway proxy -//! - No direct service dependencies (proper microservice architecture) - -use anyhow::Result; -use chrono; -use clap::{Args, Subcommand}; -use colored::Colorize; -use comfy_table::{Cell, Color, Table}; - -/// ML Trading command arguments -#[derive(Args, Debug)] -pub struct TradeMlArgs { - #[command(subcommand)] - pub command: TradeMlCommand, -} - -/// ML Trading subcommands -#[derive(Subcommand, Debug)] -pub enum TradeMlCommand { - /// Submit ML-based trade order - #[clap(long_about = "Execute ML-generated trading order.\n\n\ - Supports:\n\ - - Ensemble voting (DQN+PPO+MAMBA2+TFT)\n\ - - Single model selection (--model flag)\n\ - - Real-time confidence scoring\n\n\ - Examples:\n\ - tli trade ml submit --symbol ES.FUT --account main\n\ - tli trade ml submit --symbol ES.FUT --account main --model DQN")] - Submit { - /// Trading symbol (e.g., ES.FUT, NQ.FUT) - #[arg(short, long, required = true)] - symbol: String, - - /// Account ID - #[arg(short, long, required = true)] - account: String, - - /// Use specific model (default: ensemble) - #[arg(short, long)] - model: Option, - }, - - /// View ML prediction history - #[clap(long_about = "View historical ML predictions with outcomes.\n\n\ - Shows:\n\ - - Predicted action (BUY/SELL/HOLD)\n\ - - Confidence levels\n\ - - Actual P&L (if executed)\n\ - - Individual model predictions\n\n\ - Examples:\n\ - tli trade ml predictions --symbol ES.FUT\n\ - tli trade ml predictions --symbol ES.FUT --model MAMBA2 --limit 5")] - Predictions { - /// Symbol to filter by - #[arg(short, long, required = true)] - symbol: String, - - /// Filter by model name - #[arg(short, long)] - model: Option, - - /// Max predictions to return - #[arg(short, long, default_value = "10")] - limit: i32, - }, - - /// View ML model performance metrics - #[clap(long_about = "View ML model performance statistics.\n\n\ - Metrics:\n\ - - Accuracy (profitable predictions / total predictions)\n\ - - Sharpe ratio (risk-adjusted returns)\n\ - - Average P&L per prediction\n\ - - Total predictions made\n\n\ - Examples:\n\ - tli trade ml performance\n\ - tli trade ml performance --model PPO")] - Performance { - /// Filter by model name - #[arg(short, long)] - model: Option, - }, - - /// View current regime state (Wave D) - #[clap(long_about = "View current regime state for a symbol.\n\n\ - Shows:\n\ - - Current regime (TRENDING/RANGING/VOLATILE/CRISIS)\n\ - - Confidence level\n\ - - CUSUM statistics (S+, S-)\n\ - - ADX (Average Directional Index)\n\ - - Stability and entropy scores\n\n\ - Examples:\n\ - tli trade ml regime --symbol ES.FUT\n\ - tli trade ml regime --symbol NQ.FUT")] - Regime { - /// Symbol to query - #[arg(short, long, required = true)] - symbol: String, - }, - - /// View regime transition history (Wave D) - #[clap(long_about = "View regime transition history for a symbol.\n\n\ - Shows:\n\ - - Transition timestamps\n\ - - From/to regime changes\n\ - - Duration in previous regime\n\ - - Transition probability\n\n\ - Examples:\n\ - tli trade ml transitions --symbol ES.FUT\n\ - tli trade ml transitions --symbol NQ.FUT --limit 20")] - Transitions { - /// Symbol to query - #[arg(short, long, required = true)] - symbol: String, - - /// Max transitions to return - #[arg(short, long, default_value = "100")] - limit: i32, - }, -} - -impl TradeMlArgs { - /// Execute ML trading command - /// - /// Routes to appropriate subcommand handler. - /// All commands connect to API Gateway (). - pub async fn execute(&self, api_gateway_url: &str, jwt_token: &str) -> Result<()> { - match &self.command { - TradeMlCommand::Submit { - symbol, - account, - model, - } => { - self.submit_ml_order( - symbol, - account, - model.as_deref(), - api_gateway_url, - jwt_token, - ) - .await - }, - TradeMlCommand::Predictions { - symbol, - model, - limit, - } => { - self.get_ml_predictions( - symbol, - model.as_deref(), - *limit, - api_gateway_url, - jwt_token, - ) - .await - }, - TradeMlCommand::Performance { model } => { - self.get_ml_performance(model.as_deref(), api_gateway_url, jwt_token) - .await - }, - TradeMlCommand::Regime { symbol } => { - self.get_regime_state(symbol, api_gateway_url, jwt_token) - .await - }, - TradeMlCommand::Transitions { symbol, limit } => { - self.get_regime_transitions(symbol, *limit, api_gateway_url, jwt_token) - .await - }, - } - } - - /// Submit ML-generated order - /// - /// # Arguments - /// * `symbol` - Trading symbol (e.g., ES.FUT) - /// * `account` - Account ID - /// * `model` - Optional specific model name (None = ensemble) - /// * `api_gateway_url` - API Gateway URL - /// * `jwt_token` - JWT authentication token - /// - /// Connects to API Gateway via gRPC and submits ML order request. - async fn submit_ml_order( - &self, - symbol: &str, - account: &str, - model: Option<&str>, - api_gateway_url: &str, - jwt_token: &str, - ) -> Result<()> { - use crate::proto::trading::OrderSide; - - // Step 1: Get ML prediction from API Gateway - let prediction_result = self - .get_ml_prediction(symbol, model, api_gateway_url, jwt_token) - .await; - - let (predicted_action, confidence, model_display) = prediction_result - .map_err(|e| anyhow::anyhow!("ML prediction failed: {e}\n\nEnsure the API Gateway is reachable and ML models are loaded."))?; - - // Step 2: Submit order based on ML prediction - let order_side = match predicted_action.as_str() { - "BUY" | "STRONG_BUY" => OrderSide::Buy, - "SELL" | "STRONG_SELL" => OrderSide::Sell, - "HOLD" | _ => { - println!( - "{}", - format!( - "\u{2139}\u{fe0f} ML prediction is HOLD (confidence: {:.2}%)", - confidence * 100.0 - ) - .cyan() - ); - println!("{}", "No order submitted.".cyan()); - return Ok(()); - }, - }; - - let order_result = self - .submit_order_to_gateway( - symbol, - account, - order_side, - 1.0, // Default quantity: 1 contract - api_gateway_url, - jwt_token, - ) - .await; - - let order_id = order_result - .map_err(|e| anyhow::anyhow!("Order submission failed: {e}\n\nEnsure the API Gateway and Trading Service are running."))?; - - println!("{}", "\u{2705} ML order submitted successfully!".green()); - println!(); - println!("Order ID: {}", order_id.bright_green()); - println!("Symbol: {}", symbol.bright_cyan()); - println!("Model: {}", model_display.bright_magenta()); - println!( - "Predicted Action: {}", - predicted_action.bright_white().bold() - ); - println!( - "Confidence: {} ({:.1}%)", - format!("{:.2}", confidence).bright_green(), - confidence * 100.0 - ); - println!("Quantity: 1 contract"); - println!("Account: {}", account.bright_yellow()); - - Ok(()) - } - - /// Get ML prediction from API Gateway - /// - /// # Returns - /// Tuple of (`predicted_action`, confidence, `model_display_name`) - async fn get_ml_prediction( - &self, - symbol: &str, - model: Option<&str>, - api_gateway_url: &str, - jwt_token: &str, - ) -> Result<(String, f64, String)> { - use crate::proto::ml::{ml_service_client::MlServiceClient, EnsembleRequest}; - - // Connect to API Gateway - let mut client = MlServiceClient::connect(api_gateway_url.to_owned()) - .await - .map_err(|e| anyhow::anyhow!("Failed to connect to API Gateway: {}", e))?; - - // Create ensemble request - let model_names = if let Some(m) = model { - vec![m.to_owned()] - } else { - vec![] // Empty = all models (ensemble) - }; - - let mut request = tonic::Request::new(EnsembleRequest { - symbols: vec![symbol.to_owned()], - model_names, - method: 1, // ENSEMBLE_METHOD_WEIGHTED_AVERAGE - }); - - // Add JWT token to metadata - request.metadata_mut().insert( - "authorization", - format!("Bearer {}", jwt_token) - .parse() - .map_err(|e| anyhow::anyhow!("Invalid JWT token: {}", e))?, - ); - - // Make gRPC call - let response = client - .get_ensemble_vote(request) - .await - .map_err(|e| anyhow::anyhow!("Failed to get ensemble vote: {}", e))?; - - let ensemble_response = response.into_inner(); - - // Extract prediction from first symbol - let vote = ensemble_response - .votes - .first() - .ok_or_else(|| anyhow::anyhow!("No predictions returned for symbol"))?; - - let predicted_action = match vote.consensus { - 1 => "BUY".to_owned(), - 2 => "SELL".to_owned(), - 3 => "HOLD".to_owned(), - 4 => "STRONG_BUY".to_owned(), - 5 => "STRONG_SELL".to_owned(), - _ => "HOLD".to_owned(), - }; - - let confidence = vote.confidence; - - let model_display = model.unwrap_or("Ensemble").to_owned(); - - Ok((predicted_action, confidence, model_display)) - } - - /// Submit order to Trading Service via API Gateway - /// - /// # Returns - /// Order ID on success - async fn submit_order_to_gateway( - &self, - symbol: &str, - account: &str, - side: crate::proto::trading::OrderSide, - quantity: f64, - api_gateway_url: &str, - jwt_token: &str, - ) -> Result { - use crate::proto::trading::{ - trading_service_client::TradingServiceClient, SubmitOrderRequest, - }; - - // Connect to API Gateway - let mut client = TradingServiceClient::connect(api_gateway_url.to_owned()) - .await - .map_err(|e| anyhow::anyhow!("Failed to connect to API Gateway: {}", e))?; - - // Create order request - let mut request = tonic::Request::new(SubmitOrderRequest { - symbol: symbol.to_owned(), - side: side as i32, - order_type: 1, // ORDER_TYPE_MARKET - quantity, - price: None, - stop_price: None, - time_in_force: "GTC".to_owned(), - client_order_id: format!("ml_order_{}", chrono::Utc::now().timestamp_millis()), - }); - - // Add JWT token and account metadata - request.metadata_mut().insert( - "authorization", - format!("Bearer {}", jwt_token) - .parse() - .map_err(|e| anyhow::anyhow!("Invalid JWT token: {}", e))?, - ); - - request.metadata_mut().insert( - "account_id", - account - .parse() - .map_err(|e| anyhow::anyhow!("Invalid account ID: {}", e))?, - ); - - // Make gRPC call - let response = client - .submit_order(request) - .await - .map_err(|e| anyhow::anyhow!("Failed to submit order: {}", e))?; - - let order_response = response.into_inner(); - - if !order_response.success { - return Err(anyhow::anyhow!( - "Order rejected: {}", - order_response.message - )); - } - - Ok(order_response.order_id) - } - - /// Get ML prediction history - /// - /// # Arguments - /// * `symbol` - Trading symbol to filter by - /// * `model` - Optional model name filter - /// * `limit` - Maximum predictions to return - /// * `api_gateway_url` - API Gateway URL - /// * `jwt_token` - JWT authentication token - /// - /// # Production Implementation - /// Fetches prediction history from API Gateway via gRPC. - async fn get_ml_predictions( - &self, - symbol: &str, - model: Option<&str>, - limit: i32, - api_gateway_url: &str, - jwt_token: &str, - ) -> Result<()> { - use crate::proto::trading::{ - trading_service_client::TradingServiceClient, GetMlPredictionsRequest, - }; - let predictions_result = async { - let mut client = TradingServiceClient::connect(api_gateway_url.to_owned()) - .await - .map_err(|e| anyhow::anyhow!("Failed to connect to API Gateway: {}", e))?; - - let mut request = tonic::Request::new(GetMlPredictionsRequest { - symbol: symbol.to_owned(), - model_filter: model.map(|m| m.to_owned()), - limit: Some(limit), - }); - - request.metadata_mut().insert( - "authorization", - format!("Bearer {}", jwt_token) - .parse() - .map_err(|e| anyhow::anyhow!("Invalid JWT token: {}", e))?, - ); - - let response = client - .get_ml_predictions(request) - .await - .map_err(|e| anyhow::anyhow!("Failed to get ML predictions: {}", e))?; - - Ok::<_, anyhow::Error>(response.into_inner().predictions) - } - .await; - - let predictions_response = predictions_result - .map_err(|e| anyhow::anyhow!("Failed to fetch predictions: {e}\n\nEnsure the API Gateway and Trading Service are running."))?; - - // Display header - println!(); - println!("{} {}", "ML Predictions for".bold(), symbol.bright_cyan()); - if let Some(model_name) = model { - println!("Model Filter: {}", model_name.bright_magenta()); - } - println!(); - - // Check if we have predictions - if predictions_response.is_empty() { - println!("{}", "No predictions found for this symbol.".yellow()); - println!("Try running `tli trade ml submit --symbol {} --account ` to generate predictions.", symbol); - return Ok(()); - } - - // Print table header - println!("{}", "\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}".bold()); - println!( - "{:<20} {:<10} {:<10} {:<15} {:<12} {:<15}", - "Timestamp".bold(), - "Model".bold(), - "Symbol".bold(), - "Predicted Action".bold(), - "Confidence".bold(), - "Outcome".bold() - ); - println!("{}", "\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}".bold()); - - // Add prediction rows - for pred in &predictions_response { - // Format timestamp - let timestamp = pred.timestamp.clone(); - - // Get model ID - let model_id = pred.model_id.clone(); - - // Color code action - let action_str = match pred.predicted_action.as_str() { - "BUY" => pred.predicted_action.green().to_string(), - "SELL" => pred.predicted_action.red().to_string(), - "HOLD" => pred.predicted_action.yellow().to_string(), - _ => pred.predicted_action.white().to_string(), - }; - - // Format confidence as percentage - let confidence_val = pred.confidence * 100.0; - let confidence_str = format!("{:.1}%", confidence_val); - let confidence_colored = if pred.confidence >= 0.75 { - confidence_str.green().to_string() - } else if pred.confidence >= 0.60 { - confidence_str.yellow().to_string() - } else { - confidence_str.red().to_string() - }; - - // Format outcome (actual return if available) - let outcome_str = if let Some(actual_return) = pred.actual_return { - let return_pct = actual_return * 100.0; - let formatted = format!("{:+.2}%", return_pct); - if actual_return > 0.0 { - formatted.green().to_string() - } else if actual_return < 0.0 { - formatted.red().to_string() - } else { - formatted.white().to_string() - } - } else { - "N/A".white().to_string() - }; - - println!( - "{:<20} {:<10} {:<10} {:<15} {:<12} {:<15}", - timestamp, model_id, pred.symbol, action_str, confidence_colored, outcome_str - ); - } - - // Footer - println!("{}", "\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}".bold()); - - // Summary - let count = predictions_response.len(); - println!( - "Showing {} prediction{}", - count, - if count != 1 { "s" } else { "" } - ); - println!(); - - Ok(()) - } - - /// Get ML model performance metrics - /// - /// # Arguments - /// * `model` - Optional model name filter (None = all models) - /// * `api_gateway_url` - API Gateway URL - /// * `jwt_token` - JWT authentication token - /// - /// # Production Implementation - /// Fetches performance metrics from Trading Service via API Gateway gRPC proxy. - /// Displays ML model performance in formatted table with color-coded metrics. - async fn get_ml_performance( - &self, - model: Option<&str>, - api_gateway_url: &str, - jwt_token: &str, - ) -> Result<()> { - use crate::proto::trading::{ - trading_service_client::TradingServiceClient, GetMlPerformanceRequest, - }; - let performance_result = async { - let mut client = TradingServiceClient::connect(api_gateway_url.to_owned()) - .await - .map_err(|e| anyhow::anyhow!("Failed to connect to API Gateway: {}", e))?; - - let mut request = tonic::Request::new(GetMlPerformanceRequest { - model_filter: model.map(|s| s.to_owned()), - }); - - request.metadata_mut().insert( - "authorization", - format!("Bearer {}", jwt_token) - .parse() - .map_err(|e| anyhow::anyhow!("Invalid JWT token: {}", e))?, - ); - - let response = client - .get_ml_performance(request) - .await - .map_err(|e| anyhow::anyhow!("GetMLPerformance RPC failed: {}", e))?; - - Ok::<_, anyhow::Error>(response.into_inner().models) - } - .await; - - let models = performance_result - .map_err(|e| anyhow::anyhow!("Failed to fetch performance metrics: {e}\n\nEnsure the API Gateway and Trading Service are running."))?; - - // Display ML Model Performance header - println!("\n{}", "ML Model Performance (Last 30 days)".bold()); - println!(); - - // Display table header - println!("{}", "\u{250c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{252c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{252c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{252c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{252c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{252c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2510}".bright_black()); - println!("\u{2502} {:<6} \u{2502} {:<8} \u{2502} {:<12} \u{2502} {:<12} \u{2502} {:<9} \u{2502} {:<10} \u{2502}", - "Model".bold(), - "Accuracy".bold(), - "Predictions".bold(), - "Sharpe Ratio".bold(), - "Avg Return".bold(), - "Max Drawdown".bold() - ); - println!("{}", "\u{251c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{253c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{253c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{253c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{253c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{253c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2524}".bright_black()); - - // Display each model's performance - for model_perf in &models { - let accuracy = model_perf.accuracy * 100.0; // Convert to percentage - let accuracy_str = format!("{:.1}%", accuracy); - let accuracy_colored = if accuracy > 70.0 { - accuracy_str.green().to_string() - } else if accuracy > 65.0 { - accuracy_str.yellow().to_string() - } else { - accuracy_str.red().to_string() - }; - - let sharpe_str = format!("{:.2}", model_perf.sharpe_ratio); - let sharpe_colored = if model_perf.sharpe_ratio > 2.0 { - sharpe_str.green().to_string() - } else if model_perf.sharpe_ratio > 1.5 { - sharpe_str.yellow().to_string() - } else { - sharpe_str.red().to_string() - }; - - let avg_return = model_perf.avg_return * 100.0; // Convert to percentage - let return_str = if avg_return >= 0.0 { - format!("+{:.1}%", avg_return) - } else { - format!("{:.1}%", avg_return) - }; - let return_colored = if avg_return > 2.0 { - return_str.green().to_string() - } else if avg_return > 0.0 { - return_str.yellow().to_string() - } else { - return_str.red().to_string() - }; - - let drawdown = model_perf.max_drawdown * 100.0; // Convert to percentage - let drawdown_str = format!("{:.1}%", drawdown); - let drawdown_colored = if drawdown.abs() < 3.0 { - drawdown_str.green().to_string() - } else if drawdown.abs() < 5.0 { - drawdown_str.yellow().to_string() - } else { - drawdown_str.red().to_string() - }; - - println!("\u{2502} {:<6} \u{2502} {:<8} \u{2502} {:<12} \u{2502} {:<12} \u{2502} {:<9} \u{2502} {:<10} \u{2502}", - model_perf.model_id.bright_magenta(), - accuracy_colored, - model_perf.total_predictions.to_string().bright_cyan(), - sharpe_colored, - return_colored, - drawdown_colored - ); - } - - println!("{}", "\u{2514}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2534}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2534}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2534}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2534}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2534}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2518}".bright_black()); - - // Display ensemble summary if showing all models - if model.is_none() { - println!(); - // Calculate ensemble metrics from models - let active_models = models.len(); - let total_models = 4; // DQN, PPO, MAMBA2, TFT - let ensemble_threshold = 0.70; // Default confidence threshold - - println!( - "Ensemble Confidence Threshold: {}", - format!("{:.2}", ensemble_threshold).bright_green() - ); - println!( - "Active Models: {} ({}/{} models operational)", - format!("{}/{}", active_models, total_models).bright_yellow(), - active_models, - total_models - ); - } - - Ok(()) - } - - /// Get current regime state for a symbol (Wave D) - /// - /// # Arguments - /// * `symbol` - Trading symbol to query - /// * `api_gateway_url` - API Gateway URL - /// * `jwt_token` - JWT authentication token - async fn get_regime_state( - &self, - symbol: &str, - api_gateway_url: &str, - jwt_token: &str, - ) -> Result<()> { - use crate::proto::trading::{ - trading_service_client::TradingServiceClient, GetRegimeStateRequest, - }; - - let mut client = TradingServiceClient::connect(api_gateway_url.to_owned()) - .await - .map_err(|e| anyhow::anyhow!("Failed to connect to API Gateway: {}", e))?; - - let mut request = tonic::Request::new(GetRegimeStateRequest { - symbol: symbol.to_owned(), - }); - - request.metadata_mut().insert( - "authorization", - format!("Bearer {}", jwt_token) - .parse() - .map_err(|e| anyhow::anyhow!("Invalid JWT token: {}", e))?, - ); - - let response = client - .get_regime_state(request) - .await - .map_err(|e| anyhow::anyhow!("GetRegimeState RPC failed: {}", e))?; - - let regime_state = response.into_inner(); - - // Display regime state - println!(); - println!( - "{}", - format!("\u{1f4ca} Regime State: {}", regime_state.symbol) - .bright_cyan() - .bold() - ); - println!("{}", "\u{2500}".repeat(80).bright_black()); - - let regime_colored = match regime_state.current_regime.as_str() { - "TRENDING" => regime_state.current_regime.bright_green(), - "RANGING" => regime_state.current_regime.bright_yellow(), - "VOLATILE" => regime_state.current_regime.bright_red(), - "CRISIS" => regime_state.current_regime.red().bold(), - _ => regime_state.current_regime.white(), - }; - - println!("Current Regime: {}", regime_colored); - println!( - "Confidence: {:.2}%", - (regime_state.confidence * 100.0) - ); - println!(); - println!("Statistics:"); - println!(" CUSUM S+: {:.4}", regime_state.cusum_s_plus); - println!(" CUSUM S-: {:.4}", regime_state.cusum_s_minus); - println!(" ADX: {:.2}", regime_state.adx); - println!( - " Stability: {:.2}%", - (regime_state.stability * 100.0) - ); - println!(" Entropy: {:.4}", regime_state.entropy); - - let timestamp = chrono::DateTime::from_timestamp_nanos(regime_state.updated_at_unix_nanos); - println!(); - println!( - "Last Updated: {}", - timestamp.format("%Y-%m-%d %H:%M:%S UTC") - ); - println!("{}", "\u{2500}".repeat(80).bright_black()); - println!(); - - Ok(()) - } - - /// Get regime transition history for a symbol (Wave D) - /// - /// # Arguments - /// * `symbol` - Trading symbol to query - /// * `limit` - Maximum transitions to return - /// * `api_gateway_url` - API Gateway URL - /// * `jwt_token` - JWT authentication token - async fn get_regime_transitions( - &self, - symbol: &str, - limit: i32, - api_gateway_url: &str, - jwt_token: &str, - ) -> Result<()> { - use crate::proto::trading::{ - trading_service_client::TradingServiceClient, GetRegimeTransitionsRequest, - }; - - let mut client = TradingServiceClient::connect(api_gateway_url.to_owned()) - .await - .map_err(|e| anyhow::anyhow!("Failed to connect to API Gateway: {}", e))?; - - let mut request = tonic::Request::new(GetRegimeTransitionsRequest { - symbol: symbol.to_owned(), - limit, - }); - - request.metadata_mut().insert( - "authorization", - format!("Bearer {}", jwt_token) - .parse() - .map_err(|e| anyhow::anyhow!("Invalid JWT token: {}", e))?, - ); - - let response = client - .get_regime_transitions(request) - .await - .map_err(|e| anyhow::anyhow!("GetRegimeTransitions RPC failed: {}", e))?; - - let transitions_response = response.into_inner(); - - // Display header - println!(); - println!( - "{}", - format!("\u{1f504} Regime Transitions: {}", symbol) - .bright_cyan() - .bold() - ); - println!("{}", "\u{2500}".repeat(95).bright_black()); - println!( - "{:<20} {:<15} {:<15} {:<12} {:<15}", - "Timestamp".bold(), - "From".bold(), - "To".bold(), - "Duration".bold(), - "Probability".bold() - ); - println!("{}", "\u{2500}".repeat(95).bright_black()); - - // Display transitions - for trans in &transitions_response.transitions { - let timestamp = chrono::DateTime::from_timestamp_nanos(trans.timestamp_unix_nanos); - let timestamp_str = timestamp.format("%Y-%m-%d %H:%M:%S").to_string(); - - let from_colored = match trans.from_regime.as_str() { - "TRENDING" => trans.from_regime.bright_green(), - "RANGING" => trans.from_regime.bright_yellow(), - "VOLATILE" => trans.from_regime.bright_red(), - "CRISIS" => trans.from_regime.red().bold(), - _ => trans.from_regime.white(), - }; - - let to_colored = match trans.to_regime.as_str() { - "TRENDING" => trans.to_regime.bright_green(), - "RANGING" => trans.to_regime.bright_yellow(), - "VOLATILE" => trans.to_regime.bright_red(), - "CRISIS" => trans.to_regime.red().bold(), - _ => trans.to_regime.white(), - }; - - let duration_str = format!("{} bars", trans.duration_bars); - let prob_str = format!("{:.2}%", trans.transition_probability * 100.0); - - println!( - "{:<20} {:<15} {:<15} {:<12} {:<15}", - timestamp_str, from_colored, to_colored, duration_str, prob_str - ); - } - - println!("{}", "\u{2500}".repeat(95).bright_black()); - println!( - "Showing {} transition{}", - transitions_response.transitions.len(), - if transitions_response.transitions.len() != 1 { - "s" - } else { - "" - } - ); - println!(); - - Ok(()) - } -} - -/// Execute ML trading command (public interface for main.rs) -/// -/// # Arguments -/// * `args` - ML trading command arguments -/// * `api_gateway_url` - API Gateway URL -/// * `jwt_token` - JWT authentication token -pub async fn execute_trade_ml_command( - args: TradeMlArgs, - api_gateway_url: &str, - jwt_token: &str, -) -> Result<()> { - args.execute(api_gateway_url, jwt_token).await -} - -// ============================================================================ -// Rich Terminal Formatting Functions -// ============================================================================ -// -// These functions provide rich terminal formatting for ML trading command outputs -// using comfy-table for ASCII tables and owo-colors for advanced color coding. -// -// Used by Agents 2-4 to display: -// - ML order submission results with color-coded confidence levels -// - ML prediction history with action colors and outcomes -// - ML model performance metrics with threshold-based color coding - -/// Response type for ML order submission (mirrors gRPC proto) -#[derive(Debug, Clone)] -pub struct SubmitMLOrderResponse { - pub order_id: String, - pub symbol: String, - pub model_used: String, - pub predicted_action: String, - pub confidence: f64, - pub quantity: f64, - pub account_id: String, -} - -/// Single ML prediction entry (mirrors gRPC proto) -#[derive(Debug, Clone)] -pub struct MLPrediction { - pub timestamp: String, - pub model_id: String, - pub symbol: String, - pub predicted_action: String, - pub confidence: f64, - pub actual_return: Option, -} - -/// Response type for ML predictions (mirrors gRPC proto) -#[derive(Debug, Clone)] -pub struct GetMLPredictionsResponse { - pub predictions: Vec, -} - -/// Single model performance entry (mirrors gRPC proto) -#[derive(Debug, Clone)] -pub struct ModelPerformance { - pub model_id: String, - pub accuracy: f64, - pub total_predictions: i64, - pub sharpe_ratio: f64, - pub avg_return: f64, - pub max_drawdown: f64, -} - -/// Response type for ML performance (mirrors gRPC proto) -#[derive(Debug, Clone)] -pub struct GetMLPerformanceResponse { - pub models: Vec, - pub ensemble_threshold: f64, - pub active_models: i32, - pub total_models: i32, -} - -/// Format ML order submission response with rich terminal colors -/// -/// Displays: -/// - Success message in green/bold -/// - Order details with cyan labels -/// - Model name color-coded (yellow for Ensemble, blue for single models) -/// - Action color-coded (green=BUY, red=SELL, yellow=HOLD) -/// - Confidence with threshold-based colors (>80%=green, >60%=yellow, else=red) -/// -/// # Arguments -/// * `response` - ML order submission response data -/// -/// # Example Output -/// ```text -/// ✅ ML order submitted successfully! -/// -/// Order ID: order_12345 -/// Symbol: ES.FUT -/// Model: Ensemble -/// Predicted Action: BUY -/// Confidence: 85.0% -/// Quantity: 1 -/// Account: main_account -/// ``` -pub fn format_ml_order_submission(response: &SubmitMLOrderResponse) { - use owo_colors::OwoColorize; - - println!( - "{}", - "\u{2705} ML order submitted successfully!".green().bold() - ); - println!(); - - println!("{}: {}", "Order ID".cyan().bold(), response.order_id); - println!("{}: {}", "Symbol".cyan().bold(), response.symbol); - println!( - "{}: {}", - "Model".cyan().bold(), - if response.model_used.contains("Ensemble") { - response.model_used.yellow().to_string() - } else { - response.model_used.blue().to_string() - } - ); - - let action_colored = match response.predicted_action.as_str() { - "BUY" => response.predicted_action.green().to_string(), - "SELL" => response.predicted_action.red().to_string(), - _ => response.predicted_action.yellow().to_string(), - }; - println!("{}: {}", "Predicted Action".cyan().bold(), action_colored); - - let confidence_pct = (response.confidence * 100.0).round(); - let confidence_colored = if confidence_pct >= 80.0 { - format!("{:.1}%", confidence_pct).green().to_string() - } else if confidence_pct >= 60.0 { - format!("{:.1}%", confidence_pct).yellow().to_string() - } else { - format!("{:.1}%", confidence_pct).red().to_string() - }; - println!("{}: {}", "Confidence".cyan().bold(), confidence_colored); - - println!("{}: {}", "Quantity".cyan().bold(), response.quantity); - println!("{}: {}", "Account".cyan().bold(), response.account_id); -} - -/// Format ML predictions history with rich terminal table -/// -/// Displays: -/// - Header with symbol and prediction count -/// - ASCII table with comfy-table -/// - Action column color-coded (green=BUY, red=SELL, yellow=HOLD) -/// - Confidence column with threshold-based colors -/// - Outcome column showing actual P&L (green=profit, red=loss, grey=N/A) -/// -/// # Arguments -/// * `response` - ML predictions response data -/// * `symbol` - Trading symbol for display -/// -/// # Example Output -/// ```text -/// ML Predictions for ES.FUT (Last 10) -/// -/// ┌────────────┬────────┬────────┬─────────┬────────────┬─────────┐ -/// │ Timestamp │ Model │ Symbol │ Action │ Confidence │ Outcome │ -/// ├────────────┼────────┼────────┼─────────┼────────────┼─────────┤ -/// │ 2025-10-16 │ MAMBA2 │ ES.FUT │ BUY │ 85.0% │ +2.50% │ -/// │ 2025-10-16 │ DQN │ ES.FUT │ SELL │ 72.5% │ -1.20% │ -/// └────────────┴────────┴────────┴─────────┴────────────┴─────────┘ -/// ``` -pub fn format_ml_predictions(response: &GetMLPredictionsResponse, symbol: &str) { - use owo_colors::OwoColorize; - - println!( - "{}", - format!( - "ML Predictions for {} (Last {})", - symbol, - response.predictions.len() - ) - .cyan() - .bold() - ); - println!(); - - let mut table = Table::new(); - table.set_header(vec![ - Cell::new("Timestamp").fg(Color::Cyan), - Cell::new("Model").fg(Color::Cyan), - Cell::new("Symbol").fg(Color::Cyan), - Cell::new("Action").fg(Color::Cyan), - Cell::new("Confidence").fg(Color::Cyan), - Cell::new("Outcome").fg(Color::Cyan), - ]); - - for pred in &response.predictions { - let action_cell = match pred.predicted_action.as_str() { - "BUY" => Cell::new(&pred.predicted_action).fg(Color::Green), - "SELL" => Cell::new(&pred.predicted_action).fg(Color::Red), - _ => Cell::new(&pred.predicted_action).fg(Color::Yellow), - }; - - let confidence_pct = format!("{:.1}%", pred.confidence * 100.0); - let confidence_cell = if pred.confidence >= 0.8 { - Cell::new(confidence_pct).fg(Color::Green) - } else if pred.confidence >= 0.6 { - Cell::new(confidence_pct).fg(Color::Yellow) - } else { - Cell::new(confidence_pct).fg(Color::Red) - }; - - let outcome_str = match pred.actual_return { - Some(ret) => format!("{:+.2}%", ret * 100.0), - None => "N/A".to_owned(), - }; - let outcome_cell = match pred.actual_return { - Some(ret) if ret > 0.0 => Cell::new(outcome_str).fg(Color::Green), - Some(ret) if ret < 0.0 => Cell::new(outcome_str).fg(Color::Red), - _ => Cell::new(outcome_str).fg(Color::Grey), - }; - - table.add_row(vec![ - Cell::new(&pred.timestamp), - Cell::new(&pred.model_id), - Cell::new(&pred.symbol), - action_cell, - confidence_cell, - outcome_cell, - ]); - } - - println!("{table}"); -} - -/// Format ML model performance metrics with rich terminal table -/// -/// Displays: -/// - Header with time period (Last 30 days) -/// - ASCII table with comfy-table -/// - Accuracy column with threshold-based colors (>70%=green, >60%=yellow, else=red) -/// - Sharpe ratio with threshold-based colors (>1.5=green, >1.0=yellow, else=red) -/// - Returns with sign prefix and color coding -/// - Ensemble summary with active model count -/// -/// # Arguments -/// * `response` - ML performance response data -/// -/// # Example Output -/// ```text -/// ML Model Performance (Last 30 days) -/// -/// ┌────────┬──────────┬──────────────┬──────────────┬────────────┬──────────────┐ -/// │ Model │ Accuracy │ Predictions │ Sharpe Ratio │ Avg Return │ Max Drawdown │ -/// ├────────┼──────────┼──────────────┼──────────────┼────────────┼──────────────┤ -/// │ MAMBA2 │ 72.5% │ 150 │ 1.82 │ +2.3% │ 3.1% │ -/// │ DQN │ 68.2% │ 200 │ 1.45 │ +1.8% │ 4.5% │ -/// │ PPO │ 71.0% │ 180 │ 1.67 │ +2.1% │ 3.8% │ -/// │ TFT │ 69.5% │ 175 │ 1.52 │ +1.9% │ 4.2% │ -/// └────────┴──────────┴──────────────┴──────────────┴────────────┴──────────────┘ -/// -/// Ensemble Confidence Threshold: 0.70 -/// Active Models: 4/4 -/// ``` -pub fn format_ml_performance(response: &GetMLPerformanceResponse) { - println!("{}", "ML Model Performance (Last 30 days)".cyan().bold()); - println!(); - - let mut table = Table::new(); - table.set_header(vec![ - Cell::new("Model").fg(Color::Cyan), - Cell::new("Accuracy").fg(Color::Cyan), - Cell::new("Predictions").fg(Color::Cyan), - Cell::new("Sharpe Ratio").fg(Color::Cyan), - Cell::new("Avg Return").fg(Color::Cyan), - Cell::new("Max Drawdown").fg(Color::Cyan), - ]); - - for model in &response.models { - let accuracy_cell = if model.accuracy >= 70.0 { - Cell::new(format!("{:.1}%", model.accuracy)).fg(Color::Green) - } else if model.accuracy >= 60.0 { - Cell::new(format!("{:.1}%", model.accuracy)).fg(Color::Yellow) - } else { - Cell::new(format!("{:.1}%", model.accuracy)).fg(Color::Red) - }; - - let sharpe_cell = if model.sharpe_ratio >= 1.5 { - Cell::new(format!("{:.2}", model.sharpe_ratio)).fg(Color::Green) - } else if model.sharpe_ratio >= 1.0 { - Cell::new(format!("{:.2}", model.sharpe_ratio)).fg(Color::Yellow) - } else { - Cell::new(format!("{:.2}", model.sharpe_ratio)).fg(Color::Red) - }; - - table.add_row(vec![ - Cell::new(&model.model_id), - accuracy_cell, - Cell::new(model.total_predictions.to_string()), - sharpe_cell, - Cell::new(format!("{:+.1}%", model.avg_return * 100.0)), - Cell::new(format!("{:.1}%", model.max_drawdown * 100.0)), - ]); - } - - println!("{table}"); - println!(); - println!( - "{}: {:.2}", - "Ensemble Confidence Threshold".cyan(), - response.ensemble_threshold - ); - println!( - "{}: {}/{}", - "Active Models".cyan(), - response.active_models, - response.total_models - ); -} - -#[cfg(test)] -#[allow(clippy::unwrap_used, clippy::expect_used)] -mod tests { - use super::*; - - #[tokio::test] - #[ignore] // Requires running API Gateway - async fn test_submit_command_parses() { - let args = TradeMlArgs { - command: TradeMlCommand::Submit { - symbol: "ES.FUT".to_owned(), - account: "test_account".to_owned(), - model: None, - }, - }; - - let result = args.execute("http://localhost:50051", "test-token").await; - result.unwrap(); - } - - #[tokio::test] - #[ignore] // Requires running API Gateway - async fn test_predictions_command_parses() { - let args = TradeMlArgs { - command: TradeMlCommand::Predictions { - symbol: "ES.FUT".to_owned(), - model: Some("MAMBA2".to_owned()), - limit: 5, - }, - }; - - let result = args.execute("http://localhost:50051", "test-token").await; - result.unwrap(); - } - - #[tokio::test] - #[ignore] // Requires running API Gateway - async fn test_performance_command_parses() { - let args = TradeMlArgs { - command: TradeMlCommand::Performance { - model: Some("PPO".to_owned()), - }, - }; - - let result = args.execute("http://localhost:50051", "test-token").await; - result.unwrap(); - } - - #[test] - fn test_format_ml_order_submission() { - // Test formatting function with sample data - let response = SubmitMLOrderResponse { - order_id: "order_12345".to_owned(), - symbol: "ES.FUT".to_owned(), - model_used: "Ensemble".to_owned(), - predicted_action: "BUY".to_owned(), - confidence: 0.85, - quantity: 1.0, - account_id: "main_account".to_owned(), - }; - - // Should not panic - format_ml_order_submission(&response); - } - - #[test] - fn test_format_ml_predictions() { - // Test formatting function with sample data - let response = GetMLPredictionsResponse { - predictions: vec![ - MLPrediction { - timestamp: "2025-10-16T12:00:00Z".to_owned(), - model_id: "MAMBA2".to_owned(), - symbol: "ES.FUT".to_owned(), - predicted_action: "BUY".to_owned(), - confidence: 0.85, - actual_return: Some(0.025), - }, - MLPrediction { - timestamp: "2025-10-16T11:00:00Z".to_owned(), - model_id: "DQN".to_owned(), - symbol: "ES.FUT".to_owned(), - predicted_action: "SELL".to_owned(), - confidence: 0.72, - actual_return: Some(-0.012), - }, - ], - }; - - // Should not panic - format_ml_predictions(&response, "ES.FUT"); - } - - #[test] - fn test_format_ml_performance() { - // Test formatting function with sample data - let response = GetMLPerformanceResponse { - models: vec![ - ModelPerformance { - model_id: "MAMBA2".to_owned(), - accuracy: 72.5, - total_predictions: 150, - sharpe_ratio: 1.82, - avg_return: 0.023, - max_drawdown: 0.031, - }, - ModelPerformance { - model_id: "DQN".to_owned(), - accuracy: 68.2, - total_predictions: 200, - sharpe_ratio: 1.45, - avg_return: 0.018, - max_drawdown: 0.045, - }, - ], - ensemble_threshold: 0.70, - active_models: 2, - total_models: 4, - }; - - // Should not panic - format_ml_performance(&response); - } -} diff --git a/bin/fxt/src/commands/train.rs b/bin/fxt/src/commands/train.rs new file mode 100644 index 000000000..c3f9ec06c --- /dev/null +++ b/bin/fxt/src/commands/train.rs @@ -0,0 +1,62 @@ +//! `fxt train` -- ML training lifecycle. + +use anyhow::Result; +use clap::{Parser, Subcommand}; + +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; + +#[derive(Parser, Debug)] +pub struct TrainCommand { + #[command(subcommand)] + action: TrainAction, +} + +#[derive(Subcommand, Debug)] +enum TrainAction { + /// Start a training job + Start { + /// Model type (e.g. dqn, ppo, tft, mamba2) + #[arg(long)] + model: String, + /// Path to training config YAML + #[arg(long)] + config: Option, + /// Enable GPU training + #[arg(long)] + gpu: bool, + }, + /// Stop a running training job + Stop { + /// Training job ID + job_id: String, + }, + /// Show training job status + Status { + /// Training job ID + job_id: String, + }, + /// List training jobs + List { + /// Filter by status (running, completed, failed) + #[arg(long)] + status: Option, + /// Filter by model type + #[arg(long)] + model: Option, + }, + /// Stream training logs + Logs { + /// Training job ID + job_id: String, + /// Follow log output + #[arg(long, short)] + follow: bool, + }, +} + +impl TrainCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("train command not yet implemented") + } +} diff --git a/bin/fxt/src/commands/train/list.rs b/bin/fxt/src/commands/train/list.rs deleted file mode 100644 index e7a30f229..000000000 --- a/bin/fxt/src/commands/train/list.rs +++ /dev/null @@ -1,348 +0,0 @@ -//! TLI Train List Command - Display Training Job History -//! -//! Provides CLI interface for listing and filtering ML model training jobs. -//! -//! # Usage -//! -//! ```bash -//! # List all jobs (default: last 50) -//! tli train list -//! -//! # Filter by status -//! tli train list --status RUNNING -//! tli train list --status COMPLETED -//! -//! # Filter by model type -//! tli train list --model TFT -//! tli train list --model DQN -//! -//! # Filter by asset -//! tli train list --asset ES.FUT -//! -//! # Sort by different fields -//! tli train list --sort-by start_time --sort-order desc -//! tli train list --sort-by duration --sort-order asc -//! -//! # Limit results -//! tli train list --limit 10 -//! -//! # Show only batch or single jobs -//! tli train list --batch-only -//! tli train list --single-only -//! -//! # Combined filters -//! tli train list --status RUNNING --model TFT --limit 5 -//! ``` - -use anyhow::{Context, Result}; -use chrono::Utc; -use clap::Parser; -use comfy_table::{presets::UTF8_FULL, Attribute, Cell, Color, ContentArrangement, Table}; -use crate::proto::ml_training::{ - ml_training_service_client::MlTrainingServiceClient, ListTrainingJobsRequest, - TrainingJobSummary, TrainingStatus, -}; -use tonic::{metadata::MetadataValue, Request}; - -/// List training jobs with filtering and sorting -#[derive(Parser, Debug)] -pub struct ListCommand { - /// Filter by job status (PENDING, RUNNING, COMPLETED, FAILED, STOPPED) - #[clap(long)] - pub status: Option, - - /// Filter by model type (DQN, PPO, `MAMBA_2`, TFT, TLOB, LIQUID) - #[clap(long)] - pub model: Option, - - /// Filter by asset (ES.FUT, NQ.FUT, etc.) - #[clap(long)] - pub asset: Option, - - /// Sort by field (`start_time`, duration, status) - #[clap(long, default_value = "start_time")] - pub sort_by: String, - - /// Sort order (asc, desc) - #[clap(long, default_value = "desc")] - pub sort_order: String, - - /// Maximum results to display - #[clap(long, default_value = "50")] - pub limit: u32, - - /// Show only batch jobs - #[clap(long, conflicts_with = "single_only")] - pub batch_only: bool, - - /// Show only single-model jobs - #[clap(long, conflicts_with = "batch_only")] - pub single_only: bool, -} - -impl ListCommand { - /// Execute the list command - pub async fn run(&self, api_gateway_url: &str, jwt_token: &str) -> Result<()> { - // Connect to API Gateway - let channel = crate::client::connect_channel(api_gateway_url).await?; - - let mut client = MlTrainingServiceClient::new(channel); - - // Build request with filters - let status_filter = self.parse_status_filter()?; - let list_request = self.build_request(status_filter); - - // Add JWT authentication - let mut request = Request::new(list_request); - let token_value = MetadataValue::try_from(format!("Bearer {}", jwt_token)) - .context("Invalid JWT token format")?; - request - .metadata_mut() - .insert("authorization", token_value); - - // Call gRPC service - let response = client - .list_training_jobs(request) - .await - .context("Failed to list training jobs")? - .into_inner(); - - // Apply client-side filtering and sorting - let mut jobs = response.jobs.clone(); - jobs = self.apply_filters(jobs); - jobs = self.apply_sorting(jobs); - jobs.truncate(self.limit as usize); - - // Display results - if jobs.is_empty() { - println!("No training jobs found."); - return Ok(()); - } - - let total_count = response.total_count; - self.display_jobs_table(&jobs)?; - self.display_summary(&jobs, total_count); - - Ok(()) - } - - /// Parse status filter string to enum - fn parse_status_filter(&self) -> Result { - if let Some(status_str) = &self.status { - match status_str.to_uppercase().as_str() { - "PENDING" => Ok(TrainingStatus::Pending), - "RUNNING" => Ok(TrainingStatus::Running), - "COMPLETED" => Ok(TrainingStatus::Completed), - "FAILED" => Ok(TrainingStatus::Failed), - "STOPPED" => Ok(TrainingStatus::Stopped), - _ => Ok(TrainingStatus::Unknown), - } - } else { - Ok(TrainingStatus::Unknown) - } - } - - /// Build gRPC request - fn build_request(&self, status_filter: TrainingStatus) -> ListTrainingJobsRequest { - ListTrainingJobsRequest { - page: 1, - page_size: 500, // Fetch more for client-side filtering - status_filter: status_filter as i32, - model_type_filter: self.model.clone().unwrap_or_default(), - start_time: 0, - end_time: 0, - } - } - - /// Apply client-side filters - fn apply_filters(&self, mut jobs: Vec) -> Vec { - // Filter by asset (check tags or description) - if let Some(asset) = &self.asset { - jobs.retain(|job| { - (job.tags.get("asset") == Some(asset)) - || job.description.contains(asset) - }); - } - - // Filter batch vs single jobs - if self.batch_only { - jobs.retain(|job| job.model_type.to_uppercase() == "BATCH" || job.job_id.starts_with("batch_")); - } else if self.single_only { - jobs.retain(|job| job.model_type.to_uppercase() != "BATCH" && !job.job_id.starts_with("batch_")); - } else { - // No batch/single filter applied - } - - jobs - } - - /// Apply sorting - fn apply_sorting(&self, mut jobs: Vec) -> Vec { - match self.sort_by.as_str() { - "start_time" => { - jobs.sort_by_key(|job| job.started_at); - } - "duration" => { - jobs.sort_by_key(|job| { - if job.completed_at > 0 { - job.completed_at - job.started_at - } else { - 0 - } - }); - } - "status" => { - jobs.sort_by_key(|job| job.status); - } - _ => { - jobs.sort_by_key(|job| job.started_at); - } - } - - // Apply sort order - if self.sort_order == "desc" { - jobs.reverse(); - } - - jobs - } - - /// Display jobs in formatted table - fn display_jobs_table(&self, jobs: &[TrainingJobSummary]) -> Result<()> { - let mut table = Table::new(); - table - .load_preset(UTF8_FULL) - .set_content_arrangement(ContentArrangement::Dynamic); - - // Table header - table.set_header(vec![ - Cell::new("Job ID") - .add_attribute(Attribute::Bold) - .fg(Color::Cyan), - Cell::new("Status") - .add_attribute(Attribute::Bold) - .fg(Color::Cyan), - Cell::new("Type") - .add_attribute(Attribute::Bold) - .fg(Color::Cyan), - Cell::new("Model") - .add_attribute(Attribute::Bold) - .fg(Color::Cyan), - Cell::new("Asset(s)") - .add_attribute(Attribute::Bold) - .fg(Color::Cyan), - Cell::new("Duration") - .add_attribute(Attribute::Bold) - .fg(Color::Cyan), - ]); - - // Add rows - for job in jobs { - let status_str = self.format_status(job.status); - let job_type = self.determine_job_type(&job.job_id, &job.model_type); - let model = job.model_type.clone(); - let assets = job - .tags - .get("asset") - .cloned() - .unwrap_or_else(|| "N/A".to_owned()); - let duration = self.format_duration(job.started_at, job.completed_at); - - table.add_row(vec![ - Cell::new(&job.job_id), - Cell::new(status_str), - Cell::new(job_type), - Cell::new(model), - Cell::new(assets), - Cell::new(duration), - ]); - } - - println!("{}", table); - Ok(()) - } - - /// Format status with emoji - fn format_status(&self, status: i32) -> String { - match TrainingStatus::try_from(status).unwrap_or(TrainingStatus::Unknown) { - TrainingStatus::Pending => "\u{23f3} PENDING".to_owned(), - TrainingStatus::Running => "\u{23f3} RUNNING".to_owned(), - TrainingStatus::Completed => "\u{2705} COMPLETE".to_owned(), - TrainingStatus::Failed => "\u{274c} FAILED".to_owned(), - TrainingStatus::Stopped => "\u{1f6d1} STOPPED".to_owned(), - TrainingStatus::Paused => "\u{23f8} PAUSED".to_owned(), - TrainingStatus::Unknown => "\u{2753} UNKNOWN".to_owned(), - } - } - - /// Determine if job is batch or single - fn determine_job_type(&self, job_id: &str, model_type: &str) -> &'static str { - if job_id.starts_with("batch_") || model_type.to_uppercase() == "BATCH" { - "Batch" - } else { - "Single" - } - } - - /// Format duration - fn format_duration(&self, started_at: i64, completed_at: i64) -> String { - if started_at == 0 { - return "N/A".to_owned(); - } - - let duration_secs = if completed_at > 0 { - completed_at - started_at - } else { - // Job still running, calculate elapsed time - let now = Utc::now().timestamp(); - now - started_at - }; - - let abs_secs = duration_secs.unsigned_abs(); - if abs_secs < 60 { - format!("{}s", abs_secs) - } else if abs_secs < 3600 { - let mins = abs_secs / 60; - let secs = abs_secs % 60; - format!("{}m {}s", mins, secs) - } else { - let hours = abs_secs / 3600; - let mins = (abs_secs % 3600) / 60; - format!("{}h {}m", hours, mins) - } - } - - /// Display summary - fn display_summary( - &self, - jobs: &[TrainingJobSummary], - total_count: u32, - ) { - println!(); - println!("Total: {} jobs (showing {} results)", total_count, jobs.len()); - - // Show filter info - let mut filters = vec![]; - if let Some(status) = &self.status { - filters.push(format!("status={}", status)); - } - if let Some(model) = &self.model { - filters.push(format!("model={}", model)); - } - if let Some(asset) = &self.asset { - filters.push(format!("asset={}", asset)); - } - if self.batch_only { - filters.push("type=batch".to_owned()); - } - if self.single_only { - filters.push("type=single".to_owned()); - } - - if !filters.is_empty() { - println!("Filters: {}", filters.join(", ")); - } else { - println!("(Showing last {} jobs)", self.limit); - } - } -} diff --git a/bin/fxt/src/commands/train/mod.rs b/bin/fxt/src/commands/train/mod.rs deleted file mode 100644 index fcc2a863d..000000000 --- a/bin/fxt/src/commands/train/mod.rs +++ /dev/null @@ -1,93 +0,0 @@ -//! TLI Train Command Module -//! -//! Provides CLI interface for managing ML model training jobs. -//! -//! # Subcommands -//! - `list` - List training jobs with filtering and sorting -//! - `status` - Query individual training job status (stub for Wave 2 Agent 5) - -pub mod list; -pub mod monitor; -pub mod progress_tracker; -pub mod start; -pub mod status; - -// Re-export progress tracker types -pub use progress_tracker::{JobProgress, JobStatus, ProgressTracker}; - -pub use list::ListCommand; -pub use monitor::MonitorCommand; -pub use start::StartCommand; -pub use status::StatusCommand; - -use anyhow::Result; -use clap::Subcommand; - -/// Train command arguments -#[derive(Subcommand, Debug)] -pub enum TrainCommand { - /// List training jobs with filtering - List { - #[clap(flatten)] - list_args: ListCommand, - }, - - /// Start an on-demand training job - Start { - #[clap(flatten)] - start_args: start::StartCommand, - }, - - /// Query training job status (snapshot) - #[clap( - long_about = "Query current status of a training job.\n\n\ - Shows:\n\ - - Job status (PENDING/RUNNING/COMPLETED/FAILED/STOPPED)\n\ - - Training progress and timing\n\ - - Performance metrics (Sharpe, hit rate, drawdown)\n\ - - Resource usage (GPU memory, elapsed time)\n\ - - Model checkpoint path (for completed jobs)\n\ - - Error details (for failed jobs)\n\n\ - Examples:\n\ - tli train status train_tft_es_fut_20251022_143021\n\ - tli train status train_batch_001 --verbose" - )] - Status { - #[clap(flatten)] - status_args: StatusCommand, - }, - - /// Live training metrics monitor (Prometheus via gRPC) - #[clap( - long_about = "Stream live training metrics from the monitoring service.\n\n\ - Shows:\n\ - - Per-model epoch, loss, validation loss, throughput\n\ - - GPU utilization, VRAM, temperature, power\n\ - - Hyperopt trial progress and best objective\n\ - - Health counters (NaN, gradient explosions, checkpoints)\n\n\ - Examples:\n\ - fxt train monitor\n\ - fxt train monitor --once\n\ - fxt train monitor --model dqn --interval 5" - )] - Monitor { - #[clap(flatten)] - monitor_args: MonitorCommand, - }, -} - -/// Execute train command -pub async fn execute_train_command( - command: TrainCommand, - api_gateway_url: &str, - jwt_token: &str, -) -> Result<()> { - match command { - TrainCommand::List { list_args } => list_args.run(api_gateway_url, jwt_token).await, - TrainCommand::Start { start_args } => start_args.run(api_gateway_url, jwt_token).await, - TrainCommand::Status { status_args } => status_args.run(api_gateway_url, jwt_token).await, - TrainCommand::Monitor { monitor_args } => { - monitor_args.run(api_gateway_url, jwt_token).await - } - } -} diff --git a/bin/fxt/src/commands/train/monitor.rs b/bin/fxt/src/commands/train/monitor.rs deleted file mode 100644 index 56c32d899..000000000 --- a/bin/fxt/src/commands/train/monitor.rs +++ /dev/null @@ -1,274 +0,0 @@ -//! fxt train monitor -- live training metrics from monitoring service -//! -//! Connects to the MonitoringService gRPC endpoint and displays a live TUI -//! with per-model epoch progress, GPU telemetry, and health counters. -//! Use `--once` for a single snapshot instead of continuous streaming. - -use anyhow::{Context, Result}; -use clap::Args; -use colored::Colorize; -use std::io::{self, Write}; - -use crate::proto::monitoring::{ - monitoring_service_client::MonitoringServiceClient, GetLiveTrainingMetricsRequest, - GetLiveTrainingMetricsResponse, StreamTrainingMetricsRequest, TrainingSession, -}; - -/// Live training metrics monitor -#[derive(Args, Debug)] -pub struct MonitorCommand { - /// Print a single snapshot and exit (no live TUI) - #[arg(long)] - pub once: bool, - - /// Filter to a specific model (e.g. "dqn", "ppo", "tft") - #[arg(long)] - pub model: Option, - - /// Streaming interval in seconds (default: 3) - #[arg(long, default_value = "3")] - pub interval: u32, -} - -impl MonitorCommand { - /// Execute the monitor command -- connect to monitoring service and display metrics - pub async fn run(&self, api_gateway_url: &str, _jwt_token: &str) -> Result<()> { - // monitoring_service runs on a different port -- derive URL from env or default. - // When connecting via Tailscale (*.fxhnt.ai), use monitor.fxhnt.ai:443 (nginx proxy). - // Otherwise, replace the api-gateway port with the monitoring port. - let monitoring_url = std::env::var("MONITORING_SERVICE_URL").unwrap_or_else(|_| { - if api_gateway_url.contains("fxhnt.ai") { - "https://monitor.fxhnt.ai".to_owned() - } else { - api_gateway_url.replace(":50050", ":50057") - } - }); - - let mut client = MonitoringServiceClient::connect(monitoring_url) - .await - .context("Failed to connect to monitoring service")?; - - let model_filter = self.model.clone().unwrap_or_default(); - - if self.once { - let response = client - .get_live_training_metrics(tonic::Request::new(GetLiveTrainingMetricsRequest { - model_filter, - })) - .await - .context("GetLiveTrainingMetrics failed")?; - render_snapshot(&response.into_inner()); - } else { - println!("Connecting to monitoring service..."); - let response = client - .stream_training_metrics(tonic::Request::new(StreamTrainingMetricsRequest { - model_filter, - interval_seconds: self.interval, - })) - .await - .context("StreamTrainingMetrics failed")?; - - let mut stream = response.into_inner(); - while let Some(msg) = stream.message().await? { - render_tui(&msg); - } - } - - Ok(()) - } -} - -fn render_snapshot(resp: &GetLiveTrainingMetricsResponse) { - if let Some(gpu) = &resp.gpu { - println!( - "GPU: {:.0}% util | {:.1}/{:.1} GB VRAM | {:.0}C | {:.0}W", - gpu.utilization_percent, - gpu.memory_used_mb / 1024.0, - gpu.memory_total_mb / 1024.0, - gpu.temperature_celsius, - gpu.power_watts, - ); - } - println!("Active K8s training jobs: {}", resp.active_k8s_jobs); - println!(); - - if resp.sessions.is_empty() { - println!("No active training sessions with metrics."); - return; - } - - print_session_table(&resp.sessions); - print_rl_diagnostics(&resp.sessions); - print_hyperopt_summary(&resp.sessions); - print_health_summary(&resp.sessions); - print_financial_metrics(&resp.sessions); -} - -fn render_tui(resp: &GetLiveTrainingMetricsResponse) { - // Clear screen - print!("\x1B[2J\x1B[H"); - // Intentionally ignoring flush errors on stdout (terminal output, non-critical) - drop(io::stdout().flush()); - - println!("{}", "=".repeat(72).bright_black()); - println!( - "{}", - " Foxhunt Training Monitor" - .bright_white() - .bold() - ); - println!("{}", "=".repeat(72).bright_black()); - println!(); - - render_snapshot(resp); - - println!(); - println!("{}", "Ctrl+C to exit".bright_black()); -} - -fn print_session_table(sessions: &[TrainingSession]) { - println!( - "{:<10} {:<6} {:<7} {:<10} {:<10} {:<9} {:<10} {:<10}", - "Model".bright_cyan(), - "Fold".bright_cyan(), - "Epoch".bright_cyan(), - "Loss".bright_cyan(), - "Val Loss".bright_cyan(), - "Batch/s".bright_cyan(), - "Grad Norm".bright_cyan(), - "Eval Acc".bright_cyan(), - ); - println!("{}", "-".repeat(82).bright_black()); - - for s in sessions { - let acc = if s.eval_accuracy > 0.0 { - format!("{:.1}%", s.eval_accuracy * 100.0) - } else { - "-".to_owned() - }; - let grad = if s.gradient_norm > 0.0 { - format!("{:.4}", s.gradient_norm) - } else { - "-".to_owned() - }; - println!( - "{:<10} {:<6} {:<7.0} {:<10.4} {:<10.4} {:<9.1} {:<10} {:<10}", - s.model, s.fold, s.current_epoch, s.epoch_loss, s.validation_loss, - s.batches_per_second, grad, acc, - ); - } -} - -fn print_rl_diagnostics(sessions: &[TrainingSession]) { - let rl_sessions: Vec<_> = sessions - .iter() - .filter(|s| s.q_value_mean != 0.0 || s.policy_entropy != 0.0) - .collect(); - if rl_sessions.is_empty() { - return; - } - println!(); - println!("{}", "RL Diagnostics:".bright_cyan()); - for s in &rl_sessions { - if s.q_value_mean != 0.0 || s.q_value_max != 0.0 { - println!( - " {}: Q-mean={:.2} Q-max={:.2} | buffer={}", - s.model.bright_white(), - s.q_value_mean, - s.q_value_max, - s.replay_buffer_size, - ); - } - if s.policy_entropy != 0.0 || s.kl_divergence != 0.0 { - println!( - " {}: entropy={:.3} KL={:.4} adv-mean={:.4}", - s.model.bright_white(), - s.policy_entropy, - s.kl_divergence, - s.advantage_mean, - ); - } - } -} - -fn print_hyperopt_summary(sessions: &[TrainingSession]) { - let hyperopt: Vec<_> = sessions.iter().filter(|s| s.is_hyperopt).collect(); - if hyperopt.is_empty() { - return; - } - println!(); - for s in &hyperopt { - let trial_detail = if s.hyperopt_trial_epoch > 0 { - format!( - " (epoch {}, loss {:.4})", - s.hyperopt_trial_epoch, s.hyperopt_trial_best_loss - ) - } else { - String::new() - }; - let elapsed = if s.hyperopt_elapsed_seconds > 0.0 { - format!(" | {:.0}s elapsed", s.hyperopt_elapsed_seconds) - } else { - String::new() - }; - println!( - "Hyperopt ({}): trial {}/{}{} | best Sharpe {:.2} | {} failures{}", - s.model.bright_magenta(), - s.hyperopt_trial_current, - s.hyperopt_trial_total, - trial_detail, - s.hyperopt_best_objective, - s.hyperopt_trials_failed, - elapsed, - ); - } -} - -fn print_health_summary(sessions: &[TrainingSession]) { - let total_nan: u32 = sessions.iter().map(|s| s.nan_detected).sum(); - let total_grad: u32 = sessions.iter().map(|s| s.gradient_explosions).sum(); - let total_ckpt: u32 = sessions.iter().map(|s| s.checkpoint_saves).sum(); - println!( - "Health: {} NaN | {} grad explosions | {} checkpoints", - total_nan, total_grad, total_ckpt - ); -} - -fn print_financial_metrics(sessions: &[TrainingSession]) { - let financial: Vec<_> = sessions - .iter() - .filter(|s| s.epoch_sharpe != 0.0 || s.epoch_win_rate > 0.0) - .collect(); - if financial.is_empty() { - return; - } - println!(); - println!("{}", "Epoch Financial Metrics:".bright_cyan()); - for s in &financial { - let sharpe_colored = if s.epoch_sharpe >= 2.0 { - format!("{:.2}", s.epoch_sharpe).green() - } else if s.epoch_sharpe >= 1.0 { - format!("{:.2}", s.epoch_sharpe).yellow() - } else { - format!("{:.2}", s.epoch_sharpe).red() - }; - println!( - " {}: Sharpe={} WinRate={:.1}% MaxDD={:.1}% PF={:.2} Return={:+.2}% Trades={}", - s.model.bright_white(), - sharpe_colored, - s.epoch_win_rate * 100.0, - s.epoch_max_drawdown * 100.0, - s.epoch_profit_factor, - s.epoch_total_return * 100.0, - s.epoch_total_trades, - ); - if s.action_buy_pct > 0.0 || s.action_sell_pct > 0.0 { - println!( - " Actions: BUY {:.0}% | SELL {:.0}% | HOLD {:.0}%", - s.action_buy_pct * 100.0, - s.action_sell_pct * 100.0, - s.action_hold_pct * 100.0, - ); - } - } -} diff --git a/bin/fxt/src/commands/train/progress_tracker.rs b/bin/fxt/src/commands/train/progress_tracker.rs deleted file mode 100644 index f60ac4c2c..000000000 --- a/bin/fxt/src/commands/train/progress_tracker.rs +++ /dev/null @@ -1,445 +0,0 @@ -//! Hierarchical Progress Tracker for ML Training Jobs -//! -//! This module provides thread-safe progress tracking for batch ML training operations, -//! with support for hierarchical parent-child job relationships and weighted progress -//! calculations. -//! -//! # Features -//! -//! - Thread-safe progress tracking using `Arc>` -//! - Pull-based rendering with ANSI screen clearing -//! - Hierarchical display: parent batch job → child model jobs -//! - Weighted progress calculation based on model training complexity -//! -//! # Architecture -//! -//! The progress tracker maintains two levels of jobs: -//! - **Parent Jobs**: Represent batch training operations (e.g., "train all models") -//! - **Child Jobs**: Individual model training jobs (DQN, PPO, MAMBA-2, TFT) -//! -//! Progress is calculated using weighted averaging based on model complexity: -//! - DQN: 10% (fastest, simplest model) -//! - PPO: 30% (moderate complexity) -//! - MAMBA-2: 40% (high complexity, GPU-intensive) -//! - TFT: 20% (moderate complexity with INT8 quantization) - -use std::collections::HashMap; -use std::io::{self, Write}; -use std::sync::Arc; -use std::time::Duration; -use tokio::sync::Mutex; -use tokio::time; - -/// Job status for tracking progress. -#[derive(Debug, Clone, PartialEq)] -pub enum JobStatus { - /// Job is queued but not started - Pending, - /// Job is currently running - Running, - /// Job completed successfully - Completed, - /// Job failed with error - Failed(String), -} - -impl std::fmt::Display for JobStatus { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - JobStatus::Pending => write!(f, "PENDING"), - JobStatus::Running => write!(f, "RUNNING"), - JobStatus::Completed => write!(f, "\u{2713} DONE"), - JobStatus::Failed(err) => write!(f, "\u{2717} FAILED: {}", err), - } - } -} - -/// Progress information for a single job. -#[derive(Debug, Clone)] -pub struct JobProgress { - /// Job ID - pub job_id: String, - /// Job name/description - pub name: String, - /// Current status - pub status: JobStatus, - /// Progress percentage (0-100) - pub progress_pct: u8, - /// Optional status message - pub message: Option, - /// Parent job ID (if this is a child job) - pub parent_id: Option, - /// Weight for weighted progress calculation (0.0-1.0) - pub weight: f32, -} - -impl JobProgress { - /// Create a new job progress tracker. - pub const fn new(job_id: String, name: String, weight: f32) -> Self { - Self { - job_id, - name, - status: JobStatus::Pending, - progress_pct: 0, - message: None, - parent_id: None, - weight, - } - } - - /// Create a child job with a parent. - pub fn with_parent(mut self, parent_id: String) -> Self { - self.parent_id = Some(parent_id); - self - } -} - -/// Thread-safe hierarchical progress tracker. -/// -/// This tracker maintains a map of all jobs and provides methods for updating -/// progress, calculating weighted parent progress, and rendering the progress -/// tree to the terminal. -#[derive(Clone)] -pub struct ProgressTracker { - /// Map of job ID to progress information - jobs: Arc>>, -} - -impl ProgressTracker { - /// Create a new progress tracker. - pub fn new() -> Self { - Self { - jobs: Arc::new(Mutex::new(HashMap::new())), - } - } - - /// Add a new job to track. - /// - /// # Arguments - /// - /// * `job` - Job progress to track - pub async fn add_job(&self, job: JobProgress) { - let mut jobs = self.jobs.lock().await; - jobs.insert(job.job_id.clone(), job); - } - - /// Update job progress. - /// - /// # Arguments - /// - /// * `job_id` - ID of the job to update - /// * `progress_pct` - New progress percentage (0-100) - pub async fn update_progress(&self, job_id: &str, progress_pct: u8) { - let mut jobs = self.jobs.lock().await; - if let Some(job) = jobs.get_mut(job_id) { - job.progress_pct = progress_pct.min(100); - if progress_pct >= 100 { - job.status = JobStatus::Completed; - } else if matches!(job.status, JobStatus::Pending) { - job.status = JobStatus::Running; - } else { - // Status unchanged - } - } - } - - /// Update job status. - /// - /// # Arguments - /// - /// * `job_id` - ID of the job to update - /// * `status` - New status - pub async fn update_status(&self, job_id: &str, status: JobStatus) { - let mut jobs = self.jobs.lock().await; - if let Some(job) = jobs.get_mut(job_id) { - job.status = status; - } - } - - /// Update job message. - /// - /// # Arguments - /// - /// * `job_id` - ID of the job to update - /// * `message` - Status message - pub async fn update_message(&self, job_id: &str, message: String) { - let mut jobs = self.jobs.lock().await; - if let Some(job) = jobs.get_mut(job_id) { - job.message = Some(message); - } - } - - /// Calculate weighted progress for a parent job based on its children. - /// - /// This method sums the weighted progress of all child jobs to compute - /// the parent's overall progress. - /// - /// # Arguments - /// - /// * `parent_id` - ID of the parent job - /// - /// # Returns - /// - /// Returns the weighted progress percentage (0-100). - async fn calculate_parent_progress(&self, parent_id: &str) -> u8 { - let jobs = self.jobs.lock().await; - - let children: Vec<_> = jobs - .values() - .filter(|j| j.parent_id.as_deref() == Some(parent_id)) - .collect(); - - if children.is_empty() { - return 0; - } - - let weighted_sum: f32 = children - .iter() - .map(|child| (child.progress_pct as f32 / 100.0) * child.weight) - .sum(); - - (weighted_sum * 100.0) as u8 - } - - /// Render the progress tree to stdout. - /// - /// This method uses ANSI escape codes to clear the screen and render - /// a hierarchical view of all jobs with their current progress. - pub async fn render(&self) -> io::Result<()> { - // Clear screen using ANSI escape codes - print!("\x1B[2J\x1B[H"); - io::stdout().flush()?; - - println!("\u{2554}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2557}"); - println!("\u{2551} Foxhunt ML Training Progress Tracker \u{2551}"); - println!("\u{255a}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{2550}\u{255d}"); - println!(); - - // Collect all job data in a single lock acquisition - let jobs = self.jobs.lock().await; - let mut parent_jobs: Vec<_> = jobs - .values() - .filter(|j| j.parent_id.is_none()) - .cloned() - .collect(); - parent_jobs.sort_by(|a, b| a.job_id.cmp(&b.job_id)); - - let child_jobs: Vec<_> = jobs.values().cloned().collect(); - drop(jobs); // Release lock - - for parent in parent_jobs { - // Calculate parent progress from children - let parent_progress = self.calculate_parent_progress(&parent.job_id).await; - - // Render parent job - println!("\u{250c}\u{2500} {} [{}%]", parent.name, parent_progress); - println!("\u{2502} Status: {}", parent.status); - if let Some(msg) = &parent.message { - println!("\u{2502} {}", msg); - } - - // Find and render children - let mut children: Vec<_> = child_jobs - .iter() - .filter(|j| j.parent_id.as_deref() == Some(&parent.job_id)) - .collect(); - children.sort_by(|a, b| a.job_id.cmp(&b.job_id)); - - for (idx, child) in children.iter().enumerate() { - let is_last = idx == children.len() - 1; - let prefix = if is_last { "\u{2514}\u{2500}\u{2500}" } else { "\u{251c}\u{2500}\u{2500}" }; - - println!( - "\u{2502} {} {} [{}%] {} (weight: {:.0}%)", - prefix, - child.name, - child.progress_pct, - child.status, - child.weight * 100.0 - ); - - if let Some(msg) = &child.message { - let cont = if is_last { " " } else { "\u{2502} " }; - println!("\u{2502} {} {}", cont, msg); - } - } - - println!(); - } - - Ok(()) - } - - /// Start a background rendering loop that updates the display every 2 seconds. - /// - /// This method spawns a tokio task that continuously renders the progress - /// tree until the tracker is dropped. - /// - /// # Returns - /// - /// Returns a `tokio::task::JoinHandle` for the rendering task. - #[allow(clippy::infinite_loop)] - pub fn start_rendering(self) -> tokio::task::JoinHandle<()> { - tokio::spawn(async move { - let mut interval = time::interval(Duration::from_secs(2)); - loop { - interval.tick().await; - if let Err(e) = self.render().await { - eprintln!("Progress rendering error: {}", e); - } - } - }) - } -} - -impl Default for ProgressTracker { - fn default() -> Self { - Self::new() - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[tokio::test] - async fn test_progress_tracker_new() { - let tracker = ProgressTracker::new(); - let jobs = tracker.jobs.lock().await; - assert!(jobs.is_empty()); - } - - #[tokio::test] - async fn test_add_job() { - let tracker = ProgressTracker::new(); - let job = JobProgress::new( - "job1".to_string(), - "Test Job".to_string(), - 1.0, - ); - - tracker.add_job(job).await; - - let jobs = tracker.jobs.lock().await; - assert_eq!(jobs.len(), 1); - assert!(jobs.contains_key("job1")); - } - - #[tokio::test] - async fn test_update_progress() { - let tracker = ProgressTracker::new(); - let job = JobProgress::new( - "job1".to_string(), - "Test Job".to_string(), - 1.0, - ); - - tracker.add_job(job).await; - tracker.update_progress("job1", 50).await; - - let jobs = tracker.jobs.lock().await; - let job = jobs.get("job1").unwrap(); - assert_eq!(job.progress_pct, 50); - assert_eq!(job.status, JobStatus::Running); - } - - #[tokio::test] - async fn test_update_progress_completion() { - let tracker = ProgressTracker::new(); - let job = JobProgress::new( - "job1".to_string(), - "Test Job".to_string(), - 1.0, - ); - - tracker.add_job(job).await; - tracker.update_progress("job1", 100).await; - - let jobs = tracker.jobs.lock().await; - let job = jobs.get("job1").unwrap(); - assert_eq!(job.progress_pct, 100); - assert_eq!(job.status, JobStatus::Completed); - } - - #[tokio::test] - async fn test_update_status() { - let tracker = ProgressTracker::new(); - let job = JobProgress::new( - "job1".to_string(), - "Test Job".to_string(), - 1.0, - ); - - tracker.add_job(job).await; - tracker.update_status("job1", JobStatus::Failed("Test error".to_string())).await; - - let jobs = tracker.jobs.lock().await; - let job = jobs.get("job1").unwrap(); - assert!(matches!(job.status, JobStatus::Failed(_))); - } - - #[tokio::test] - async fn test_weighted_progress_calculation() { - let tracker = ProgressTracker::new(); - - // Create parent job - let parent = JobProgress::new( - "parent".to_string(), - "Batch Training".to_string(), - 1.0, - ); - tracker.add_job(parent).await; - - // Create child jobs with different weights matching model complexity - let dqn = JobProgress::new( - "dqn".to_string(), - "DQN".to_string(), - 0.10, // 10% weight - ).with_parent("parent".to_string()); - - let ppo = JobProgress::new( - "ppo".to_string(), - "PPO".to_string(), - 0.30, // 30% weight - ).with_parent("parent".to_string()); - - let mamba2 = JobProgress::new( - "mamba2".to_string(), - "MAMBA-2".to_string(), - 0.40, // 40% weight - ).with_parent("parent".to_string()); - - let tft = JobProgress::new( - "tft".to_string(), - "TFT-INT8".to_string(), - 0.20, // 20% weight - ).with_parent("parent".to_string()); - - tracker.add_job(dqn).await; - tracker.add_job(ppo).await; - tracker.add_job(mamba2).await; - tracker.add_job(tft).await; - - // Set progress: DQN 100%, PPO 50%, MAMBA-2 25%, TFT 0% - tracker.update_progress("dqn", 100).await; - tracker.update_progress("ppo", 50).await; - tracker.update_progress("mamba2", 25).await; - tracker.update_progress("tft", 0).await; - - // Calculate weighted progress - // Expected: (100 * 0.10) + (50 * 0.30) + (25 * 0.40) + (0 * 0.20) = 10 + 15 + 10 + 0 = 35% - let parent_progress = tracker.calculate_parent_progress("parent").await; - assert_eq!(parent_progress, 35); - } - - #[tokio::test] - async fn test_job_status_display() { - assert_eq!(JobStatus::Pending.to_string(), "PENDING"); - assert_eq!(JobStatus::Running.to_string(), "RUNNING"); - assert_eq!(JobStatus::Completed.to_string(), "✓ DONE"); - assert_eq!( - JobStatus::Failed("error".to_string()).to_string(), - "✗ FAILED: error" - ); - } -} diff --git a/bin/fxt/src/commands/train/start.rs b/bin/fxt/src/commands/train/start.rs deleted file mode 100644 index e2f0ec0ff..000000000 --- a/bin/fxt/src/commands/train/start.rs +++ /dev/null @@ -1,194 +0,0 @@ -//! fxt train start -- trigger on-demand training job -//! -//! Sends a `StartTraining` gRPC request to the ML Training Service via the API Gateway -//! and prints the resulting job ID so the user can track progress with `fxt train status`. - -use anyhow::{Context, Result}; -use clap::Args; -use colored::Colorize; -use std::collections::HashMap; - -use crate::proto::ml_training::{ - data_source::Source, ml_training_service_client::MlTrainingServiceClient, DataSource, - StartTrainingRequest, TrainingStatus, -}; -use tonic::{metadata::MetadataValue, Request}; - -/// Start an on-demand training job -#[derive(Args, Debug)] -pub struct StartCommand { - /// Model type (dqn, ppo, tft, mamba2, tggn, tlob, liquid, kan, xlstm, diffusion) - pub model: String, - - /// Symbol to train on (e.g. ES.FUT) - pub symbol: String, - - /// Number of training epochs - #[arg(long, default_value = "50")] - pub epochs: u32, - - /// Use GPU acceleration - #[arg(long, default_value = "true")] - pub use_gpu: bool, - - /// Optional job description - #[arg(long)] - pub description: Option, - - /// Data source file path (defaults to Databento cache for the symbol) - #[arg(long)] - pub data_path: Option, - - /// Run hyperopt (PSO hyperparameter optimization) instead of training. - /// Uses hyperopt_baseline_rl for DQN/PPO, hyperopt_baseline_supervised for others. - #[arg(long, default_value = "false")] - pub hyperopt: bool, - - /// Number of PSO trials (only used with --hyperopt) - #[arg(long, default_value = "20")] - pub trials: u32, - - /// Parallel trial evaluations (only with --hyperopt, 0 = auto-detect) - #[arg(long, default_value = "0")] - pub parallel: u32, -} - -impl StartCommand { - /// Execute the start training command via gRPC - pub async fn run(&self, api_gateway_url: &str, jwt_token: &str) -> Result<()> { - let model_upper = self.model.to_uppercase(); - validate_model_type(&model_upper)?; - - println!( - "Requesting training: model={}, symbol={}, epochs={}", - model_upper.bright_magenta(), - self.symbol.bright_cyan(), - self.epochs, - ); - - // Connect to API Gateway - let channel = crate::client::connect_channel(api_gateway_url).await?; - - let mut client = MlTrainingServiceClient::new(channel); - - // Build data source -- default to Databento cache path for the symbol - let default_path = format!("data/cache/futures-baseline/{}", self.symbol); - let file_path = self - .data_path - .clone() - .unwrap_or_else(|| default_path.clone()); - - let data_source = DataSource { - source: Some(Source::FilePath(file_path)), - start_time: 0, - end_time: 0, - }; - - // Build tags - let mut tags: HashMap = HashMap::new(); - tags.insert("asset".to_owned(), self.symbol.clone()); - tags.insert("source".to_owned(), "fxt-cli".to_owned()); - if self.hyperopt { - tags.insert("mode".to_owned(), "hyperopt".to_owned()); - tags.insert("trials".to_owned(), self.trials.to_string()); - tags.insert("parallel".to_owned(), self.parallel.to_string()); - } - - let description = self.description.clone().unwrap_or_else(|| { - if self.hyperopt { - format!("Hyperopt {} on {} ({} trials)", model_upper, self.symbol, self.trials) - } else { - format!("On-demand {} training on {}", model_upper, self.symbol) - } - }); - - // Hyperparameters are left empty -- the service will apply defaults for the model type. - // Per-model params can be added later via --config flag. - let request_body = StartTrainingRequest { - model_type: model_upper.clone(), - data_source: Some(data_source), - hyperparameters: None, - use_gpu: self.use_gpu, - description, - tags, - }; - - // Add JWT authentication - let mut request = Request::new(request_body); - let token_value = MetadataValue::try_from(format!("Bearer {}", jwt_token)) - .context("Invalid JWT token format")?; - request - .metadata_mut() - .insert("authorization", token_value); - - // Call gRPC service - let response = client - .start_training(request) - .await - .context("Failed to start training job")? - .into_inner(); - - // Display result - let status_str = format_status(response.status); - println!("\nTraining job started: {}", response.job_id.bright_green()); - println!(" Status: {}", status_str); - if self.hyperopt { - println!(" Mode: {}", "HYPEROPT".bright_yellow()); - println!(" Trials: {}", self.trials); - } - println!(" Model: {}", model_upper.bright_magenta()); - println!(" Symbol: {}", self.symbol.bright_cyan()); - if !self.hyperopt { - println!(" Epochs: {}", self.epochs); - } - println!(" GPU: {}", self.use_gpu); - if !response.message.is_empty() { - println!(" Message: {}", response.message); - } - println!( - "\nTrack progress: fxt train status {}", - response.job_id.bright_yellow() - ); - - Ok(()) - } -} - -/// Validate that the model type is one of the supported types -fn validate_model_type(model: &str) -> Result<()> { - const VALID_MODELS: &[&str] = &[ - "DQN", - "PPO", - "TFT", - "MAMBA2", - "MAMBA_2", - "TGGN", - "TLOB", - "LIQUID", - "KAN", - "XLSTM", - "DIFFUSION", - ]; - - if VALID_MODELS.contains(&model) { - Ok(()) - } else { - anyhow::bail!( - "Unsupported model type '{}'. Valid types: dqn, ppo, tft, mamba2, tggn, tlob, liquid, kan, xlstm, diffusion", - model, - ); - } -} - -/// Format TrainingStatus enum value to display string -fn format_status(status: i32) -> String { - match TrainingStatus::try_from(status).ok() { - Some(TrainingStatus::Pending) => "PENDING".to_owned(), - Some(TrainingStatus::Running) => "RUNNING".to_owned(), - Some(TrainingStatus::Completed) => "COMPLETED".to_owned(), - Some(TrainingStatus::Failed) => "FAILED".to_owned(), - Some(TrainingStatus::Stopped) => "STOPPED".to_owned(), - Some(TrainingStatus::Paused) => "PAUSED".to_owned(), - Some(TrainingStatus::Unknown) | None => "UNKNOWN".to_owned(), - } -} diff --git a/bin/fxt/src/commands/train/status.rs b/bin/fxt/src/commands/train/status.rs deleted file mode 100644 index bfdb52c79..000000000 --- a/bin/fxt/src/commands/train/status.rs +++ /dev/null @@ -1,311 +0,0 @@ -//! TLI Train Status Command - Training Job Status Queries -//! -//! Provides CLI interface for querying ML training job status through the API Gateway. - -use anyhow::{Context, Result as AnyhowResult}; -use chrono::{DateTime, Utc}; -use clap::Parser; -use colored::Colorize; -use comfy_table::{Cell, Color, Table}; - -// Import ML training proto types -use crate::proto::ml_training::{ - ml_training_service_client::MlTrainingServiceClient, GetTrainingJobDetailsRequest, - TrainingJobDetails, TrainingStatus, -}; - -/// Train Status command arguments -#[derive(Parser, Debug)] -pub struct StatusCommand { - /// Training job ID to query - #[arg(value_name = "JOB_ID")] - pub job_id: String, - - /// Show detailed child job breakdown (for batch jobs) - #[arg(long, short = 'v')] - pub verbose: bool, -} - -impl StatusCommand { - /// Execute the train status command - pub async fn run(&self, api_gateway_url: &str, jwt_token: &str) -> AnyhowResult<()> { - println!("\u{1f50d} Fetching training job status..."); - println!(" Job ID: {}", self.job_id.bright_cyan()); - - // Query job status from API Gateway - let job_details = - query_job_status(api_gateway_url, jwt_token, &self.job_id).await?; - - // Display job summary - display_job_summary(&job_details)?; - - // Display metrics (if available) - if let Some(metrics) = &job_details.final_financial_metrics { - display_financial_metrics(metrics); - } - - // Display error details (if failed) - if job_details.status == TrainingStatus::Failed as i32 - && !job_details.error_message.is_empty() - { - display_error_details(&job_details.error_message); - } - - // Display checkpoint path (if completed) - if job_details.status == TrainingStatus::Completed as i32 - && !job_details.model_artifact_path.is_empty() - { - display_checkpoint_info(&job_details.model_artifact_path); - } - - Ok(()) - } -} - -/// Query job status from ML Training Service via API Gateway -pub async fn query_job_status( - api_gateway_url: &str, - jwt_token: &str, - job_id: &str, -) -> AnyhowResult { - // Connect to API Gateway - let mut client = MlTrainingServiceClient::connect(api_gateway_url.to_owned()) - .await - .context("Failed to connect to API Gateway")?; - - // Create request with job ID - let mut request = tonic::Request::new(GetTrainingJobDetailsRequest { - job_id: job_id.to_owned(), - }); - - // Add JWT token to metadata - request.metadata_mut().insert( - "authorization", - format!("Bearer {}", jwt_token).parse()?, - ); - - // Make gRPC call - let response = client - .get_training_job_details(request) - .await - .context("Failed to get training job details")?; - - let job_response = response.into_inner(); - - job_response - .job_details - .ok_or_else(|| anyhow::anyhow!("No job details returned from server")) -} - -/// Display job summary with status, timing, and progress -fn display_job_summary(job: &TrainingJobDetails) -> AnyhowResult<()> { - println!("\n\u{1f4ca} Training Job Status"); - println!("{}", "\u{2500}".repeat(80).bright_black()); - - // Job ID and Model - println!("Job ID: {}", job.job_id.bright_green()); - println!("Model: {}", job.model_type.bright_magenta()); - - // Status with color coding - let status_str = format_training_status(job.status); - let status_colored = format_status_colored(&status_str); - println!("Status: {}", status_colored); - - // Description (if available) - if !job.description.is_empty() { - println!("Description: {}", job.description.bright_white()); - } - - // Timing information - if job.started_at > 0 { - let started = DateTime::from_timestamp(job.started_at, 0) - .ok_or_else(|| anyhow::anyhow!("Invalid started_at timestamp"))?; - println!("Started: {} UTC", started.format("%Y-%m-%d %H:%M:%S")); - - if job.completed_at > 0 { - let completed = DateTime::from_timestamp(job.completed_at, 0) - .ok_or_else(|| anyhow::anyhow!("Invalid completed_at timestamp"))?; - println!("Completed: {} UTC", completed.format("%Y-%m-%d %H:%M:%S")); - - // Calculate duration - let duration = job.completed_at - job.started_at; - println!("Duration: {}", format_duration(duration)); - } else if job.status == TrainingStatus::Running as i32 { - // Calculate elapsed time for running jobs - let now = Utc::now().timestamp(); - let elapsed = now - job.started_at; - println!("Elapsed: {}", format_duration(elapsed)); - } else { - // No timing info available - } - } else if job.status == TrainingStatus::Pending as i32 { - println!("Status: {}", "Waiting in queue...".yellow()); - } else { - // No status display needed - } - - println!("{}", "\u{2500}".repeat(80).bright_black()); - - Ok(()) -} - -/// Display financial performance metrics -fn display_financial_metrics(metrics: &crate::proto::ml_training::FinancialMetrics) { - println!("\n\u{1f3c6} Performance Metrics"); - println!("{}", "\u{2500}".repeat(80).bright_black()); - - let mut table = Table::new(); - table.set_header(vec![ - Cell::new("Metric").fg(Color::Cyan), - Cell::new("Value").fg(Color::Cyan), - ]); - - // Sharpe Ratio - let sharpe_str = format!("{:.4}", metrics.sharpe_ratio); - let sharpe_cell = if metrics.sharpe_ratio >= 2.0 { - Cell::new(sharpe_str).fg(Color::Green) - } else if metrics.sharpe_ratio >= 1.5 { - Cell::new(sharpe_str).fg(Color::Yellow) - } else { - Cell::new(sharpe_str).fg(Color::Red) - }; - table.add_row(vec![Cell::new("Sharpe Ratio"), sharpe_cell]); - - // Hit Rate (Accuracy) - let hit_rate_str = format!("{:.1}%", metrics.hit_rate * 100.0); - let hit_rate_cell = if metrics.hit_rate >= 0.70 { - Cell::new(hit_rate_str).fg(Color::Green) - } else if metrics.hit_rate >= 0.60 { - Cell::new(hit_rate_str).fg(Color::Yellow) - } else { - Cell::new(hit_rate_str).fg(Color::Red) - }; - table.add_row(vec![Cell::new("Hit Rate"), hit_rate_cell]); - - // Simulated Return - let return_str = format!("{:+.2}%", metrics.simulated_return * 100.0); - let return_cell = if metrics.simulated_return > 0.0 { - Cell::new(return_str).fg(Color::Green) - } else { - Cell::new(return_str).fg(Color::Red) - }; - table.add_row(vec![Cell::new("Simulated Return"), return_cell]); - - // Max Drawdown - let drawdown_str = format!("{:.2}%", metrics.max_drawdown * 100.0); - let drawdown_cell = if metrics.max_drawdown < 0.05 { - Cell::new(drawdown_str).fg(Color::Green) - } else if metrics.max_drawdown < 0.10 { - Cell::new(drawdown_str).fg(Color::Yellow) - } else { - Cell::new(drawdown_str).fg(Color::Red) - }; - table.add_row(vec![Cell::new("Max Drawdown"), drawdown_cell]); - - // Risk-Adjusted Return - table.add_row(vec![ - Cell::new("Risk-Adjusted Return"), - Cell::new(format!("{:+.2}%", metrics.risk_adjusted_return * 100.0)), - ]); - - // VaR 5% - table.add_row(vec![ - Cell::new("VaR (5%)"), - Cell::new(format!("{:.2}%", metrics.var_5pct * 100.0)), - ]); - - println!("{}", table); -} - -/// Display error details for failed jobs -fn display_error_details(error_message: &str) { - println!("\n\u{274c} Error Details"); - println!("{}", "\u{2500}".repeat(80).bright_black()); - println!("{}", error_message.red()); - println!("{}", "\u{2500}".repeat(80).bright_black()); -} - -/// Display checkpoint path for completed jobs -fn display_checkpoint_info(checkpoint_path: &str) { - println!("\n\u{1f4be} Model Checkpoint"); - println!("{}", "\u{2500}".repeat(80).bright_black()); - println!("Path: {}", checkpoint_path.bright_green()); - - // Estimate file size (if path exists) - if let Ok(metadata) = std::fs::metadata(checkpoint_path) { - let size_mb = metadata.len() as f64 / (1024.0 * 1024.0); - println!("Size: {:.1} MB", size_mb); - } - - println!("{}", "\u{2500}".repeat(80).bright_black()); -} - -// ============================================================================ -// Helper Functions -// ============================================================================ - -/// Convert `TrainingStatus` enum to string -pub fn format_training_status(status: i32) -> String { - match TrainingStatus::try_from(status).ok() { - Some(TrainingStatus::Unknown) => "UNKNOWN".to_owned(), - Some(TrainingStatus::Pending) => "PENDING".to_owned(), - Some(TrainingStatus::Running) => "RUNNING".to_owned(), - Some(TrainingStatus::Completed) => "COMPLETED".to_owned(), - Some(TrainingStatus::Failed) => "FAILED".to_owned(), - Some(TrainingStatus::Stopped) => "STOPPED".to_owned(), - Some(TrainingStatus::Paused) => "PAUSED".to_owned(), - None => format!("UNKNOWN({})", status), - } -} - -/// Format status with color coding -fn format_status_colored(status: &str) -> colored::ColoredString { - match status { - "RUNNING" => format!("\u{23f3} {}", status).green(), - "COMPLETED" => format!("\u{2705} {}", status).bright_green().bold(), - "FAILED" => format!("\u{274c} {}", status).red().bold(), - "STOPPED" => format!("\u{1f6d1} {}", status).yellow(), - "PENDING" => format!("\u{23f8}\u{fe0f} {}", status).bright_blue(), - "PAUSED" => format!("\u{23f8}\u{fe0f} {}", status).yellow(), - _ => status.white(), - } -} - -/// Format duration in human-readable format -pub fn format_duration(seconds: i64) -> String { - let abs = seconds.unsigned_abs(); - let hours = abs / 3600; - let minutes = (abs % 3600) / 60; - let secs = abs % 60; - - if hours > 0 { - format!("{}h {}m {}s", hours, minutes, secs) - } else if minutes > 0 { - format!("{}m {}s", minutes, secs) - } else { - format!("{}s", secs) - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_format_training_status_enum() { - assert_eq!(format_training_status(TrainingStatus::Pending as i32), "PENDING"); - assert_eq!(format_training_status(TrainingStatus::Running as i32), "RUNNING"); - assert_eq!(format_training_status(TrainingStatus::Completed as i32), "COMPLETED"); - assert_eq!(format_training_status(TrainingStatus::Failed as i32), "FAILED"); - assert_eq!(format_training_status(TrainingStatus::Stopped as i32), "STOPPED"); - assert_eq!(format_training_status(TrainingStatus::Paused as i32), "PAUSED"); - } - - #[test] - fn test_format_duration_seconds() { - assert_eq!(format_duration(30), "30s"); - assert_eq!(format_duration(90), "1m 30s"); - assert_eq!(format_duration(204), "3m 24s"); - assert_eq!(format_duration(3661), "1h 1m 1s"); - } -} diff --git a/bin/fxt/src/commands/tune.rs b/bin/fxt/src/commands/tune.rs index aee2d1a83..4fbadc2db 100644 --- a/bin/fxt/src/commands/tune.rs +++ b/bin/fxt/src/commands/tune.rs @@ -1,996 +1,53 @@ -//! TLI Tune Command - Hyperparameter Tuning Management -//! -//! Provides CLI interface for managing ML model hyperparameter tuning jobs through the API Gateway. -//! Supports starting, monitoring, and stopping Optuna-based hyperparameter optimization runs. -//! -//! # Architecture Flow -//! -//! ```text -//! TLI Client (tune.rs) -//! │ -//! │ gRPC + JWT auth -//! ▼ -//! API Gateway (port 50051) -//! │ -//! │ Proxy with metadata forwarding -//! ▼ -//! ML Training Service (port 50054) -//! │ -//! ├─ Optuna hyperparameter optimization -//! ├─ Trial management & tracking -//! └─ Best parameters persistence -//! ``` -//! -//! # Authentication -//! -//! All commands require JWT authentication obtained through TLI's auth flow: -//! - JWT token passed in gRPC metadata header: `authorization: Bearer ` -//! - API Gateway validates token and forwards to ML Training Service -//! - Token must have valid `jti`, `roles`, and `permissions` fields -//! -//! # Supported Operations -//! -//! ## 1. Start Tuning Job -//! Initiates a new hyperparameter optimization run with Optuna. -//! -//! ```bash -//! # Basic usage -//! tli tune start --model DQN --trials 50 --config tuning_config.yaml -//! -//! # With GPU acceleration -//! tli tune start --model MAMBA_2 --trials 100 --config tuning.yaml --gpu -//! -//! # With custom data source -//! tli tune start --model TFT --trials 75 --data-source /path/to/training_data.parquet -//! -//! # With live progress monitoring (polling every 5s) -//! tli tune start --model DQN --trials 50 --config tuning.yaml --watch -//! ``` -//! -//! ## 2. Check Job Status -//! Query current progress and metrics of a running tuning job. -//! -//! ```bash -//! tli tune status --job-id 550e8400-e29b-41d4-a716-446655440000 -//! ``` -//! -//! Output: -//! - Current status (PENDING, RUNNING, COMPLETED, FAILED, STOPPED) -//! - Progress: trials completed / total trials -//! - Best Sharpe ratio found so far -//! - Elapsed time -//! -//! ## 3. Get Best Parameters -//! Retrieve the best hyperparameters found during optimization. -//! -//! ```bash -//! # Display best parameters -//! tli tune best --job-id -//! -//! # Export to YAML file -//! tli tune best --job-id --export best_params.yaml -//! ``` -//! -//! Returns: -//! - Best hyperparameters (`learning_rate`, `batch_size`, etc.) -//! - Best performance metrics (`sharpe_ratio`, `training_loss`, etc.) -//! -//! ## 4. Stop Tuning Job -//! Gracefully stop a running optimization job. -//! -//! ```bash -//! # Stop with reason -//! tli tune stop --job-id --reason "Sufficient trials completed" -//! -//! # Stop without reason -//! tli tune stop --job-id -//! ``` -//! -//! # Supported Models -//! -//! - **DQN**: Deep Q-Network (reinforcement learning) -//! - **PPO**: Proximal Policy Optimization (reinforcement learning) -//! - **`MAMBA_2`**: State space model with selective attention -//! - **TLOB**: Time-aware Limit Order Book model -//! - **TFT**: Temporal Fusion Transformer (time series) -//! - **LIQUID**: Liquid neural networks (continuous-time RNN) -//! -//! # Configuration File Format -//! -//! Tuning configuration YAML example: -//! -//! ```yaml -//! search_space: -//! learning_rate: -//! type: loguniform -//! low: 1e-5 -//! high: 1e-2 -//! batch_size: -//! type: categorical -//! choices: [64, 128, 256, 512] -//! gamma: -//! type: uniform -//! low: 0.95 -//! high: 0.999 -//! -//! objective: -//! metric: sharpe_ratio -//! direction: maximize -//! -//! pruning: -//! enabled: true -//! strategy: median -//! warmup_trials: 10 -//! ``` -//! -//! # Job Tracking -//! -//! TLI automatically saves job IDs to `~/.foxhunt/tuning_jobs.json` for easy tracking: -//! -//! ```json -//! { -//! "550e8400-e29b-41d4-a716-446655440000": { -//! "job_id": "550e8400-e29b-41d4-a716-446655440000", -//! "model": "DQN", -//! "trials": 50, -//! "started_at": "2025-10-13T14:30:00Z", -//! "status": "RUNNING" -//! } -//! } -//! ``` -//! -//! # Troubleshooting -//! -//! ## Authentication Errors -//! -//! **Error**: `Status { code: Unauthenticated, message: "Missing or invalid JWT token" }` -//! -//! **Solution**: Ensure you're authenticated with TLI: -//! ```bash -//! tli auth login --username --password -//! ``` -//! -//! ## Connection Errors -//! -//! **Error**: `Failed to connect to API Gateway` -//! -//! **Solution**: Verify API Gateway is running: -//! ```bash -//! # Check API Gateway health -//! grpc_health_probe -addr=localhost:50051 -//! -//! # Start API Gateway if not running -//! cargo run -p api_gateway -//! ``` -//! -//! ## ML Training Service Unavailable -//! -//! **Error**: `Service unavailable: ML Training Service` -//! -//! **Solution**: Verify ML Training Service is running: -//! ```bash -//! # Check ML Training Service health -//! curl http://localhost:8095/health -//! -//! # Start ML Training Service if not running -//! cargo run -p ml_training_service -//! ``` -//! -//! ## Invalid Job ID -//! -//! **Error**: `Invalid job ID format (expected UUID)` -//! -//! **Solution**: Ensure job ID is valid UUID format: -//! ```bash -//! # Valid: 550e8400-e29b-41d4-a716-446655440000 -//! # Invalid: abc123, test-job-id -//! ``` -//! -//! Check `~/.foxhunt/tuning_jobs.json` for valid job IDs. -//! -//! ## Config File Not Found -//! -//! **Error**: `Config file not found: tuning_config.yaml` -//! -//! **Solution**: Create tuning config YAML or specify absolute path: -//! ```bash -//! tli tune start --model DQN --trials 50 --config /absolute/path/to/tuning.yaml -//! ``` -//! -//! ## GPU Not Available -//! -//! **Warning**: Using `--gpu` flag when GPU unavailable falls back to CPU. -//! -//! **Verify GPU**: Check CUDA availability: -//! ```bash -//! nvidia-smi # Check GPU status -//! nvcc --version # Check CUDA installation -//! ``` -//! -//! # Performance Characteristics -//! -//! - **Latency**: ~21-488μs (API Gateway proxy overhead) -//! - **Throughput**: Supports multiple concurrent tuning jobs -//! - **Scalability**: Job tracking persists across TLI sessions -//! - **Reliability**: Graceful degradation if ML Training Service unavailable -//! -//! # Implementation Status -//! -//! - ✅ Start tuning job (fully implemented, Wave 132) -//! - ✅ Get job status (fully implemented, Wave 132) -//! - ✅ Get best parameters (fully implemented, Wave 132) -//! - ✅ Stop tuning job (fully implemented, Wave 132) -//! - ⚠️ Real-time streaming (not yet implemented, use polling with `--watch`) -//! -//! # Notes -//! -//! - TLI is a pure client with zero direct service access -//! - All operations route through API Gateway (port 50051) -//! - ML Training Service runs on port 50054 (gRPC) + 8095 (health) -//! - Job IDs are UUIDs generated by ML Training Service -//! - Best parameters exported in YAML format for easy integration +//! `fxt tune` -- hyperparameter optimization. -use anyhow::{Context, Result as AnyhowResult}; -use chrono; -use clap::Subcommand; -use colored::Colorize; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::fmt::Write as _; -use tabled::{Table, Tabled}; -use uuid::Uuid; +use anyhow::Result; +use clap::{Parser, Subcommand}; -// Import ML training proto types -use crate::proto::ml_training::{ - data_source::Source, ml_training_service_client::MlTrainingServiceClient, DataSource, - GetTuningJobStatusRequest, StartTuningJobRequest, StopTuningJobRequest, TrialResult, - TrialState, TuningJobStatus, -}; +use crate::grpc::FoxhuntClient; +use crate::output::OutputFormat; -use super::tune_stream; +#[derive(Parser, Debug)] +pub struct TuneCommand { + #[command(subcommand)] + action: TuneAction, +} -/// Tuning command subcommands -#[derive(Debug, Subcommand)] -pub enum TuneCommand { - /// Start a new hyperparameter tuning job +#[derive(Subcommand, Debug)] +enum TuneAction { + /// Start a hyperparameter optimization run Start { - /// Model type to tune (DQN, PPO, `MAMBA_2`, TLOB, TFT, LIQUID) - #[clap(long, value_name = "MODEL")] + /// Model type + #[arg(long)] model: String, - - /// Number of tuning trials to run - #[clap(long, default_value = "50")] + /// Number of trials + #[arg(long, default_value = "50")] trials: u32, - - /// Path to tuning configuration file (YAML) - #[clap(long, default_value = "tuning_config.yaml", value_name = "PATH")] - config: String, - - /// Training data source (file path or database query) - #[clap(long, value_name = "SOURCE")] - data_source: Option, - - /// Use GPU acceleration - #[clap(long)] + /// Path to search-space config + #[arg(long)] + config: Option, + /// Enable GPU + #[arg(long)] gpu: bool, - - /// Optional job description - #[clap(long, value_name = "DESCRIPTION")] - description: Option, - - /// Watch live progress with polling (5s interval) - #[clap(long)] - watch: bool, }, - - /// Get tuning job status and progress - Status { - /// Tuning job ID (UUID) - #[clap(long, value_name = "UUID")] - job_id: String, - }, - - /// Show best hyperparameters found so far - Best { - /// Tuning job ID (UUID) - #[clap(long, value_name = "UUID")] - job_id: String, - - /// Export best parameters to file - #[clap(long, value_name = "PATH")] - export: Option, - }, - - /// Stop a running tuning job + /// Stop a running optimization Stop { - /// Tuning job ID (UUID) - #[clap(long, value_name = "UUID")] + /// Optimization job ID + job_id: String, + }, + /// Show optimization status and best trial so far + Status { + /// Optimization job ID + job_id: String, + }, + /// Approve best hyperparameters for training + Approve { + /// Optimization job ID job_id: String, - - /// Reason for stopping - #[clap(long, value_name = "REASON")] - reason: Option, }, } -/// Tuning job status representation -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TuningJobStatusDisplay { - pub job_id: String, - pub status: String, - pub current_trial: u32, - pub total_trials: u32, - pub progress_percent: f32, - pub best_sharpe_ratio: f32, - pub elapsed_time_seconds: i64, -} - -/// Trial result display (for history table) -#[derive(Debug, Clone, Tabled)] -struct TrialDisplay { - #[tabled(rename = "Trial")] - trial_number: u32, - #[tabled(rename = "Sharpe Ratio")] - sharpe_ratio: String, - #[tabled(rename = "Training Loss")] - training_loss: String, - #[tabled(rename = "State")] - state: String, - #[tabled(rename = "Duration (s)")] - duration: i64, -} - -/// Best hyperparameters display (for table output) -#[derive(Debug, Clone, Tabled)] -struct BestParamDisplay { - #[tabled(rename = "Parameter")] - name: String, - #[tabled(rename = "Value")] - value: String, - #[tabled(rename = "Type")] - param_type: String, -} - -/// Execute the tune command -/// -/// # Arguments -/// * `command` - Tune subcommand to execute -/// * `api_gateway_url` - API Gateway endpoint (default: ) -/// * `jwt_token` - JWT authentication token from TLI auth flow -/// -/// # Returns -/// `Result<(), anyhow::Error>` - Ok if command executed successfully -/// -/// # Errors -/// - Authentication failures (invalid/expired JWT) -/// - Service unavailable (API Gateway down) -/// - Invalid `job_id` (malformed UUID or job not found) -/// - Network errors (connection timeout) -pub async fn execute_tune_command( - command: TuneCommand, - api_gateway_url: &str, - jwt_token: &str, -) -> AnyhowResult<()> { - match command { - TuneCommand::Start { - model, - trials, - config, - data_source, - gpu, - description, - watch, - } => { - start_tuning_job( - api_gateway_url, - jwt_token, - &model, - trials, - &config, - data_source.as_deref(), - gpu, - description.as_deref(), - watch, - ) - .await - }, - TuneCommand::Status { job_id } => { - get_tuning_status(api_gateway_url, jwt_token, &job_id).await - }, - TuneCommand::Best { job_id, export } => { - get_best_params(api_gateway_url, jwt_token, &job_id, export.as_deref()).await - }, - TuneCommand::Stop { job_id, reason } => { - stop_tuning_job(api_gateway_url, jwt_token, &job_id, reason.as_deref()).await - }, - } -} - -/// Start a new hyperparameter tuning job -async fn start_tuning_job( - api_gateway_url: &str, - jwt_token: &str, - model: &str, - trials: u32, - config_path: &str, - data_source: Option<&str>, - use_gpu: bool, - description: Option<&str>, - watch: bool, -) -> AnyhowResult<()> { - // Validate model type - validate_model_type(model)?; - - // Validate config file exists - if !std::path::Path::new(config_path).exists() { - anyhow::bail!("\u{274c} Config file not found: {}", config_path); - } - - println!("\u{1f680} Starting hyperparameter tuning job..."); - println!(" Model: {}", model.bright_cyan()); - println!(" Trials: {}", trials.to_string().bright_yellow()); - println!(" Config: {}", config_path.bright_white()); - println!( - " GPU: {}", - if use_gpu { - "\u{2705} Enabled".green() - } else { - "\u{274c} Disabled".red() - } - ); - if watch { - println!( - " Watch: {}", - "\u{2705} Enabled (polling every 5s)".green() - ); - } - - // Connect to API Gateway and start tuning job - let mut client = MlTrainingServiceClient::connect(api_gateway_url.to_owned()) - .await - .context("Failed to connect to API Gateway")?; - - let mut request = tonic::Request::new(StartTuningJobRequest { - model_type: model.to_owned(), - num_trials: trials, - config_path: config_path.to_owned(), - data_source: data_source.map(|s| DataSource { - source: Some(Source::FilePath(s.to_owned())), - start_time: 0, - end_time: 0, - }), - use_gpu, - description: description.map(String::from).unwrap_or_default(), - tags: HashMap::new(), - }); - - // Add JWT token to metadata - request - .metadata_mut() - .insert("authorization", format!("Bearer {}", jwt_token).parse()?); - - let response = client - .start_tuning_job(request) - .await - .context("Failed to start tuning job")?; - - let job_id_str = response.into_inner().job_id; - let job_id = Uuid::parse_str(&job_id_str).context("Invalid job ID returned from server")?; - - println!("\n\u{2705} Tuning job started successfully!"); - println!(" Job ID: {}", job_id.to_string().bright_green()); - - // Save job ID to ~/.foxhunt/tuning_jobs.json for later queries - if let Err(e) = save_tuning_job_id(&job_id, model, trials) { - println!( - "\u{26a0}\u{fe0f} Warning: Failed to save job ID to ~/.foxhunt/tuning_jobs.json: {}", - e - ); - println!(" (Job is still running, but manual tracking required)"); - } else { - println!(" Saved to ~/.foxhunt/tuning_jobs.json"); - } - - // If --watch flag is set, stream real-time progress updates - if watch { - tune_stream::watch_tuning_progress_streaming( - api_gateway_url, - jwt_token, - &job_id.to_string(), - ) - .await?; - } else { - println!("\n\u{1f4a1} Monitor progress with:"); - println!(" tli tune status --job-id {}", job_id); - } - - Ok(()) -} - -/// Get tuning job status with rich progress display -async fn get_tuning_status( - api_gateway_url: &str, - jwt_token: &str, - job_id_str: &str, -) -> AnyhowResult<()> { - // Validate job ID format - let job_id = - Uuid::parse_str(job_id_str).context("\u{274c} Invalid job ID format (expected UUID)")?; - - println!("\u{1f50d} Fetching tuning job status..."); - println!(" Job ID: {}", job_id.to_string().bright_cyan()); - - // Connect to API Gateway - let mut client = MlTrainingServiceClient::connect(api_gateway_url.to_owned()) - .await - .context("Failed to connect to API Gateway")?; - - // Create request with job ID - let mut request = tonic::Request::new(GetTuningJobStatusRequest { - job_id: job_id.to_string(), - }); - - // Add JWT token to metadata - request - .metadata_mut() - .insert("authorization", format!("Bearer {}", jwt_token).parse()?); - - // Make gRPC call - let response = client - .get_tuning_job_status(request) - .await - .context("Failed to get tuning job status")?; - - let status_response = response.into_inner(); - - // Convert to display format - let progress_percent = if status_response.total_trials > 0 { - (status_response.current_trial as f32 / status_response.total_trials as f32) * 100.0 - } else { - 0.0 - }; - - let elapsed_time_seconds = if status_response.started_at > 0 { - let now = chrono::Utc::now().timestamp(); - (now - status_response.started_at).max(0) - } else { - 0 - }; - - // Get best Sharpe ratio from best_metrics - let best_sharpe_ratio = status_response - .best_metrics - .get("sharpe_ratio") - .copied() - .unwrap_or(0.0); - - // Convert TuningJobStatus enum to string - let status_str = format_tuning_job_status(status_response.status); - - // Display status with color coding - println!("\n\u{1f4ca} Tuning Job Status"); - println!(" Status: {}", format_status_colored(&status_str)); - println!( - " Progress: {}/{} trials ({:.1}%)", - status_response.current_trial, status_response.total_trials, progress_percent - ); - - // Progress bar visualization - let progress_bar = create_progress_bar(progress_percent); - println!(" {}", progress_bar); - - println!("\n\u{1f3c6} Best Results So Far"); - println!( - " Sharpe Ratio: {}", - format!("{:.4}", best_sharpe_ratio).bright_green() - ); - println!(" Elapsed Time: {} seconds", elapsed_time_seconds); - - // Display best metrics if available - if !status_response.best_metrics.is_empty() { - println!("\n\u{1f4c8} Best Metrics"); - for (metric_name, metric_value) in &status_response.best_metrics { - println!( - " {}: {}", - metric_name.bright_white(), - format!("{:.6}", metric_value).bright_cyan() - ); - } - } - - // Display trial history if available - if !status_response.trial_history.is_empty() { - display_trial_history(&status_response.trial_history); - } - - Ok(()) -} - -/// Get best hyperparameters found so far -async fn get_best_params( - api_gateway_url: &str, - jwt_token: &str, - job_id_str: &str, - export_path: Option<&str>, -) -> AnyhowResult<()> { - // Validate job ID - let job_id = - Uuid::parse_str(job_id_str).context("\u{274c} Invalid job ID format (expected UUID)")?; - - println!("\u{1f50d} Fetching best hyperparameters..."); - println!(" Job ID: {}", job_id.to_string().bright_cyan()); - - // Connect to API Gateway - let mut client = MlTrainingServiceClient::connect(api_gateway_url.to_owned()) - .await - .context("Failed to connect to API Gateway")?; - - // Create request with job ID - let mut request = tonic::Request::new(GetTuningJobStatusRequest { - job_id: job_id.to_string(), - }); - - // Add JWT token to metadata - request - .metadata_mut() - .insert("authorization", format!("Bearer {}", jwt_token).parse()?); - - // Make gRPC call (reuse GetTuningJobStatus which returns best_params and best_metrics) - let response = client - .get_tuning_job_status(request) - .await - .context("Failed to get tuning job status")?; - - let status_response = response.into_inner(); - - // Extract best params and metrics - let best_params = status_response.best_params; - let best_metrics = status_response.best_metrics; - - // Display best metrics - println!("\n\u{1f3c6} Best Performance Metrics"); - for (metric_name, metric_value) in &best_metrics { - println!( - " {}: {}", - metric_name.bright_white(), - format!("{:.4}", metric_value).bright_green() - ); - } - - // Display best hyperparameters as table - println!("\n\u{1f4cb} Best Hyperparameters"); - let param_rows: Vec = best_params - .iter() - .map(|(name, value)| BestParamDisplay { - name: name.clone(), - value: format!("{:.6}", value), - param_type: infer_param_type(*value), - }) - .collect(); - - let table = Table::new(param_rows).to_string(); - println!("{}", table); - - // Export to file if requested - if let Some(path) = export_path { - export_best_params(&best_params, &best_metrics, path)?; - println!( - "\n\u{2705} Best parameters exported to: {}", - path.bright_green() - ); - } - - println!("\n\u{1f4a1} Use these parameters in your training configuration."); - - Ok(()) -} - -/// Stop a running tuning job -async fn stop_tuning_job( - api_gateway_url: &str, - jwt_token: &str, - job_id_str: &str, - reason: Option<&str>, -) -> AnyhowResult<()> { - // Validate job ID - let job_id = - Uuid::parse_str(job_id_str).context("\u{274c} Invalid job ID format (expected UUID)")?; - - println!("\u{1f6d1} Stopping tuning job..."); - println!(" Job ID: {}", job_id.to_string().bright_cyan()); - if let Some(reason_text) = reason { - println!(" Reason: {}", reason_text.bright_yellow()); - } - - // Connect to API Gateway - let mut client = MlTrainingServiceClient::connect(api_gateway_url.to_owned()) - .await - .context("Failed to connect to API Gateway")?; - - // Create stop request - let mut request = tonic::Request::new(StopTuningJobRequest { - job_id: job_id.to_string(), - reason: reason.map(String::from).unwrap_or_default(), - }); - - // Add JWT token to metadata - request - .metadata_mut() - .insert("authorization", format!("Bearer {}", jwt_token).parse()?); - - // Call stop_tuning_job method - let response = client - .stop_tuning_job(request) - .await - .context("Failed to stop tuning job")?; - - let stop_response = response.into_inner(); - - // Convert final_status enum to string - let final_status_str = format_tuning_job_status(stop_response.final_status); - - println!("\n\u{2705} Tuning job stopped successfully!"); - println!(" Message: {}", stop_response.message.bright_green()); - println!( - " Final Status: {}", - format_status_colored(&final_status_str) - ); - - println!("\n\u{1f4a1} Get final results with:"); - println!(" tli tune best --job-id {}", job_id); - - Ok(()) -} - -// ============================================================================ -// Helper Functions -// ============================================================================ - -/// Save tuning job ID to ~/.`foxhunt/tuning_jobs.json` for later tracking -fn save_tuning_job_id(job_id: &Uuid, model: &str, trials: u32) -> AnyhowResult<()> { - use std::fs; - use std::io::Write; - use std::path::PathBuf; - - // Get home directory - let home_dir = std::env::var("HOME") - .or_else(|_| std::env::var("USERPROFILE")) - .context("Failed to determine home directory")?; - - // Create ~/.foxhunt directory if it doesn't exist - let foxhunt_dir = PathBuf::from(home_dir).join(".foxhunt"); - fs::create_dir_all(&foxhunt_dir).context("Failed to create ~/.foxhunt directory")?; - - let jobs_file = foxhunt_dir.join("tuning_jobs.json"); - - // Load existing jobs or create new structure - let mut jobs: HashMap = if jobs_file.exists() { - let contents = fs::read_to_string(&jobs_file).context("Failed to read tuning_jobs.json")?; - serde_json::from_str(&contents).unwrap_or_default() - } else { - HashMap::new() - }; - - // Add new job entry - let job_entry = serde_json::json!({ - "job_id": job_id.to_string(), - "model": model, - "trials": trials, - "started_at": chrono::Utc::now().to_rfc3339(), - "status": "RUNNING" - }); - - jobs.insert(job_id.to_string(), job_entry); - - // Write back to file with pretty formatting - let json_str = - serde_json::to_string_pretty(&jobs).context("Failed to serialize jobs to JSON")?; - - let mut file = fs::File::create(&jobs_file).context("Failed to create tuning_jobs.json")?; - file.write_all(json_str.as_bytes()) - .context("Failed to write to tuning_jobs.json")?; - - Ok(()) -} - -/// Display trial history as a table -fn display_trial_history(trial_history: &[TrialResult]) { - println!("\n\u{1f4ca} Trial History"); - - let trial_rows: Vec = trial_history - .iter() - .map(|trial| { - let sharpe = trial - .metrics - .get("sharpe_ratio") - .or(Some(&trial.objective_value)) - .map(|v| format!("{:.4}", v)) - .unwrap_or_else(|| "N/A".to_owned()); - - let loss = trial - .metrics - .get("training_loss") - .map(|v| format!("{:.6}", v)) - .unwrap_or_else(|| "N/A".to_owned()); - - let duration = if trial.completed_at > trial.started_at { - trial.completed_at - trial.started_at - } else { - 0 - }; - - TrialDisplay { - trial_number: trial.trial_number, - sharpe_ratio: sharpe, - training_loss: loss, - state: format_trial_state(trial.state), - duration, - } - }) - .collect(); - - let table = Table::new(trial_rows).to_string(); - println!("{}", table); -} - -/// Convert `TuningJobStatus` enum to string -fn format_tuning_job_status(status: i32) -> String { - match TuningJobStatus::try_from(status).ok() { - Some(TuningJobStatus::TuningUnknown) => "TUNING_UNKNOWN".to_owned(), - Some(TuningJobStatus::TuningPending) => "TUNING_PENDING".to_owned(), - Some(TuningJobStatus::TuningRunning) => "TUNING_RUNNING".to_owned(), - Some(TuningJobStatus::TuningCompleted) => "TUNING_COMPLETED".to_owned(), - Some(TuningJobStatus::TuningFailed) => "TUNING_FAILED".to_owned(), - Some(TuningJobStatus::TuningStopped) => "TUNING_STOPPED".to_owned(), - None => format!("UNKNOWN({})", status), - } -} - -/// Convert `TrialState` enum to string -fn format_trial_state(state: i32) -> String { - match TrialState::try_from(state).ok() { - Some(TrialState::TrialUnknown) => "UNKNOWN".to_owned(), - Some(TrialState::TrialRunning) => "RUNNING".to_owned(), - Some(TrialState::TrialComplete) => "COMPLETE".to_owned(), - Some(TrialState::TrialPruned) => "PRUNED".to_owned(), - Some(TrialState::TrialFailed) => "FAILED".to_owned(), - None => format!("UNKNOWN({})", state), - } -} - -/// Validate model type is supported -fn validate_model_type(model: &str) -> AnyhowResult<()> { - const VALID_MODELS: &[&str] = &["DQN", "PPO", "MAMBA_2", "TLOB", "TFT", "LIQUID"]; - - if !VALID_MODELS.contains(&model) { - anyhow::bail!( - "\u{274c} Invalid model type: {}. Valid options: {}", - model, - VALID_MODELS.join(", ") - ); - } - - Ok(()) -} - -/// Format status with color coding -fn format_status_colored(status: &str) -> colored::ColoredString { - match status { - "RUNNING" | "TUNING_RUNNING" => status.green(), - "COMPLETED" | "TUNING_COMPLETED" => status.bright_green().bold(), - "FAILED" | "TUNING_FAILED" => status.red().bold(), - "STOPPED" | "TUNING_STOPPED" => status.yellow(), - "PENDING" | "TUNING_PENDING" => status.bright_blue(), - _ => status.white(), - } -} - -/// Create ASCII progress bar -fn create_progress_bar(progress_percent: f32) -> String { - let bar_width = 50; - let filled = ((progress_percent / 100.0) * bar_width as f32) as usize; - let empty = bar_width - filled; - - let filled_str = "\u{2588}".repeat(filled).green(); - let empty_str = "\u{2591}".repeat(empty).white(); - - format!("[{}{}] {:.1}%", filled_str, empty_str, progress_percent) -} - -/// Infer parameter type from value -fn infer_param_type(value: f32) -> String { - if value == value.floor() && (1.0..=10000.0).contains(&value) { - "Integer".to_owned() - } else if value > 0.0 && value < 1.0 { - "Learning Rate".to_owned() - } else { - "Float".to_owned() - } -} - -/// Export best parameters to YAML file -fn export_best_params( - params: &HashMap, - metrics: &HashMap, - export_path: &str, -) -> AnyhowResult<()> { - use std::io::Write; - - let mut content = String::from("# Best Hyperparameters from Tuning Job\n\n"); - content.push_str("hyperparameters:\n"); - - for (name, value) in params { - _ = writeln!(content, " {}: {}", name, value); - } - - content.push_str("\nmetrics:\n"); - for (name, value) in metrics { - _ = writeln!(content, " {}: {:.6}", name, value); - } - - let mut file = std::fs::File::create(export_path).context("Failed to create export file")?; - file.write_all(content.as_bytes()) - .context("Failed to write to export file")?; - - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_validate_model_type_valid() { - validate_model_type("DQN").unwrap(); - validate_model_type("PPO").unwrap(); - validate_model_type("MAMBA_2").unwrap(); - } - - #[test] - fn test_validate_model_type_invalid() { - assert!(validate_model_type("INVALID").is_err()); - assert!(validate_model_type("dqn").is_err()); // Case sensitive - } - - #[test] - fn test_uuid_validation() { - let valid_uuid = "550e8400-e29b-41d4-a716-446655440000"; - Uuid::parse_str(valid_uuid).unwrap(); - - let invalid_uuid = "not-a-uuid"; - Uuid::parse_str(invalid_uuid).unwrap_err(); - } - - #[test] - fn test_progress_bar_generation() { - let bar_0 = create_progress_bar(0.0); - assert!(bar_0.contains("0.0%")); - - let bar_50 = create_progress_bar(50.0); - assert!(bar_50.contains("50.0%")); - - let bar_100 = create_progress_bar(100.0); - assert!(bar_100.contains("100.0%")); - } - - #[test] - fn test_param_type_inference() { - assert_eq!(infer_param_type(128.0), "Integer"); - assert_eq!(infer_param_type(0.001), "Learning Rate"); - assert_eq!(infer_param_type(0.99), "Learning Rate"); - assert_eq!(infer_param_type(1.5), "Float"); - } - - #[test] - fn test_validate_model_type_all_valid() { - let valid_models = ["DQN", "PPO", "MAMBA_2", "TLOB", "TFT", "LIQUID"]; - - for model in &valid_models { - assert!( - validate_model_type(model).is_ok(), - "Model {} should be valid", - model - ); - } +impl TuneCommand { + pub async fn execute(&self, _client: &FoxhuntClient, _format: OutputFormat) -> Result<()> { + anyhow::bail!("tune command not yet implemented") } } diff --git a/bin/fxt/src/commands/tune_impl.rs b/bin/fxt/src/commands/tune_impl.rs deleted file mode 100644 index a58a870ed..000000000 --- a/bin/fxt/src/commands/tune_impl.rs +++ /dev/null @@ -1,377 +0,0 @@ -//! Implementation for tune status and tune best commands with real gRPC integration - -use anyhow::{Context, Result as AnyhowResult}; -use chrono; -use colored::Colorize; -use std::collections::HashMap; -use tabled::{Table, Tabled}; -use uuid::Uuid; - -use crate::proto::ml_training::{ - ml_training_service_client::MlTrainingServiceClient, - GetTuningJobStatusRequest, - TuningJobStatus, TrialState, -}; - -/// Best hyperparameters display (for table output) -#[derive(Debug, Clone, Tabled)] -struct BestParamDisplay { - #[tabled(rename = "Parameter")] - name: String, - #[tabled(rename = "Value")] - value: String, - #[tabled(rename = "Type")] - param_type: String, -} - -/// Trial result display (for history table) -#[derive(Debug, Clone, Tabled)] -struct TrialDisplay { - #[tabled(rename = "Trial")] - trial_number: u32, - #[tabled(rename = "Sharpe Ratio")] - sharpe_ratio: String, - #[tabled(rename = "State")] - state: String, - #[tabled(rename = "Duration (s)")] - duration: i64, -} - -/// Get tuning job status with rich progress display -pub async fn get_tuning_status( - api_gateway_url: &str, - jwt_token: &str, - job_id_str: &str, -) -> AnyhowResult<()> { - // Validate job ID format - let job_id = Uuid::parse_str(job_id_str) - .context("❌ Invalid job ID format (expected UUID)")?; - - println!("🔍 Fetching tuning job status..."); - println!(" Job ID: {}", job_id.to_string().bright_cyan()); - - // Create gRPC client - let mut client = MlTrainingServiceClient::connect(api_gateway_url.to_string()) - .await - .context("Failed to connect to API Gateway")?; - - // Create request with JWT metadata - let mut request = tonic::Request::new(GetTuningJobStatusRequest { - job_id: job_id.to_string(), - }); - - request.metadata_mut().insert( - "authorization", - format!("Bearer {}", jwt_token) - .parse() - .context("Failed to parse JWT token")? - ); - - // Execute gRPC call - let response = client - .get_tuning_job_status(request) - .await - .context("Failed to get tuning job status")?; - - let status_response = response.into_inner(); - - // Calculate progress percentage and elapsed time - let progress_percent = if status_response.total_trials > 0 { - (status_response.current_trial as f32 / status_response.total_trials as f32) * 100.0 - } else { - 0.0 - }; - - let elapsed_seconds = if status_response.started_at > 0 { - chrono::Utc::now().timestamp() - status_response.started_at - } else { - 0 - }; - - let status_str = format_tuning_status(status_response.status()); - - // Display status with color coding - println!("\n📊 Tuning Job Status"); - println!(" Job ID: {}", job_id.to_string().bright_cyan()); - println!(" Status: {}", format_status_colored(&status_str)); - println!(" Progress: {}/{} trials ({:.1}%)", - status_response.current_trial, - status_response.total_trials, - progress_percent - ); - - // Progress bar visualization - let progress_bar = create_progress_bar(progress_percent); - println!(" {}", progress_bar); - - // Display best results if available - if !status_response.best_metrics.is_empty() { - println!("\n🏆 Best Results So Far"); - - // Extract Sharpe ratio from best_metrics - if let Some(sharpe_ratio) = status_response.best_metrics.get("sharpe_ratio") { - println!(" Best Sharpe Ratio: {}", format!("{:.4}", sharpe_ratio).bright_green()); - } - - // Find which trial achieved the best result - if let Some(best_trial) = status_response.trial_history.iter() - .filter(|t| t.state() == TrialState::TrialComplete) - .max_by(|a, b| a.objective_value.partial_cmp(&b.objective_value).unwrap_or(std::cmp::Ordering::Equal)) - { - println!(" Best Trial: #{}", best_trial.trial_number.to_string().bright_yellow()); - } - } - - // Display elapsed time - let elapsed_minutes = elapsed_seconds / 60; - let elapsed_seconds_remainder = elapsed_seconds % 60; - - // Estimate remaining time if in progress - if status_response.status() == TuningJobStatus::TuningRunning && status_response.current_trial > 0 { - let avg_trial_time = elapsed_seconds as f64 / status_response.current_trial as f64; - let remaining_trials = status_response.total_trials.saturating_sub(status_response.current_trial); - let est_remaining_seconds = (avg_trial_time * remaining_trials as f64) as i64; - let est_remaining_minutes = est_remaining_seconds / 60; - - println!("\n⏱️ Time Information"); - println!(" Elapsed Time: {}m {}s", elapsed_minutes, elapsed_seconds_remainder); - println!(" Estimated Time Remaining: {} minutes", est_remaining_minutes); - } else { - println!("\n⏱️ Elapsed Time: {}m {}s", elapsed_minutes, elapsed_seconds_remainder); - } - - // Display trial history if available - if !status_response.trial_history.is_empty() { - display_trial_history(&status_response.trial_history)?; - } - - // Display message if available - if !status_response.message.is_empty() { - println!("\n💬 Message: {}", status_response.message); - } - - Ok(()) -} - -/// Get best hyperparameters found so far -pub async fn get_best_params( - api_gateway_url: &str, - jwt_token: &str, - job_id_str: &str, - export_path: Option<&str>, -) -> AnyhowResult<()> { - // Validate job ID - let job_id = Uuid::parse_str(job_id_str) - .context("❌ Invalid job ID format (expected UUID)")?; - - println!("🔍 Fetching best hyperparameters..."); - println!(" Job ID: {}", job_id.to_string().bright_cyan()); - - // Create gRPC client - let mut client = MlTrainingServiceClient::connect(api_gateway_url.to_string()) - .await - .context("Failed to connect to API Gateway")?; - - // Create request with JWT metadata - let mut request = tonic::Request::new(GetTuningJobStatusRequest { - job_id: job_id.to_string(), - }); - - request.metadata_mut().insert( - "authorization", - format!("Bearer {}", jwt_token) - .parse() - .context("Failed to parse JWT token")? - ); - - // Execute gRPC call - let response = client - .get_tuning_job_status(request) - .await - .context("Failed to get tuning job status")?; - - let status_response = response.into_inner(); - - // Display best metrics - if !status_response.best_metrics.is_empty() { - println!("\n🏆 Best Performance Metrics"); - - // Extract and display Sharpe ratio prominently - if let Some(sharpe_ratio) = status_response.best_metrics.get("sharpe_ratio") { - println!(" Sharpe Ratio: {}", format!("{:.4}", sharpe_ratio).bright_green().bold()); - } - - // Display other metrics - for (metric_name, metric_value) in &status_response.best_metrics { - if metric_name != "sharpe_ratio" { - println!(" {}: {}", - metric_name.bright_white(), - format!("{:.6}", metric_value).bright_green() - ); - } - } - } - - // Display best hyperparameters as table - if !status_response.best_params.is_empty() { - println!("\n📋 Best Hyperparameters"); - let param_rows: Vec = status_response.best_params - .iter() - .map(|(name, value)| BestParamDisplay { - name: name.clone(), - value: format!("{:.6}", value), - param_type: infer_param_type(*value), - }) - .collect(); - - let table = Table::new(param_rows).to_string(); - println!("{}", table); - - // Display checkpoint path if available - println!("\n💾 Model Checkpoint"); - println!(" (Checkpoint location will be provided after training completes)"); - } else { - println!("\n⚠️ No best parameters available yet"); - println!(" Job may still be initializing or no trials have completed successfully"); - } - - // Export to file if requested - if let Some(export_path) = export_path { - if !status_response.best_params.is_empty() { - export_best_params( - &status_response.best_params, - &status_response.best_metrics, - export_path - )?; - println!("\n✅ Best parameters exported to: {}", export_path.bright_green()); - } else { - println!("\n⚠️ Cannot export: no parameters available yet"); - } - } - - println!("\n💡 Use these parameters in your training configuration."); - - Ok(()) -} - -// ============================================================================ -// Helper Functions -// ============================================================================ - -/// Display trial history table -fn display_trial_history(trial_history: &[crate::proto::ml_training::TrialResult]) -> AnyhowResult<()> { - println!("\n📊 Trial History (last 10 trials)"); - - // Take last 10 trials - let recent_trials: Vec = trial_history - .iter() - .rev() - .take(10) - .rev() - .map(|trial| { - let duration_seconds = if trial.completed_at > trial.started_at { - trial.completed_at - trial.started_at - } else { - 0 - }; - - TrialDisplay { - trial_number: trial.trial_number, - sharpe_ratio: format!("{:.4}", trial.objective_value), - state: format_trial_state(trial.state()), - duration: duration_seconds, - } - }) - .collect(); - - let table = Table::new(recent_trials).to_string(); - println!("{}", table); - - Ok(()) -} - -/// Format tuning job status as string -fn format_tuning_status(status: TuningJobStatus) -> String { - match status { - TuningJobStatus::TuningUnknown => "UNKNOWN", - TuningJobStatus::TuningPending => "PENDING", - TuningJobStatus::TuningRunning => "RUNNING", - TuningJobStatus::TuningCompleted => "COMPLETED", - TuningJobStatus::TuningFailed => "FAILED", - TuningJobStatus::TuningStopped => "STOPPED", - }.to_string() -} - -/// Format trial state as string -fn format_trial_state(state: TrialState) -> String { - match state { - TrialState::TrialUnknown => "UNKNOWN", - TrialState::TrialRunning => "RUNNING", - TrialState::TrialComplete => "COMPLETE", - TrialState::TrialPruned => "PRUNED", - TrialState::TrialFailed => "FAILED", - }.to_string() -} - -/// Format status with color coding -fn format_status_colored(status: &str) -> colored::ColoredString { - match status { - "RUNNING" | "TUNING_RUNNING" => status.green(), - "COMPLETED" | "TUNING_COMPLETED" => status.bright_green().bold(), - "FAILED" | "TUNING_FAILED" => status.red().bold(), - "STOPPED" | "TUNING_STOPPED" => status.yellow(), - "PENDING" | "TUNING_PENDING" => status.bright_blue(), - _ => status.white(), - } -} - -/// Create ASCII progress bar -fn create_progress_bar(progress_percent: f32) -> String { - let bar_width = 50; - let filled = ((progress_percent / 100.0) * bar_width as f32) as usize; - let empty = bar_width.saturating_sub(filled); - - let filled_str = "█".repeat(filled).green(); - let empty_str = "░".repeat(empty).white(); - - format!("[{}{}] {:.1}%", filled_str, empty_str, progress_percent) -} - -/// Infer parameter type from value -fn infer_param_type(value: f32) -> String { - if value == value.floor() && value >= 1.0 && value <= 10000.0 { - "Integer".to_string() - } else if value > 0.0 && value < 1.0 { - "Learning Rate".to_string() - } else { - "Float".to_string() - } -} - -/// Export best parameters to YAML file -fn export_best_params( - params: &HashMap, - metrics: &HashMap, - export_path: &str, -) -> AnyhowResult<()> { - use std::io::Write; - - let mut content = String::from("# Best Hyperparameters from Tuning Job\n\n"); - content.push_str("hyperparameters:\n"); - - for (name, value) in params { - content.push_str(&format!(" {}: {}\n", name, value)); - } - - content.push_str("\nmetrics:\n"); - for (name, value) in metrics { - content.push_str(&format!(" {}: {:.6}\n", name, value)); - } - - let mut file = std::fs::File::create(export_path) - .context("Failed to create export file")?; - file.write_all(content.as_bytes()) - .context("Failed to write to export file")?; - - Ok(()) -} diff --git a/bin/fxt/src/commands/tune_stream.rs b/bin/fxt/src/commands/tune_stream.rs deleted file mode 100644 index 45148c48a..000000000 --- a/bin/fxt/src/commands/tune_stream.rs +++ /dev/null @@ -1,257 +0,0 @@ -//! TLI Streaming Client for Hyperparameter Tuning Progress -//! -//! This module implements real-time streaming progress updates for tuning jobs. -//! Uses gRPC streaming to receive trial completion notifications from ML Training Service. - -use anyhow::{Context, Result as AnyhowResult}; -use colored::Colorize; -use tokio_stream::StreamExt; -use tracing::{debug, error, info}; -use uuid::Uuid; - -use crate::proto::ml_training::{ - ml_training_service_client::MlTrainingServiceClient, - StreamProgressRequest, UpdateType, -}; - -/// Watch tuning progress with real-time streaming updates -pub async fn watch_tuning_progress_streaming( - api_gateway_url: &str, - jwt_token: &str, - job_id: &str, -) -> AnyhowResult<()> { - // Validate job ID format - let job_id_uuid = Uuid::parse_str(job_id) - .context("\u{274c} Invalid job ID format (expected UUID)")?; - - info!("Subscribing to tuning progress stream for job: {}", job_id_uuid); - - // Create gRPC client with reconnect capability - let mut client = MlTrainingServiceClient::connect(api_gateway_url.to_owned()) - .await - .context("Failed to connect to API Gateway")?; - - // Create streaming request with JWT metadata - let mut request = tonic::Request::new(StreamProgressRequest { - job_id: job_id.to_owned(), - }); - - request.metadata_mut().insert( - "authorization", - format!("Bearer {}", jwt_token) - .parse() - .context("Failed to parse JWT token")? - ); - - // Subscribe to progress stream - let response = client - .stream_tuning_progress(request) - .await - .context("Failed to subscribe to progress stream")?; - - let mut stream = response.into_inner(); - - println!("\n\u{1f440} Watching tuning progress (press Ctrl+C to stop)...\n"); - - // Track progress state - let mut last_displayed_trial = 0_u32; - let mut iteration = 0_u32; - - // Process stream updates - while let Some(result) = stream.next().await { - match result { - Ok(update) => { - iteration += 1; - let update_type = UpdateType::try_from(update.update_type) - .unwrap_or(UpdateType::UpdateUnknown); - - // Skip heartbeats unless we want to show keepalive - if update_type == UpdateType::UpdateHeartbeat { - debug!("Received heartbeat for job {}", job_id); - continue; - } - - // Clear previous display if not first iteration - if iteration > 1 && update.current_trial != last_displayed_trial { - // Clear previous display (9 lines) - print!("\x1B[9A\x1B[J"); - } - - // Display rich progress UI - display_progress_ui(&update); - - last_displayed_trial = update.current_trial; - - // Check if job is complete - if update_type == UpdateType::UpdateJobComplete { - let status_str = format_status_from_i32(update.status); - match status_str.as_str() { - "TUNING_COMPLETED" => { - println!("\n\u{2705} Tuning job completed successfully!"); - println!(" Best Sharpe Ratio: {}", format!("{:.4}", update.best_sharpe_so_far).bright_green()); - println!("\n\u{1f4a1} Get best parameters with:"); - println!(" tli tune best --job-id {}", job_id); - } - "TUNING_FAILED" => { - println!("\n\u{274c} Tuning job failed!"); - println!(" Message: {}", update.message); - } - "TUNING_STOPPED" => { - println!("\n\u{1f6d1} Tuning job stopped by user"); - println!("\n\u{1f4a1} Get partial results with:"); - println!(" tli tune best --job-id {}", job_id); - } - _ => { - println!("\n\u{26a0}\u{fe0f} Tuning job ended with status: {}", status_str); - } - } - break; - } - } - Err(e) => { - error!("Stream error: {}", e); - println!("\n\u{26a0}\u{fe0f} Stream error: {}", e); - println!(" Connection lost, attempting to reconnect..."); - - // Try to reconnect with exponential backoff - tokio::time::sleep(tokio::time::Duration::from_secs(5)).await; - - // Attempt reconnect (caller should handle retry logic) - return Err(anyhow::anyhow!("Stream disconnected: {}", e)); - } - } - } - - println!("\n\u{1f4ca} Progress stream ended"); - - Ok(()) -} - -/// Display rich progress UI for tuning job -fn display_progress_ui(update: &crate::proto::ml_training::ProgressUpdate) { - let progress_percent = if update.total_trials > 0 { - (update.current_trial as f32 / update.total_trials as f32) * 100.0 - } else { - 0.0 - }; - - let status_str = format_status_from_i32(update.status); - - println!("\u{250c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2510}"); - println!("\u{2502} \u{1f3af} Tuning Job: {} \u{2502}", - update.job_id.chars().take(8).collect::() - ); - println!("\u{251c}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2524}"); - println!("\u{2502} Progress: {}/{} ({:.1}%) \u{2502}", - update.current_trial, - update.total_trials, - progress_percent - ); - - // Progress bar - let progress_bar = create_progress_bar(progress_percent); - println!("\u{2502} {} \u{2502}", progress_bar); - - println!("\u{2502} \u{1f3c6} Best Sharpe Ratio: {} \u{2502}", - format!("{:.4}", update.best_sharpe_so_far).bright_green() - ); - - println!("\u{2502} \u{1f4c8} Trial Sharpe: {} \u{2502}", - format!("{:.4}", update.trial_sharpe).bright_cyan() - ); - - // Time remaining - if update.estimated_time_remaining > 0 { - let remaining_minutes = update.estimated_time_remaining / 60; - let remaining_seconds = update.estimated_time_remaining % 60; - println!("\u{2502} \u{23f1}\u{fe0f} Estimated Time: {}m {}s remaining \u{2502}", - remaining_minutes, remaining_seconds - ); - } else { - println!("\u{2502} \u{23f1}\u{fe0f} Estimated Time: calculating... \u{2502}"); - } - - println!("\u{2502} \u{1f4ca} Status: {} \u{2502}", - format_status_colored(&status_str) - ); - - // Trial parameters (show first 3) - if !update.trial_params.is_empty() { - let param_count = update.trial_params.len().min(3); - let params_display: Vec = update.trial_params - .iter() - .take(param_count) - .map(|(k, v)| format!("{}={}", k, v)) - .collect(); - let params_str = params_display.join(", "); - println!("\u{2502} \u{1f527} Params: {} \u{2502}", params_str.chars().take(43).collect::()); - } - - println!("\u{2514}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2500}\u{2518}"); -} - -/// Create ASCII progress bar -fn create_progress_bar(progress_percent: f32) -> String { - let bar_width: usize = 50; - let filled = ((progress_percent / 100.0) * bar_width as f32) as usize; - let empty = bar_width.saturating_sub(filled); - - let filled_str = "\u{2588}".repeat(filled).green(); - let empty_str = "\u{2591}".repeat(empty).white(); - - format!("[{}{}] {:.1}%", filled_str, empty_str, progress_percent) -} - -/// Format status with color coding -fn format_status_colored(status: &str) -> colored::ColoredString { - match status { - "TUNING_RUNNING" => status.green(), - "TUNING_COMPLETED" => status.bright_green().bold(), - "TUNING_FAILED" => status.red().bold(), - "TUNING_STOPPED" => status.yellow(), - "TUNING_PENDING" => status.bright_blue(), - _ => status.white(), - } -} - -/// Convert i32 status to string -fn format_status_from_i32(status: i32) -> String { - match status { - 0 => "TUNING_UNKNOWN", - 1 => "TUNING_PENDING", - 2 => "TUNING_RUNNING", - 3 => "TUNING_COMPLETED", - 4 => "TUNING_FAILED", - 5 => "TUNING_STOPPED", - _ => "TUNING_UNKNOWN", - }.to_owned() -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_progress_bar_generation() { - let bar_0 = create_progress_bar(0.0); - assert!(bar_0.contains("0.0%")); - - let bar_50 = create_progress_bar(50.0); - assert!(bar_50.contains("50.0%")); - - let bar_100 = create_progress_bar(100.0); - assert!(bar_100.contains("100.0%")); - } - - #[test] - fn test_status_formatting() { - let status = format_status_from_i32(2); - assert_eq!(status, "TUNING_RUNNING"); - - let status = format_status_from_i32(3); - assert_eq!(status, "TUNING_COMPLETED"); - - let status = format_status_from_i32(4); - assert_eq!(status, "TUNING_FAILED"); - } -} diff --git a/bin/fxt/src/commands/watch.rs b/bin/fxt/src/commands/watch.rs new file mode 100644 index 000000000..ca0663df0 --- /dev/null +++ b/bin/fxt/src/commands/watch.rs @@ -0,0 +1,17 @@ +//! `fxt watch` -- TUI cockpit dashboard. + +use anyhow::Result; +use clap::Parser; + +#[derive(Parser, Debug)] +pub struct WatchCommand; + +impl WatchCommand { + /// Launch the full-screen TUI dashboard. + /// + /// This command does not need a gRPC client at construction time -- + /// the TUI event loop manages its own connections. + pub async fn execute(&self) -> Result<()> { + anyhow::bail!("watch (TUI) not yet implemented") + } +} diff --git a/bin/fxt/src/commands/watch/event_loop.rs b/bin/fxt/src/commands/watch/event_loop.rs deleted file mode 100644 index bd87e7717..000000000 --- a/bin/fxt/src/commands/watch/event_loop.rs +++ /dev/null @@ -1,852 +0,0 @@ -//! Event loop for the `fxt watch` streaming TUI dashboard. -//! -//! [`run_dashboard`] is the public entry point that sets up the terminal, -//! runs the select-loop, and guarantees clean-up on exit or error. - -use anyhow::{Context, Result}; -use crossterm::{ - event::{Event, EventStream, KeyCode, KeyEvent, KeyModifiers}, - execute, - terminal::{disable_raw_mode, enable_raw_mode, EnterAlternateScreen, LeaveAlternateScreen}, -}; -use ratatui::prelude::*; -use std::io; -use tokio::time::{interval, Duration}; -use tokio_stream::StreamExt; - -use super::render; -use super::state::{ - CircuitBreakerRow, DashboardState, FillRow, ServiceRow, Tab, TrainingViewMode, -}; -use super::streams::{self, StreamEvent}; - -// --------------------------------------------------------------------------- -// Public entry point -// --------------------------------------------------------------------------- - -/// Initialise the terminal, run the event loop, and restore the terminal on -/// exit (even if an error occurred). -pub(super) async fn run_dashboard(api_gateway_url: &str, jwt_token: &str) -> Result<()> { - enable_raw_mode().context("failed to enable raw mode")?; - - let mut stdout = io::stdout(); - execute!(stdout, EnterAlternateScreen).context("failed to enter alternate screen")?; - - let backend = CrosstermBackend::new(stdout); - let mut terminal = - Terminal::new(backend).context("failed to create ratatui terminal")?; - - let result = run_event_loop(&mut terminal, api_gateway_url, jwt_token).await; - - // Always restore the terminal, regardless of whether the loop succeeded. - disable_raw_mode().ok(); - execute!(terminal.backend_mut(), LeaveAlternateScreen).ok(); - terminal.show_cursor().ok(); - - result -} - -// --------------------------------------------------------------------------- -// Core event loop -// --------------------------------------------------------------------------- - -async fn run_event_loop( - terminal: &mut Terminal>, - api_gateway_url: &str, - jwt_token: &str, -) -> Result<()> { - let mut state = DashboardState { - dirty: true, - ..DashboardState::default() - }; - - let mut stream_rx = streams::spawn_all_streams(api_gateway_url, jwt_token); - let mut tick = interval(Duration::from_millis(200)); - let mut event_stream = EventStream::new(); - - loop { - tokio::select! { - maybe_event = stream_rx.recv() => { - if let Some(event) = maybe_event { - apply_stream_event(&mut state, event); - } - // When None, all stream producers dropped — keep running so - // the user can still view stale data and quit manually. - } - - maybe_crossterm = event_stream.next() => { - if let Some(Ok(Event::Key(key))) = maybe_crossterm { - if handle_key(&mut state, key) { - return Ok(()); - } - } - } - - _ = tick.tick() => { - if state.dirty { - terminal - .draw(|f| render::render(f, &state)) - .context("failed to draw frame")?; - state.dirty = false; - } - } - } - } -} - -// --------------------------------------------------------------------------- -// Stream event application -// --------------------------------------------------------------------------- - -fn apply_stream_event(state: &mut DashboardState, event: StreamEvent) { - match event { - StreamEvent::TrainingUpdate { sessions, gpu } => { - // Filter out uninitialised or aggregate rows: - // - empty model name (default proto) - // - model present but empty fold with zero epoch (summary row) - let active_sessions: Vec<_> = sessions - .into_iter() - .filter(|s| !(s.model.is_empty() || s.fold.is_empty() && s.epoch == 0.0)) - .collect(); - // Push first session's epoch_loss to global loss sparkline. - if let Some(first) = active_sessions.first() { - state.training.push_loss(f64::from(first.epoch_loss)); - } - // Accumulate per-session metric history for detail sparklines. - for session in &active_sessions { - let key = format!("{}:{}", session.model, session.fold); - state - .training - .session_histories - .entry(key) - .or_default() - .push(session); - } - state.training.sessions = active_sessions; - state.training.clamp_selection(); - state.training.gpu = gpu.clone(); - state.system.gpu = gpu; - state.training.connected = true; - } - - StreamEvent::OrderUpdate { - order_id: _, - symbol, - status, - filled_qty, - last_fill_price, - timestamp_nanos, - } => { - let fill = FillRow { - time: format_nanos(timestamp_nanos), - side: if filled_qty >= 0.0 { - "BUY".to_owned() - } else { - "SELL".to_owned() - }, - symbol, - quantity: filled_qty.abs(), - price: last_fill_price, - status, - }; - state.trading.push_fill(fill); - state.trading.connected = true; - } - - StreamEvent::PositionsSnapshot { positions } => { - let unrealized: f64 = positions.iter().map(|p| p.unrealized_pnl).sum(); - state.trading.positions = positions; - state.trading.unrealized_pnl = unrealized; - state.trading.connected = true; - } - - StreamEvent::RiskMetrics { - var, - max_drawdown, - current_drawdown, - sharpe, - } => { - state.risk.portfolio_var = var; - state.risk.max_drawdown = max_drawdown; - state.risk.current_drawdown = current_drawdown; - state.risk.sharpe_ratio = sharpe; - state.risk.connected = true; - } - - StreamEvent::RiskAlert { - severity: _, - symbol: _, - message, - threshold, - current, - } => { - let status = if current < threshold { - "OK".to_owned() - } else { - "TRIPPED".to_owned() - }; - - // Upsert: find existing breaker by name or insert new. - if let Some(cb) = state - .risk - .circuit_breakers - .iter_mut() - .find(|cb| cb.name == message) - { - cb.status = status; - cb.current = format!("{current:.4}"); - cb.limit = format!("{threshold:.4}"); - } else { - state.risk.circuit_breakers.push(CircuitBreakerRow { - name: message, - status, - current: format!("{current:.4}"), - limit: format!("{threshold:.4}"), - }); - } - state.risk.connected = true; - } - - StreamEvent::SystemStatus { - service, - status, - message, - } => { - // Upsert: find existing service by name or insert new. - if let Some(svc) = state - .system - .services - .iter_mut() - .find(|s| s.name == service) - { - svc.status = status; - svc.message = message; - } else { - state.system.services.push(ServiceRow { - name: service, - status, - message, - latency_ms: 0.0, - }); - } - state.system.connected = true; - } - - StreamEvent::StreamDisconnected { stream_name } => { - match stream_name.as_str() { - "training" | "training_metrics" => state.training.connected = false, - "trading" | "orders" | "order_updates" => state.trading.connected = false, - "positions" => state.trading.connected = false, - "risk" | "risk_alerts" => state.risk.connected = false, - "system" | "system_status" => state.system.connected = false, - _ => {} - } - } - } - - state.dirty = true; -} - -// --------------------------------------------------------------------------- -// Key handling -// --------------------------------------------------------------------------- - -/// Returns `true` if the user wants to quit. -#[allow(clippy::wildcard_enum_match_arm)] // KeyCode has 25+ variants; exhaustive match is noise -fn handle_key(state: &mut DashboardState, key: KeyEvent) -> bool { - // Global quit - match key.code { - KeyCode::Char('q') => return true, - KeyCode::Char('c') if key.modifiers.contains(KeyModifiers::CONTROL) => return true, - _ => {} - } - - // Training tab has special modal key handling. - if state.current_tab() == Tab::Training { - match state.training.view_mode { - TrainingViewMode::List => match key.code { - KeyCode::Char('j') | KeyCode::Down => { - state.training.select_down(); - state.dirty = true; - } - KeyCode::Char('k') | KeyCode::Up => { - state.training.select_up(); - state.dirty = true; - } - KeyCode::Enter => { - state.training.enter_detail(); - state.dirty = true; - } - KeyCode::Char('1') => state.select_tab(0), - KeyCode::Char('2') => state.select_tab(1), - KeyCode::Char('3') => state.select_tab(2), - KeyCode::Char('4') => state.select_tab(3), - _ => {} - }, - TrainingViewMode::Detail => match key.code { - KeyCode::Esc => { - state.training.exit_detail(); - state.dirty = true; - } - KeyCode::Tab => { - state.training.next_sub_tab(); - state.dirty = true; - } - KeyCode::BackTab => { - state.training.prev_sub_tab(); - state.dirty = true; - } - KeyCode::Char('l') | KeyCode::Right => { - // Next session - state.training.select_down(); - state.dirty = true; - } - KeyCode::Char('h') | KeyCode::Left => { - // Previous session - state.training.select_up(); - state.dirty = true; - } - KeyCode::Char('j') | KeyCode::Down => { - state.training.scroll_down(); - state.dirty = true; - } - KeyCode::Char('k') | KeyCode::Up => { - state.training.scroll_up(); - state.dirty = true; - } - KeyCode::Char(c @ '1'..='4') => { - state.training.exit_detail(); - state.select_tab((c as u8 - b'1') as usize); - } - _ => {} - }, - } - return false; - } - - // Non-Training tabs: standard key handling. - match key.code { - KeyCode::Char('1') => state.select_tab(0), - KeyCode::Char('2') => state.select_tab(1), - KeyCode::Char('3') => state.select_tab(2), - KeyCode::Char('4') => state.select_tab(3), - KeyCode::Char('j') | KeyCode::Down => state.scroll_down(), - KeyCode::Char('k') | KeyCode::Up => state.scroll_up(), - _ => {} - } - false -} - -// --------------------------------------------------------------------------- -// Helpers -// --------------------------------------------------------------------------- - -/// Format a nanosecond Unix timestamp as `HH:MM:SS`. -#[allow(clippy::modulo_arithmetic)] // Intentional clock arithmetic on non-negative quotients -fn format_nanos(nanos: i64) -> String { - let secs = nanos / 1_000_000_000; - let hour = (secs / 3600).rem_euclid(24); - let min = (secs / 60).rem_euclid(60); - let sec = secs.rem_euclid(60); - format!("{hour:02}:{min:02}:{sec:02}") -} - -// --------------------------------------------------------------------------- -// Tests -// --------------------------------------------------------------------------- - -#[cfg(test)] -#[allow(clippy::unwrap_used, clippy::expect_used)] -mod tests { - use super::*; - use super::super::state::GpuInfo; - use super::super::state::PositionRow; - use super::super::state::TrainingSession; - use crossterm::event::KeyEvent; - - // -- Key handling -------------------------------------------------------- - - #[test] - fn test_handle_key_quit() { - let mut state = DashboardState::default(); - let key = KeyEvent::new(KeyCode::Char('q'), KeyModifiers::NONE); - assert!(handle_key(&mut state, key)); - } - - #[test] - fn test_handle_key_ctrl_c_quit() { - let mut state = DashboardState::default(); - let key = KeyEvent::new(KeyCode::Char('c'), KeyModifiers::CONTROL); - assert!(handle_key(&mut state, key)); - } - - #[test] - fn test_handle_key_tab_switch() { - let mut state = DashboardState::default(); - let key = KeyEvent::new(KeyCode::Char('3'), KeyModifiers::NONE); - assert!(!handle_key(&mut state, key)); - assert_eq!(state.active_tab, 2); - } - - #[test] - fn test_handle_key_scroll_non_training_tab() { - // Non-training tabs still use scroll on j/k. - let mut state = DashboardState::default(); - state.select_tab(1); // Trading tab - let j = KeyEvent::new(KeyCode::Char('j'), KeyModifiers::NONE); - handle_key(&mut state, j); - assert_eq!(state.trading.scroll_offset, 1); - - let k = KeyEvent::new(KeyCode::Char('k'), KeyModifiers::NONE); - handle_key(&mut state, k); - assert_eq!(state.trading.scroll_offset, 0); - } - - #[test] - fn test_handle_key_arrow_scroll_non_training_tab() { - let mut state = DashboardState::default(); - state.select_tab(1); // Trading tab - let down = KeyEvent::new(KeyCode::Down, KeyModifiers::NONE); - handle_key(&mut state, down); - assert_eq!(state.trading.scroll_offset, 1); - - let up = KeyEvent::new(KeyCode::Up, KeyModifiers::NONE); - handle_key(&mut state, up); - assert_eq!(state.trading.scroll_offset, 0); - } - - #[test] - fn test_handle_key_non_quit_returns_false() { - let mut state = DashboardState::default(); - let key = KeyEvent::new(KeyCode::Char('x'), KeyModifiers::NONE); - assert!(!handle_key(&mut state, key)); - } - - // -- Training list mode keys ------------------------------------------- - - fn make_state_with_sessions(n: usize) -> DashboardState { - let mut state = DashboardState::default(); - for i in 0..n { - state.training.sessions.push(TrainingSession { - model: format!("DQN_{i}"), - fold: format!("fold_{i}"), - ..Default::default() - }); - } - state - } - - #[test] - fn test_training_list_jk_selects() { - let mut state = make_state_with_sessions(3); - // j selects first row - let j = KeyEvent::new(KeyCode::Char('j'), KeyModifiers::NONE); - handle_key(&mut state, j); - assert_eq!(state.training.selected_index, Some(0)); - // j again moves to second - handle_key(&mut state, j); - assert_eq!(state.training.selected_index, Some(1)); - // k moves back - let k = KeyEvent::new(KeyCode::Char('k'), KeyModifiers::NONE); - handle_key(&mut state, k); - assert_eq!(state.training.selected_index, Some(0)); - } - - #[test] - fn test_training_list_enter_opens_detail() { - let mut state = make_state_with_sessions(3); - state.training.selected_index = Some(1); - let enter = KeyEvent::new(KeyCode::Enter, KeyModifiers::NONE); - handle_key(&mut state, enter); - assert_eq!(state.training.view_mode, TrainingViewMode::Detail); - } - - // -- Training detail mode keys ----------------------------------------- - - #[test] - fn test_training_detail_esc_exits() { - let mut state = make_state_with_sessions(3); - state.training.selected_index = Some(0); - state.training.enter_detail(); - let esc = KeyEvent::new(KeyCode::Esc, KeyModifiers::NONE); - handle_key(&mut state, esc); - assert_eq!(state.training.view_mode, TrainingViewMode::List); - } - - #[test] - fn test_training_detail_tab_cycles_sub_tabs() { - let mut state = make_state_with_sessions(1); - state.training.selected_index = Some(0); - state.training.enter_detail(); - // DQN_0 -> Overview, Loss, RL, Health - let tab = KeyEvent::new(KeyCode::Tab, KeyModifiers::NONE); - handle_key(&mut state, tab); - assert_eq!( - state.training.detail_sub_tab, - super::super::state::DetailSubTab::Loss - ); - } - - #[test] - fn test_training_detail_backtab_cycles_back() { - let mut state = make_state_with_sessions(1); - state.training.selected_index = Some(0); - state.training.enter_detail(); - state.training.detail_sub_tab = super::super::state::DetailSubTab::Loss; - let backtab = KeyEvent::new(KeyCode::BackTab, KeyModifiers::SHIFT); - handle_key(&mut state, backtab); - assert_eq!( - state.training.detail_sub_tab, - super::super::state::DetailSubTab::Overview - ); - } - - #[test] - fn test_training_detail_lr_navigates_sessions() { - let mut state = make_state_with_sessions(3); - state.training.selected_index = Some(0); - state.training.enter_detail(); - // Right moves to next session - let right = KeyEvent::new(KeyCode::Right, KeyModifiers::NONE); - handle_key(&mut state, right); - assert_eq!(state.training.selected_index, Some(1)); - // Left moves back - let left = KeyEvent::new(KeyCode::Left, KeyModifiers::NONE); - handle_key(&mut state, left); - assert_eq!(state.training.selected_index, Some(0)); - } - - #[test] - fn test_training_detail_number_exits_and_switches_tab() { - let mut state = make_state_with_sessions(3); - state.training.selected_index = Some(0); - state.training.enter_detail(); - let key2 = KeyEvent::new(KeyCode::Char('2'), KeyModifiers::NONE); - handle_key(&mut state, key2); - assert_eq!(state.training.view_mode, TrainingViewMode::List); - assert_eq!(state.active_tab, 1); // Trading tab - } - - #[test] - fn test_training_detail_q_quits() { - let mut state = make_state_with_sessions(1); - state.training.selected_index = Some(0); - state.training.enter_detail(); - let q = KeyEvent::new(KeyCode::Char('q'), KeyModifiers::NONE); - assert!(handle_key(&mut state, q)); - } - - // -- format_nanos -------------------------------------------------------- - - #[test] - fn test_format_nanos() { - assert_eq!(format_nanos(52_321_000_000_000), "14:32:01"); - } - - #[test] - fn test_format_nanos_zero() { - assert_eq!(format_nanos(0), "00:00:00"); - } - - #[test] - fn test_format_nanos_midnight() { - // Exactly 24 hours in nanos wraps to 00:00:00. - let nanos = 24 * 3600 * 1_000_000_000_i64; - assert_eq!(format_nanos(nanos), "00:00:00"); - } - - // -- apply_stream_event -------------------------------------------------- - - #[test] - fn test_apply_training_update() { - let mut state = DashboardState::default(); - let session = TrainingSession { - model: "DQN".into(), - fold: "fold0".into(), - is_hyperopt: false, - epoch: 5.0, - epoch_loss: 0.123, - validation_loss: 0.456, - learning_rate: 1e-3, - gpu_percent: 85.0, - batches_per_second: 42.0, - gradient_norm: 1.5, - epoch_duration_seconds: 12.3, - ..Default::default() - }; - let gpu = GpuInfo { - utilization_percent: 85.0, - memory_used_mb: 3000.0, - memory_total_mb: 48000.0, - temperature_celsius: 62.0, - }; - - apply_stream_event( - &mut state, - StreamEvent::TrainingUpdate { - sessions: vec![session], - gpu: gpu.clone(), - }, - ); - - assert!(state.training.connected); - assert_eq!(state.training.sessions.len(), 1); - assert_eq!(state.training.loss_history.len(), 1); - // Loss should be epoch_loss as f64. - assert!((state.training.loss_history.back().unwrap() - 0.123_f64).abs() < 1e-3); - // GPU should be cloned to both training and system. - assert!((state.system.gpu.utilization_percent - 85.0).abs() < f32::EPSILON); - // Per-session history should be accumulated. - assert_eq!(state.training.session_histories.len(), 1); - let hist = state.training.session_histories.get("DQN:fold0").unwrap(); - assert_eq!(hist.loss.len(), 1); - assert!((hist.loss[0] - 0.123_f64).abs() < 1e-3); - assert!(state.dirty); - } - - #[test] - fn test_apply_order_update() { - let mut state = DashboardState::default(); - apply_stream_event( - &mut state, - StreamEvent::OrderUpdate { - order_id: "ord-1".into(), - symbol: "ES".into(), - status: "Filled".into(), - filled_qty: 2.0, - last_fill_price: 5100.25, - timestamp_nanos: 52_321_000_000_000, - }, - ); - - assert!(state.trading.connected); - assert_eq!(state.trading.fills.len(), 1); - let fill = state.trading.fills.front().unwrap(); - assert_eq!(fill.time, "14:32:01"); - assert_eq!(fill.side, "BUY"); - assert_eq!(fill.symbol, "ES"); - assert!((fill.quantity - 2.0).abs() < f64::EPSILON); - } - - #[test] - fn test_apply_order_update_sell() { - let mut state = DashboardState::default(); - apply_stream_event( - &mut state, - StreamEvent::OrderUpdate { - order_id: "ord-2".into(), - symbol: "NQ".into(), - status: "Filled".into(), - filled_qty: -1.0, - last_fill_price: 18000.0, - timestamp_nanos: 0, - }, - ); - - let fill = state.trading.fills.front().unwrap(); - assert_eq!(fill.side, "SELL"); - assert!((fill.quantity - 1.0).abs() < f64::EPSILON); - } - - #[test] - fn test_apply_positions_snapshot() { - let mut state = DashboardState::default(); - let positions = vec![ - PositionRow { - symbol: "ES".into(), - side: "LONG".into(), - quantity: 2.0, - entry_price: 5000.0, - unrealized_pnl: 150.0, - status: "Open".into(), - }, - PositionRow { - symbol: "NQ".into(), - side: "SHORT".into(), - quantity: 1.0, - entry_price: 18000.0, - unrealized_pnl: -50.0, - status: "Open".into(), - }, - ]; - - apply_stream_event( - &mut state, - StreamEvent::PositionsSnapshot { - positions: positions.clone(), - }, - ); - - assert!(state.trading.connected); - assert_eq!(state.trading.positions.len(), 2); - assert!((state.trading.unrealized_pnl - 100.0).abs() < f64::EPSILON); - } - - #[test] - fn test_apply_risk_metrics() { - let mut state = DashboardState::default(); - apply_stream_event( - &mut state, - StreamEvent::RiskMetrics { - var: 0.02, - max_drawdown: 0.05, - current_drawdown: 0.01, - sharpe: 2.5, - }, - ); - - assert!(state.risk.connected); - assert!((state.risk.portfolio_var - 0.02).abs() < f64::EPSILON); - assert!((state.risk.max_drawdown - 0.05).abs() < f64::EPSILON); - assert!((state.risk.current_drawdown - 0.01).abs() < f64::EPSILON); - assert!((state.risk.sharpe_ratio - 2.5).abs() < f64::EPSILON); - } - - #[test] - fn test_apply_risk_alert_insert() { - let mut state = DashboardState::default(); - apply_stream_event( - &mut state, - StreamEvent::RiskAlert { - severity: "WARNING".into(), - symbol: "ES".into(), - message: "max_drawdown".into(), - threshold: 0.05, - current: 0.03, - }, - ); - - assert!(state.risk.connected); - assert_eq!(state.risk.circuit_breakers.len(), 1); - assert_eq!(state.risk.circuit_breakers[0].name, "max_drawdown"); - assert_eq!(state.risk.circuit_breakers[0].status, "OK"); - } - - #[test] - fn test_apply_risk_alert_tripped() { - let mut state = DashboardState::default(); - apply_stream_event( - &mut state, - StreamEvent::RiskAlert { - severity: "CRITICAL".into(), - symbol: "ES".into(), - message: "max_drawdown".into(), - threshold: 0.05, - current: 0.06, - }, - ); - - assert_eq!(state.risk.circuit_breakers[0].status, "TRIPPED"); - } - - #[test] - fn test_apply_system_status_upsert() { - let mut state = DashboardState::default(); - - // Insert first time. - apply_stream_event( - &mut state, - StreamEvent::SystemStatus { - service: "gateway".into(), - status: "HEALTHY".into(), - message: "ok".into(), - }, - ); - assert_eq!(state.system.services.len(), 1); - - // Update same service. - apply_stream_event( - &mut state, - StreamEvent::SystemStatus { - service: "gateway".into(), - status: "DEGRADED".into(), - message: "high latency".into(), - }, - ); - - // Should still be 1 entry, not 2. - assert_eq!(state.system.services.len(), 1); - assert_eq!(state.system.services[0].status, "DEGRADED"); - assert_eq!(state.system.services[0].message, "high latency"); - assert!(state.system.connected); - } - - #[test] - fn test_apply_stream_disconnected() { - let mut state = DashboardState::default(); - state.training.connected = true; - state.trading.connected = true; - state.risk.connected = true; - state.system.connected = true; - - apply_stream_event( - &mut state, - StreamEvent::StreamDisconnected { - stream_name: "training".into(), - }, - ); - assert!(!state.training.connected); - - apply_stream_event( - &mut state, - StreamEvent::StreamDisconnected { - stream_name: "risk".into(), - }, - ); - assert!(!state.risk.connected); - - apply_stream_event( - &mut state, - StreamEvent::StreamDisconnected { - stream_name: "system".into(), - }, - ); - assert!(!state.system.connected); - - apply_stream_event( - &mut state, - StreamEvent::StreamDisconnected { - stream_name: "orders".into(), - }, - ); - assert!(!state.trading.connected); - - // Dirty flag should be set after each event. - assert!(state.dirty); - } - - #[test] - fn test_training_update_filters_empty_model() { - let mut state = DashboardState::default(); - let valid = TrainingSession { - model: "DQN".into(), - fold: "fold0".into(), - epoch_loss: 0.5, - ..Default::default() - }; - let empty = TrainingSession { - model: String::new(), - fold: String::new(), - ..Default::default() - }; - // Summary row: model present but no fold and epoch=0 (aggregate row from proto) - let summary = TrainingSession { - model: "DQN".into(), - fold: String::new(), - epoch: 0.0, - ..Default::default() - }; - apply_stream_event( - &mut state, - StreamEvent::TrainingUpdate { - sessions: vec![empty, summary, valid], - gpu: GpuInfo::default(), - }, - ); - // Only the session with a non-empty model AND a fold (or non-zero epoch) should remain. - assert_eq!(state.training.sessions.len(), 1); - assert_eq!(state.training.sessions[0].model, "DQN"); - assert_eq!(state.training.session_histories.len(), 1); - assert!(state.training.session_histories.contains_key("DQN:fold0")); - } -} diff --git a/bin/fxt/src/commands/watch/mod.rs b/bin/fxt/src/commands/watch/mod.rs deleted file mode 100644 index 9241db720..000000000 --- a/bin/fxt/src/commands/watch/mod.rs +++ /dev/null @@ -1,10 +0,0 @@ -pub mod state; -mod streams; -mod render; -mod event_loop; - -use anyhow::Result; - -pub async fn run(api_gateway_url: &str, jwt_token: &str) -> Result<()> { - event_loop::run_dashboard(api_gateway_url, jwt_token).await -} diff --git a/bin/fxt/src/commands/watch/render.rs b/bin/fxt/src/commands/watch/render.rs deleted file mode 100644 index 6653716b6..000000000 --- a/bin/fxt/src/commands/watch/render.rs +++ /dev/null @@ -1,947 +0,0 @@ -//! Ratatui renderer for the `fxt watch` dashboard. -//! -//! One public entry point ([`render`]) dispatches to per-tab helpers. -//! No tests — visual rendering tests are brittle and not worth maintaining. - -use ratatui::{ - layout::{Constraint, Direction, Layout, Rect}, - style::{Color, Modifier, Style}, - text::{Line, Span}, - widgets::{Block, Borders, Cell, Gauge, Paragraph, Row, Sparkline, Table, TableState, Tabs}, - Frame, -}; - -use super::state::{ - DashboardState, DetailSubTab, RiskTabState, SessionHistory, SystemTabState, Tab, - TradingTabState, TrainingTabState, TrainingViewMode, -}; - -// --------------------------------------------------------------------------- -// Public entry point -// --------------------------------------------------------------------------- - -/// Render the full dashboard into the given frame. -pub(super) fn render(frame: &mut Frame, state: &DashboardState) { - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Length(3), // tab bar - Constraint::Min(0), // content - Constraint::Length(1), // status bar - ]) - .split(frame.area()); - - render_tab_bar(frame, chunks[0], state); - - match state.current_tab() { - Tab::Training => match state.training.view_mode { - TrainingViewMode::List => render_training_list(frame, chunks[1], &state.training), - TrainingViewMode::Detail => { - render_training_detail(frame, chunks[1], &state.training); - } - }, - Tab::Trading => render_trading_tab(frame, chunks[1], &state.trading), - Tab::Risk => render_risk_tab(frame, chunks[1], &state.risk), - Tab::System => render_system_tab(frame, chunks[1], &state.system), - } - - render_status_bar(frame, chunks[2], state); -} - -// --------------------------------------------------------------------------- -// Tab bar -// --------------------------------------------------------------------------- - -fn render_tab_bar(frame: &mut Frame, area: Rect, state: &DashboardState) { - let titles: Vec> = Tab::ALL - .iter() - .map(|t| Line::from(Span::raw(t.title()))) - .collect(); - - let tabs = Tabs::new(titles) - .block( - Block::default() - .borders(Borders::ALL) - .title(" fxt watch "), - ) - .select(state.active_tab) - .highlight_style( - Style::default() - .fg(Color::Cyan) - .add_modifier(Modifier::BOLD), - ); - - frame.render_widget(tabs, area); -} - -// --------------------------------------------------------------------------- -// Status bar -// --------------------------------------------------------------------------- - -fn render_status_bar(frame: &mut Frame, area: Rect, state: &DashboardState) { - let hints: Vec> = - if state.current_tab() == Tab::Training && state.training.view_mode == TrainingViewMode::Detail { - vec![ - Span::styled("Esc", Style::default().fg(Color::Cyan)), - Span::raw(" back "), - Span::styled("Tab/S-Tab", Style::default().fg(Color::Cyan)), - Span::raw(" sub-tab "), - Span::styled("h/l", Style::default().fg(Color::Cyan)), - Span::raw(" prev/next "), - Span::styled("j/k", Style::default().fg(Color::Cyan)), - Span::raw(" scroll "), - Span::styled("q", Style::default().fg(Color::Cyan)), - Span::raw(" quit"), - ] - } else if state.current_tab() == Tab::Training { - vec![ - Span::styled("1-4", Style::default().fg(Color::Cyan)), - Span::raw(" tabs "), - Span::styled("j/k", Style::default().fg(Color::Cyan)), - Span::raw(" select "), - Span::styled("Enter", Style::default().fg(Color::Cyan)), - Span::raw(" detail "), - Span::styled("q", Style::default().fg(Color::Cyan)), - Span::raw(" quit"), - ] - } else { - vec![ - Span::styled("1-4", Style::default().fg(Color::Cyan)), - Span::raw(" tabs "), - Span::styled("j/k", Style::default().fg(Color::Cyan)), - Span::raw(" scroll "), - Span::styled("q", Style::default().fg(Color::Cyan)), - Span::raw(" quit"), - ] - }; - let bar = Paragraph::new(Line::from(hints)); - frame.render_widget(bar, area); -} - -// --------------------------------------------------------------------------- -// Training tab -// --------------------------------------------------------------------------- - -fn render_training_list(frame: &mut Frame, area: Rect, tab: &TrainingTabState) { - let conn_label = if tab.connected { - "connected" - } else { - "disconnected" - }; - - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Min(0), // sessions table - Constraint::Length(6), // GPU + sparkline - ]) - .split(area); - - // -- Sessions table with selection ------------------------------------ - let header = Row::new(vec![ - Cell::from("Model"), - Cell::from("Fold"), - Cell::from("Epoch"), - Cell::from("Loss"), - Cell::from("Val Loss"), - Cell::from("Acc"), - Cell::from("F1"), - Cell::from("LR"), - Cell::from("Batch/s"), - Cell::from("Grad"), - Cell::from("Sharpe"), - Cell::from("Win%"), - ]) - .style( - Style::default() - .fg(Color::Yellow) - .add_modifier(Modifier::BOLD), - ); - - let rows: Vec> = tab - .sessions - .iter() - .map(|s| { - Row::new(vec![ - Cell::from(s.model.as_str()), - Cell::from(s.fold.as_str()), - Cell::from(format!("{:.1}", s.epoch)), - Cell::from(format!("{:.4}", s.epoch_loss)), - Cell::from(format!("{:.4}", s.validation_loss)), - Cell::from(format!("{:.4}", s.eval_accuracy)), - Cell::from(format!("{:.4}", s.eval_f1)), - Cell::from(format!("{:.2e}", s.learning_rate)), - Cell::from(format!("{:.1}", s.batches_per_second)), - Cell::from(format!("{:.3}", s.gradient_norm)), - Cell::from(if s.epoch_sharpe != 0.0 { format!("{:.2}", s.epoch_sharpe) } else { "-".to_owned() }), - Cell::from(if s.epoch_win_rate > 0.0 { format!("{:.0}%", s.epoch_win_rate * 100.0) } else { "-".to_owned() }), - ]) - }) - .collect(); - - let table = Table::new( - rows, - [ - Constraint::Length(10), - Constraint::Length(8), - Constraint::Length(8), - Constraint::Length(10), - Constraint::Length(10), - Constraint::Length(8), - Constraint::Length(8), - Constraint::Length(10), - Constraint::Length(8), - Constraint::Length(8), - Constraint::Length(8), - Constraint::Length(8), - ], - ) - .header(header) - .row_highlight_style( - Style::default() - .fg(Color::Cyan) - .add_modifier(Modifier::BOLD), - ) - .block( - Block::default() - .borders(Borders::ALL) - .title(format!(" Training [{conn_label}] ")), - ); - - let mut table_state = TableState::default().with_selected(tab.selected_index); - frame.render_stateful_widget(table, chunks[0], &mut table_state); - - // -- Bottom row: GPU info (left) + loss sparkline (right) ------------- - let bottom = Layout::default() - .direction(Direction::Horizontal) - .constraints([Constraint::Percentage(50), Constraint::Percentage(50)]) - .split(chunks[1]); - - let gpu = &tab.gpu; - let gpu_text = format!( - "GPU: {:.0}% VRAM: {:.0}/{:.0} MB Temp: {:.0}C", - gpu.utilization_percent, gpu.memory_used_mb, gpu.memory_total_mb, gpu.temperature_celsius, - ); - let gpu_para = Paragraph::new(gpu_text).block( - Block::default() - .borders(Borders::ALL) - .title(" GPU "), - ); - frame.render_widget(gpu_para, bottom[0]); - - // Sparkline: clamp to 0..10, multiply by 100 for u64 resolution. - let spark_data: Vec = tab - .loss_history - .iter() - .map(|&v| { - let clamped = v.clamp(0.0, 10.0); - (clamped * 100.0) as u64 - }) - .collect(); - - let sparkline = Sparkline::default() - .block( - Block::default() - .borders(Borders::ALL) - .title(" Loss "), - ) - .data(&spark_data) - .style(Style::default().fg(Color::Cyan)); - - frame.render_widget(sparkline, bottom[1]); -} - -// --------------------------------------------------------------------------- -// Training detail view -// --------------------------------------------------------------------------- - -/// Convert a `&[f64]` metric history to Sparkline-compatible `Vec`. -fn spark_u64(data: &[f64], max_clamp: f64) -> Vec { - data.iter() - .map(|&v| { - let clamped = v.clamp(0.0, max_clamp); - (clamped * 100.0) as u64 - }) - .collect() -} - -fn render_sparkline_row(frame: &mut Frame, area: Rect, title: &str, data: &[f64], max: f64, color: Color) { - let vals = spark_u64(data, max); - let sparkline = Sparkline::default() - .block(Block::default().borders(Borders::ALL).title(format!(" {title} "))) - .data(&vals) - .style(Style::default().fg(color)); - frame.render_widget(sparkline, area); -} - -fn render_training_detail(frame: &mut Frame, area: Rect, tab: &TrainingTabState) { - let Some(session) = tab.selected_session() else { - let msg = Paragraph::new("No session selected") - .block(Block::default().borders(Borders::ALL).title(" Detail ")); - frame.render_widget(msg, area); - return; - }; - - let vis_tabs = tab.visible_tabs(); - let history_key = format!("{}:{}", session.model, session.fold); - let empty_hist = SessionHistory::default(); - let history = tab.session_histories.get(&history_key).unwrap_or(&empty_hist); - - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Length(3), // title bar - Constraint::Length(3), // sub-tab bar - Constraint::Min(0), // content - ]) - .split(area); - - // -- Title bar -------------------------------------------------------- - let status_label = if session.is_hyperopt { - format!( - "Hyperopt trial {}/{}", - session.hyperopt_trial_current, session.hyperopt_trial_total, - ) - } else { - format!("Epoch {:.1}", session.epoch) - }; - let title_text = format!( - " {} | {} | {} | GPU {:.0}%", - session.model, session.fold, status_label, tab.gpu.utilization_percent, - ); - let title_bar = Paragraph::new(title_text).block( - Block::default() - .borders(Borders::ALL) - .title(" Session Detail "), - ); - frame.render_widget(title_bar, chunks[0]); - - // -- Sub-tab bar ------------------------------------------------------ - let tab_titles: Vec> = vis_tabs - .iter() - .map(|t| Line::from(Span::raw(t.title()))) - .collect(); - let active_idx = vis_tabs - .iter() - .position(|t| *t == tab.detail_sub_tab) - .unwrap_or(0); - let sub_tabs = Tabs::new(tab_titles) - .select(active_idx) - .highlight_style( - Style::default() - .fg(Color::Cyan) - .add_modifier(Modifier::BOLD), - ) - .block(Block::default().borders(Borders::ALL)); - frame.render_widget(sub_tabs, chunks[1]); - - // -- Content area: dispatch to sub-tab renderer ----------------------- - match tab.detail_sub_tab { - DetailSubTab::Overview => render_detail_overview(frame, chunks[2], session, history), - DetailSubTab::Loss => render_detail_loss(frame, chunks[2], history), - DetailSubTab::RlDiagnostics => render_detail_rl(frame, chunks[2], session, history), - DetailSubTab::Metrics => render_detail_metrics(frame, chunks[2], session, history), - DetailSubTab::Hyperopt => render_detail_hyperopt(frame, chunks[2], session, history), - DetailSubTab::Health => render_detail_health(frame, chunks[2], session), - } -} - -// -- Sub-tab renderers ---------------------------------------------------- - -fn render_detail_overview( - frame: &mut Frame, - area: Rect, - session: &super::state::TrainingSession, - history: &SessionHistory, -) { - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Length(12), // key stats (incl. financial metrics) - Constraint::Min(0), // loss sparkline - ]) - .split(area); - - // Key stats grid - let stats = Paragraph::new(vec![ - Line::from(format!( - " Loss: {:.4} Val Loss: {:.4} LR: {:.2e}", - session.epoch_loss, session.validation_loss, session.learning_rate, - )), - Line::from(format!( - " Batch/s: {:.1} Grad Norm: {:.3} Epoch Time: {:.1}s", - session.batches_per_second, session.gradient_norm, session.epoch_duration_seconds, - )), - Line::from(format!( - " Batches: {:.0} Iter: {:.3}s Checkpoints: {} ({} failed)", - session.batches_processed, - session.iteration_seconds, - session.checkpoint_saves, - session.checkpoint_failures, - )), - Line::from(format!( - " NaN: {} Grad Explosions: {} Feature Errors: {}", - session.nan_detected, session.gradient_explosions, session.feature_errors, - )), - Line::from(""), - Line::from(format!( - " Sharpe: {:.2} Sortino: {:.2} Win Rate: {:.1}% Max DD: {:.1}%", - session.epoch_sharpe, session.epoch_sortino, - session.epoch_win_rate * 100.0, session.epoch_max_drawdown * 100.0, - )), - Line::from(format!( - " PF: {:.2} Return: {:+.2}% Avg: {:+.4} Trades: {}", - session.epoch_profit_factor, session.epoch_total_return * 100.0, - session.epoch_avg_return, session.epoch_total_trades, - )), - ]) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Overview "), - ); - frame.render_widget(stats, chunks[0]); - - render_sparkline_row(frame, chunks[1], "Loss", &history.loss, 10.0, Color::Cyan); -} - -fn render_detail_loss(frame: &mut Frame, area: Rect, history: &SessionHistory) { - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Percentage(25), - Constraint::Percentage(25), - Constraint::Percentage(25), - Constraint::Percentage(25), - ]) - .split(area); - - render_sparkline_row(frame, chunks[0], "Loss", &history.loss, 10.0, Color::Cyan); - render_sparkline_row(frame, chunks[1], "Val Loss", &history.val_loss, 10.0, Color::Yellow); - render_sparkline_row(frame, chunks[2], "Learning Rate", &history.learning_rate, 1.0, Color::Green); - render_sparkline_row(frame, chunks[3], "Grad Norm", &history.gradient_norm, 100.0, Color::Red); -} - -fn render_detail_rl( - frame: &mut Frame, - area: Rect, - session: &super::state::TrainingSession, - history: &SessionHistory, -) { - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Percentage(20), - Constraint::Percentage(20), - Constraint::Percentage(20), - Constraint::Percentage(20), - Constraint::Percentage(20), - ]) - .split(area); - - render_sparkline_row(frame, chunks[0], "Q-Value Mean", &history.q_value_mean, 100.0, Color::Cyan); - render_sparkline_row(frame, chunks[1], "Policy Entropy", &history.policy_entropy, 10.0, Color::Yellow); - render_sparkline_row(frame, chunks[2], "KL Divergence", &history.kl_divergence, 1.0, Color::Red); - render_sparkline_row(frame, chunks[3], "Advantage Mean", &history.advantage_mean, 10.0, Color::Green); - - // Replay buffer gauge - let buf_size = session.replay_buffer_size; - let ratio = if buf_size > 0 { - (buf_size as f64 / 100_000.0).min(1.0) - } else { - 0.0 - }; - let gauge = Gauge::default() - .block( - Block::default() - .borders(Borders::ALL) - .title(format!(" Replay Buffer: {buf_size} ")), - ) - .gauge_style(Style::default().fg(Color::Cyan)) - .ratio(ratio); - frame.render_widget(gauge, chunks[4]); -} - -fn render_detail_metrics( - frame: &mut Frame, - area: Rect, - session: &super::state::TrainingSession, - history: &SessionHistory, -) { - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Length(7), // current values (incl. action distribution) - Constraint::Min(0), // sparklines - ]) - .split(area); - - let current = Paragraph::new(vec![ - Line::from(format!( - " Accuracy: {:.4} F1: {:.4}", - session.eval_accuracy, session.eval_f1, - )), - Line::from(format!( - " Precision: {:.4} Recall: {:.4}", - session.eval_precision, session.eval_recall, - )), - Line::from(format!( - " Action: BUY {:.0}% SELL {:.0}% HOLD {:.0}%", - session.action_buy_pct * 100.0, - session.action_sell_pct * 100.0, - session.action_hold_pct * 100.0, - )), - ]) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Current Metrics "), - ); - frame.render_widget(current, chunks[0]); - - let spark_area = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Ratio(1, 8), - Constraint::Ratio(1, 8), - Constraint::Ratio(1, 8), - Constraint::Ratio(1, 8), - Constraint::Ratio(1, 8), - Constraint::Ratio(1, 8), - Constraint::Ratio(1, 8), - Constraint::Ratio(1, 8), - ]) - .split(chunks[1]); - - render_sparkline_row(frame, spark_area[0], "Accuracy", &history.accuracy, 1.0, Color::Green); - render_sparkline_row(frame, spark_area[1], "Precision", &history.precision, 1.0, Color::Cyan); - render_sparkline_row(frame, spark_area[2], "Recall", &history.recall, 1.0, Color::Magenta); - render_sparkline_row(frame, spark_area[3], "F1", &history.f1, 1.0, Color::Yellow); - render_sparkline_row(frame, spark_area[4], "Sharpe", &history.sharpe, 20.0, Color::Cyan); - render_sparkline_row(frame, spark_area[5], "Win Rate", &history.win_rate, 1.0, Color::Green); - render_sparkline_row(frame, spark_area[6], "Max DD", &history.max_drawdown, 1.0, Color::Red); - render_sparkline_row(frame, spark_area[7], "Total Return", &history.total_return, 1.0, Color::Magenta); -} - -fn render_detail_hyperopt( - frame: &mut Frame, - area: Rect, - session: &super::state::TrainingSession, - history: &SessionHistory, -) { - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Length(7), // stats - Constraint::Min(0), // best objective sparkline - ]) - .split(area); - - let pct = if session.hyperopt_trial_total > 0 { - (session.hyperopt_trial_current as f32 / session.hyperopt_trial_total as f32) * 100.0 - } else { - 0.0 - }; - let stats = Paragraph::new(vec![ - Line::from(format!( - " Trial: {}/{} ({:.0}%)", - session.hyperopt_trial_current, session.hyperopt_trial_total, pct, - )), - Line::from(format!( - " Best Objective: {:.6} Trial Best Loss: {:.6}", - session.hyperopt_best_objective, session.hyperopt_trial_best_loss, - )), - Line::from(format!( - " Trial Epoch: {} Failed: {} Elapsed: {:.0}s", - session.hyperopt_trial_epoch, - session.hyperopt_trials_failed, - session.hyperopt_elapsed_seconds, - )), - ]) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Hyperopt Progress "), - ); - frame.render_widget(stats, chunks[0]); - - render_sparkline_row( - frame, - chunks[1], - "Best Objective", - &history.hyperopt_best_objective, - 10.0, - Color::Magenta, - ); -} - -fn render_detail_health( - frame: &mut Frame, - area: Rect, - session: &super::state::TrainingSession, -) { - let health = Paragraph::new(vec![ - Line::from(format!( - " NaN Detected: {}", - session.nan_detected, - )), - Line::from(format!( - " Gradient Explosions: {}", - session.gradient_explosions, - )), - Line::from(format!( - " Feature Errors: {}", - session.feature_errors, - )), - Line::from(""), - Line::from(format!( - " Checkpoint Saves: {} Failures: {}", - session.checkpoint_saves, session.checkpoint_failures, - )), - Line::from(format!( - " Checkpoint Size: {:.1} MB", - session.checkpoint_size_bytes / (1024.0 * 1024.0), - )), - ]) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Health "), - ); - frame.render_widget(health, area); -} - -// --------------------------------------------------------------------------- -// Trading tab -// --------------------------------------------------------------------------- - -fn render_trading_tab(frame: &mut Frame, area: Rect, tab: &TradingTabState) { - let conn_label = if tab.connected { - "connected" - } else { - "disconnected" - }; - - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Percentage(40), // positions - Constraint::Percentage(40), // fills - Constraint::Length(3), // PnL bar - ]) - .split(area); - - // -- Positions table -------------------------------------------------- - let pos_header = Row::new(vec![ - Cell::from("Symbol"), - Cell::from("Side"), - Cell::from("Qty"), - Cell::from("Entry"), - Cell::from("PnL"), - Cell::from("Status"), - ]) - .style( - Style::default() - .fg(Color::Yellow) - .add_modifier(Modifier::BOLD), - ); - - let pos_rows: Vec> = tab - .positions - .iter() - .map(|p| { - let pnl_color = if p.unrealized_pnl >= 0.0 { - Color::Green - } else { - Color::Red - }; - Row::new(vec![ - Cell::from(p.symbol.as_str()), - Cell::from(p.side.as_str()), - Cell::from(format!("{:.0}", p.quantity)), - Cell::from(format!("{:.2}", p.entry_price)), - Cell::from(Span::styled( - format!("{:+.2}", p.unrealized_pnl), - Style::default().fg(pnl_color), - )), - Cell::from(p.status.as_str()), - ]) - }) - .collect(); - - let pos_table = Table::new( - pos_rows, - [ - Constraint::Length(10), - Constraint::Length(6), - Constraint::Length(6), - Constraint::Length(10), - Constraint::Length(12), - Constraint::Length(8), - ], - ) - .header(pos_header) - .block( - Block::default() - .borders(Borders::ALL) - .title(format!(" Positions [{conn_label}] ")), - ); - - frame.render_widget(pos_table, chunks[0]); - - // -- Fills table ------------------------------------------------------ - let fill_header = Row::new(vec![ - Cell::from("Time"), - Cell::from("Side"), - Cell::from("Symbol"), - Cell::from("Qty"), - Cell::from("Price"), - Cell::from("Status"), - ]) - .style( - Style::default() - .fg(Color::Yellow) - .add_modifier(Modifier::BOLD), - ); - - let fill_rows: Vec> = tab - .fills - .iter() - .skip(tab.scroll_offset) - .take(10) - .map(|f| { - Row::new(vec![ - Cell::from(f.time.as_str()), - Cell::from(f.side.as_str()), - Cell::from(f.symbol.as_str()), - Cell::from(format!("{:.0}", f.quantity)), - Cell::from(format!("{:.2}", f.price)), - Cell::from(f.status.as_str()), - ]) - }) - .collect(); - - let fill_table = Table::new( - fill_rows, - [ - Constraint::Length(10), - Constraint::Length(6), - Constraint::Length(10), - Constraint::Length(6), - Constraint::Length(10), - Constraint::Length(8), - ], - ) - .header(fill_header) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Fills "), - ); - - frame.render_widget(fill_table, chunks[1]); - - // -- PnL bar ---------------------------------------------------------- - let pnl_color = if tab.day_pnl >= 0.0 { - Color::Green - } else { - Color::Red - }; - let pnl_text = Line::from(vec![Span::styled( - format!( - "Day PnL: {:+.2} | Realized: {:+.2} Unrealized: {:+.2}", - tab.day_pnl, tab.realized_pnl, tab.unrealized_pnl, - ), - Style::default().fg(pnl_color), - )]); - let pnl_bar = Paragraph::new(pnl_text).block( - Block::default() - .borders(Borders::ALL) - .title(" PnL "), - ); - frame.render_widget(pnl_bar, chunks[2]); -} - -// --------------------------------------------------------------------------- -// Risk tab -// --------------------------------------------------------------------------- - -fn render_risk_tab(frame: &mut Frame, area: Rect, tab: &RiskTabState) { - let conn_label = if tab.connected { - "connected" - } else { - "disconnected" - }; - - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Length(5), // summary - Constraint::Min(0), // circuit breakers - ]) - .split(area); - - // -- Metrics summary -------------------------------------------------- - let summary = Paragraph::new(vec![ - Line::from(format!( - " VaR: {:.2} Max Drawdown: {:.2}%", - tab.portfolio_var, - tab.max_drawdown * 100.0, - )), - Line::from(format!( - " Sharpe: {:.2} Current Drawdown: {:.2}%", - tab.sharpe_ratio, - tab.current_drawdown * 100.0, - )), - ]) - .block( - Block::default() - .borders(Borders::ALL) - .title(format!(" Risk [{conn_label}] ")), - ); - frame.render_widget(summary, chunks[0]); - - // -- Circuit breakers table ------------------------------------------- - let cb_header = Row::new(vec![ - Cell::from("Breaker"), - Cell::from("Status"), - Cell::from("Current"), - Cell::from("Limit"), - ]) - .style( - Style::default() - .fg(Color::Yellow) - .add_modifier(Modifier::BOLD), - ); - - let cb_rows: Vec> = tab - .circuit_breakers - .iter() - .skip(tab.scroll_offset) - .map(|cb| { - let status_color = if cb.status == "OK" { - Color::Green - } else { - Color::Red - }; - Row::new(vec![ - Cell::from(cb.name.as_str()), - Cell::from(Span::styled( - cb.status.as_str(), - Style::default().fg(status_color), - )), - Cell::from(cb.current.as_str()), - Cell::from(cb.limit.as_str()), - ]) - }) - .collect(); - - let cb_table = Table::new( - cb_rows, - [ - Constraint::Length(18), - Constraint::Length(10), - Constraint::Length(15), - Constraint::Length(15), - ], - ) - .header(cb_header) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Circuit Breakers "), - ); - - frame.render_widget(cb_table, chunks[1]); -} - -// --------------------------------------------------------------------------- -// System tab -// --------------------------------------------------------------------------- - -fn render_system_tab(frame: &mut Frame, area: Rect, tab: &SystemTabState) { - let conn_label = if tab.connected { - "connected" - } else { - "disconnected" - }; - - let chunks = Layout::default() - .direction(Direction::Vertical) - .constraints([ - Constraint::Min(0), // services table - Constraint::Length(5), // resources bar - ]) - .split(area); - - // -- Services table --------------------------------------------------- - let svc_header = Row::new(vec![ - Cell::from("Service"), - Cell::from("Status"), - Cell::from("Latency"), - Cell::from("Message"), - ]) - .style( - Style::default() - .fg(Color::Yellow) - .add_modifier(Modifier::BOLD), - ); - - let svc_rows: Vec> = tab - .services - .iter() - .skip(tab.scroll_offset) - .map(|s| { - let status_color = match s.status.as_str() { - "HEALTHY" | "UP" => Color::Green, - "DEGRADED" => Color::Yellow, - _ => Color::Red, - }; - Row::new(vec![ - Cell::from(s.name.as_str()), - Cell::from(Span::styled( - s.status.as_str(), - Style::default().fg(status_color), - )), - Cell::from(format!("{:.1}ms", s.latency_ms)), - Cell::from(s.message.as_str()), - ]) - }) - .collect(); - - let svc_table = Table::new( - svc_rows, - [ - Constraint::Length(22), - Constraint::Length(12), - Constraint::Length(10), - Constraint::Min(0), - ], - ) - .header(svc_header) - .block( - Block::default() - .borders(Borders::ALL) - .title(format!(" Services [{conn_label}] ")), - ); - - frame.render_widget(svc_table, chunks[0]); - - // -- Resources bar ---------------------------------------------------- - let gpu = &tab.gpu; - let resources = Paragraph::new(format!( - " GPU: {:.0}% VRAM: {:.0}/{:.0} MB CPU: {:.1}% RAM: {:.1}/{:.1} GB", - gpu.utilization_percent, - gpu.memory_used_mb, - gpu.memory_total_mb, - tab.cpu_percent, - tab.ram_used_gb, - tab.ram_total_gb, - )) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Resources "), - ); - frame.render_widget(resources, chunks[1]); -} diff --git a/bin/fxt/src/commands/watch/state.rs b/bin/fxt/src/commands/watch/state.rs deleted file mode 100644 index 60ffc534e..000000000 --- a/bin/fxt/src/commands/watch/state.rs +++ /dev/null @@ -1,846 +0,0 @@ -//! Dashboard state types for the `fxt watch` streaming TUI. -//! -//! All mutable state lives here so that the render and event-loop modules -//! can borrow it without circular dependencies. - -use std::collections::{HashMap, VecDeque}; - -// --------------------------------------------------------------------------- -// Tab enum -// --------------------------------------------------------------------------- - -/// Top-level dashboard tabs. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum Tab { - Training, - Trading, - Risk, - System, -} - -impl Tab { - /// All tabs in display order. - pub const ALL: [Tab; 4] = [Tab::Training, Tab::Trading, Tab::Risk, Tab::System]; - - /// Human-readable title for the tab bar. - pub fn title(self) -> &'static str { - match self { - Tab::Training => "Training", - Tab::Trading => "Trading", - Tab::Risk => "Risk", - Tab::System => "System", - } - } - - /// Zero-based index matching [`Tab::ALL`]. - pub fn index(self) -> usize { - match self { - Tab::Training => 0, - Tab::Trading => 1, - Tab::Risk => 2, - Tab::System => 3, - } - } -} - -// --------------------------------------------------------------------------- -// Training types -// --------------------------------------------------------------------------- - -/// Returns `true` if the model name indicates an RL model (DQN or PPO). -pub fn is_rl_model(model: &str) -> bool { - let lower = model.to_ascii_lowercase(); - lower.contains("dqn") || lower.contains("ppo") -} - -/// A single in-progress training session (all 35 proto fields + local gpu_percent). -#[derive(Debug, Clone, Default)] -pub struct TrainingSession { - pub model: String, - pub fold: String, - pub is_hyperopt: bool, - // Epoch / progress - pub epoch: f32, - pub epoch_loss: f32, - pub validation_loss: f32, - // Throughput - pub batches_per_second: f32, - pub batches_processed: f32, - pub iteration_seconds: f32, - // Eval metrics - pub eval_accuracy: f32, - pub eval_precision: f32, - pub eval_recall: f32, - pub eval_f1: f32, - // Checkpoint - pub checkpoint_size_bytes: f32, - pub checkpoint_saves: u32, - pub checkpoint_failures: u32, - // Health counters - pub nan_detected: u32, - pub gradient_explosions: u32, - pub feature_errors: u32, - // Hyperopt - pub hyperopt_trial_current: u32, - pub hyperopt_trial_total: u32, - pub hyperopt_best_objective: f32, - pub hyperopt_trials_failed: u32, - // RL diagnostics - pub q_value_mean: f32, - pub q_value_max: f32, - pub policy_entropy: f32, - pub kl_divergence: f32, - pub advantage_mean: f32, - pub replay_buffer_size: u32, - // Gradient & training health - pub gradient_norm: f32, - pub learning_rate: f32, - pub epoch_duration_seconds: f32, - // Hyperopt intra-trial - pub hyperopt_trial_epoch: u32, - pub hyperopt_trial_best_loss: f32, - pub hyperopt_elapsed_seconds: f32, - // Local (not from proto) - pub gpu_percent: f32, -} - -/// Whether the Training tab shows the session list or a single-session detail. -#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] -pub enum TrainingViewMode { - #[default] - List, - Detail, -} - -/// Sub-tabs within the training detail view. -#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] -pub enum DetailSubTab { - #[default] - Overview, - Loss, - RlDiagnostics, - Metrics, - Hyperopt, - Health, -} - -impl DetailSubTab { - /// Human-readable label for the sub-tab bar. - pub fn title(self) -> &'static str { - match self { - Self::Overview => "Overview", - Self::Loss => "Loss", - Self::RlDiagnostics => "RL", - Self::Metrics => "Metrics", - Self::Hyperopt => "Hyperopt", - Self::Health => "Health", - } - } -} - -/// Returns the sub-tabs visible for a given model/session configuration. -pub fn visible_sub_tabs(model: &str, is_hyperopt: bool) -> Vec { - let mut tabs = vec![DetailSubTab::Overview, DetailSubTab::Loss]; - if is_rl_model(model) { - tabs.push(DetailSubTab::RlDiagnostics); - } else { - tabs.push(DetailSubTab::Metrics); - } - if is_hyperopt { - tabs.push(DetailSubTab::Hyperopt); - } - tabs.push(DetailSubTab::Health); - tabs -} - -/// Per-session unbounded metric history for sparklines. -#[derive(Debug, Clone, Default)] -pub struct SessionHistory { - pub loss: Vec, - pub val_loss: Vec, - pub learning_rate: Vec, - pub gradient_norm: Vec, - pub batches_per_second: Vec, - pub accuracy: Vec, - pub precision: Vec, - pub recall: Vec, - pub f1: Vec, - pub q_value_mean: Vec, - pub policy_entropy: Vec, - pub kl_divergence: Vec, - pub advantage_mean: Vec, - pub hyperopt_best_objective: Vec, -} - -impl SessionHistory { - /// Append one sample from a training session snapshot. - pub fn push(&mut self, s: &TrainingSession) { - self.loss.push(f64::from(s.epoch_loss)); - self.val_loss.push(f64::from(s.validation_loss)); - self.learning_rate.push(f64::from(s.learning_rate)); - self.gradient_norm.push(f64::from(s.gradient_norm)); - self.batches_per_second.push(f64::from(s.batches_per_second)); - self.accuracy.push(f64::from(s.eval_accuracy)); - self.precision.push(f64::from(s.eval_precision)); - self.recall.push(f64::from(s.eval_recall)); - self.f1.push(f64::from(s.eval_f1)); - self.q_value_mean.push(f64::from(s.q_value_mean)); - self.policy_entropy.push(f64::from(s.policy_entropy)); - self.kl_divergence.push(f64::from(s.kl_divergence)); - self.advantage_mean.push(f64::from(s.advantage_mean)); - self.hyperopt_best_objective.push(f64::from(s.hyperopt_best_objective)); - } -} - -/// GPU telemetry snapshot. -#[derive(Debug, Default, Clone)] -pub struct GpuInfo { - pub utilization_percent: f32, - pub memory_used_mb: f32, - pub memory_total_mb: f32, - pub temperature_celsius: f32, -} - -/// State for the **Training** tab. -#[derive(Debug, Default)] -pub struct TrainingTabState { - pub sessions: Vec, - pub gpu: GpuInfo, - /// Ring buffer of recent loss values (max 200). - pub loss_history: VecDeque, - pub scroll_offset: usize, - pub connected: bool, - // -- Detail view state -- - pub selected_index: Option, - pub view_mode: TrainingViewMode, - pub detail_sub_tab: DetailSubTab, - pub session_histories: HashMap, -} - -const MAX_LOSS_HISTORY: usize = 200; - -impl TrainingTabState { - /// Append a loss sample, evicting the oldest if the buffer is full. - pub fn push_loss(&mut self, value: f64) { - if self.loss_history.len() >= MAX_LOSS_HISTORY { - self.loss_history.pop_front(); - } - self.loss_history.push_back(value); - } - - pub fn scroll_up(&mut self) { - self.scroll_offset = self.scroll_offset.saturating_sub(1); - } - - pub fn scroll_down(&mut self) { - self.scroll_offset = self.scroll_offset.saturating_add(1); - } - - /// Move selection cursor up (list mode). - pub fn select_up(&mut self) { - if let Some(idx) = self.selected_index { - self.selected_index = Some(idx.saturating_sub(1)); - } - } - - /// Move selection cursor down (list mode). - pub fn select_down(&mut self) { - match self.selected_index { - Some(idx) => { - let max = self.sessions.len().saturating_sub(1); - if idx < max { - self.selected_index = Some(idx.saturating_add(1)); - } - } - None => { - if !self.sessions.is_empty() { - self.selected_index = Some(0); - } - } - } - } - - /// Enter detail view for the currently selected session. - pub fn enter_detail(&mut self) { - if self.selected_index.is_some() && !self.sessions.is_empty() { - self.view_mode = TrainingViewMode::Detail; - self.detail_sub_tab = DetailSubTab::Overview; - self.scroll_offset = 0; - } - } - - /// Return to list view. - pub fn exit_detail(&mut self) { - self.view_mode = TrainingViewMode::List; - self.scroll_offset = 0; - } - - /// Advance to the next visible sub-tab. - pub fn next_sub_tab(&mut self) { - let tabs = match self.selected_index.and_then(|i| self.sessions.get(i)) { - Some(s) => visible_sub_tabs(&s.model, s.is_hyperopt), - None => return, - }; - if let Some(pos) = tabs.iter().position(|t| *t == self.detail_sub_tab) { - if let Some(&next) = tabs.get(pos + 1) { - self.detail_sub_tab = next; - } - } - } - - /// Go back to the previous visible sub-tab. - pub fn prev_sub_tab(&mut self) { - let tabs = match self.selected_index.and_then(|i| self.sessions.get(i)) { - Some(s) => visible_sub_tabs(&s.model, s.is_hyperopt), - None => return, - }; - if let Some(pos) = tabs.iter().position(|t| *t == self.detail_sub_tab) { - if let Some(&prev) = pos.checked_sub(1).and_then(|p| tabs.get(p)) { - self.detail_sub_tab = prev; - } - } - } - - /// The currently selected session, if any. - pub fn selected_session(&self) -> Option<&TrainingSession> { - self.selected_index.and_then(|i| self.sessions.get(i)) - } - - /// Sub-tabs visible for the currently selected session. - pub fn visible_tabs(&self) -> Vec { - match self.selected_session() { - Some(s) => visible_sub_tabs(&s.model, s.is_hyperopt), - None => vec![DetailSubTab::Overview], - } - } - - /// Clamp selected_index after sessions list changes. - pub fn clamp_selection(&mut self) { - match (self.sessions.is_empty(), self.selected_index) { - (true, _) => self.selected_index = None, - (false, Some(idx)) if idx >= self.sessions.len() => { - self.selected_index = Some(self.sessions.len().saturating_sub(1)); - } - _ => {} - } - } -} - -// --------------------------------------------------------------------------- -// Trading types -// --------------------------------------------------------------------------- - -/// A single open position. -#[derive(Debug, Clone)] -pub struct PositionRow { - pub symbol: String, - pub side: String, - pub quantity: f64, - pub entry_price: f64, - pub unrealized_pnl: f64, - pub status: String, -} - -/// A single fill / execution. -#[derive(Debug, Clone)] -pub struct FillRow { - pub time: String, - pub side: String, - pub symbol: String, - pub quantity: f64, - pub price: f64, - pub status: String, -} - -/// State for the **Trading** tab. -#[derive(Debug, Default)] -pub struct TradingTabState { - pub positions: Vec, - /// Recent fills ring buffer (max 50). - pub fills: VecDeque, - pub day_pnl: f64, - pub realized_pnl: f64, - pub unrealized_pnl: f64, - pub scroll_offset: usize, - pub connected: bool, -} - -const MAX_FILLS: usize = 50; - -impl TradingTabState { - /// Append a fill, evicting the oldest if the buffer is full. - pub fn push_fill(&mut self, fill: FillRow) { - if self.fills.len() >= MAX_FILLS { - self.fills.pop_front(); - } - self.fills.push_back(fill); - } - - pub fn scroll_up(&mut self) { - self.scroll_offset = self.scroll_offset.saturating_sub(1); - } - - pub fn scroll_down(&mut self) { - self.scroll_offset = self.scroll_offset.saturating_add(1); - } -} - -// --------------------------------------------------------------------------- -// Risk types -// --------------------------------------------------------------------------- - -/// A single circuit breaker status row. -#[derive(Debug, Clone)] -pub struct CircuitBreakerRow { - pub name: String, - pub status: String, - pub current: String, - pub limit: String, -} - -/// State for the **Risk** tab. -#[derive(Debug, Default)] -pub struct RiskTabState { - pub portfolio_var: f64, - pub max_drawdown: f64, - pub current_drawdown: f64, - pub sharpe_ratio: f64, - pub circuit_breakers: Vec, - pub scroll_offset: usize, - pub connected: bool, -} - -impl RiskTabState { - pub fn scroll_up(&mut self) { - self.scroll_offset = self.scroll_offset.saturating_sub(1); - } - - pub fn scroll_down(&mut self) { - self.scroll_offset = self.scroll_offset.saturating_add(1); - } -} - -// --------------------------------------------------------------------------- -// System types -// --------------------------------------------------------------------------- - -/// Health status of a single micro-service. -#[derive(Debug, Clone)] -pub struct ServiceRow { - pub name: String, - pub status: String, - pub message: String, - pub latency_ms: f64, -} - -/// State for the **System** tab. -#[derive(Debug, Default)] -pub struct SystemTabState { - pub services: Vec, - pub gpu: GpuInfo, - pub cpu_percent: f64, - pub ram_used_gb: f64, - pub ram_total_gb: f64, - pub scroll_offset: usize, - pub connected: bool, -} - -impl SystemTabState { - pub fn scroll_up(&mut self) { - self.scroll_offset = self.scroll_offset.saturating_sub(1); - } - - pub fn scroll_down(&mut self) { - self.scroll_offset = self.scroll_offset.saturating_add(1); - } -} - -// --------------------------------------------------------------------------- -// Top-level dashboard state -// --------------------------------------------------------------------------- - -/// Root state object shared by the event loop and renderer. -#[derive(Debug, Default)] -pub struct DashboardState { - pub active_tab: usize, - pub training: TrainingTabState, - pub trading: TradingTabState, - pub risk: RiskTabState, - pub system: SystemTabState, - /// Set to `true` whenever any field changes so the renderer knows to - /// repaint. The render pass resets it to `false`. - pub dirty: bool, -} - -impl DashboardState { - /// Returns the [`Tab`] variant for the currently active index. - pub fn current_tab(&self) -> Tab { - *Tab::ALL - .get(self.active_tab) - .unwrap_or(&Tab::Training) - } - - /// Switch to a tab by index (clamped to valid range). - pub fn select_tab(&mut self, index: usize) { - if index < Tab::ALL.len() { - self.active_tab = index; - self.dirty = true; - } - } - - /// Delegate scroll-up to the active tab. - pub fn scroll_up(&mut self) { - match self.current_tab() { - Tab::Training => self.training.scroll_up(), - Tab::Trading => self.trading.scroll_up(), - Tab::Risk => self.risk.scroll_up(), - Tab::System => self.system.scroll_up(), - } - self.dirty = true; - } - - /// Delegate scroll-down to the active tab. - pub fn scroll_down(&mut self) { - match self.current_tab() { - Tab::Training => self.training.scroll_down(), - Tab::Trading => self.trading.scroll_down(), - Tab::Risk => self.risk.scroll_down(), - Tab::System => self.system.scroll_down(), - } - self.dirty = true; - } -} - -// --------------------------------------------------------------------------- -// Tests -// --------------------------------------------------------------------------- - -#[cfg(test)] -#[allow(clippy::unwrap_used, clippy::expect_used)] -mod tests { - use super::*; - - #[test] - fn test_tab_selection() { - let mut state = DashboardState::default(); - assert_eq!(state.current_tab(), Tab::Training); - - state.select_tab(2); - assert_eq!(state.current_tab(), Tab::Risk); - assert_eq!(state.active_tab, 2); - assert!(state.dirty); - } - - #[test] - fn test_tab_selection_out_of_bounds() { - let mut state = DashboardState::default(); - state.select_tab(99); - // Should remain at the default (0 = Training). - assert_eq!(state.active_tab, 0); - assert_eq!(state.current_tab(), Tab::Training); - } - - #[test] - fn test_loss_history_ring_buffer() { - let mut ts = TrainingTabState::default(); - for i in 0..250 { - ts.push_loss(i as f64); - } - assert_eq!(ts.loss_history.len(), 200); - // Oldest surviving entry should be 50.0 (0..49 were evicted). - assert_eq!(*ts.loss_history.front().unwrap(), 50.0); - } - - #[test] - fn test_fill_ring_buffer() { - let mut ts = TradingTabState::default(); - for i in 0..60 { - ts.push_fill(FillRow { - time: format!("t{i}"), - side: "Buy".into(), - symbol: "ES".into(), - quantity: 1.0, - price: 5000.0, - status: "Filled".into(), - }); - } - assert_eq!(ts.fills.len(), 50); - } - - #[test] - fn test_scroll() { - let mut state = DashboardState::default(); - - // Scroll down twice. - state.scroll_down(); - state.scroll_down(); - assert_eq!(state.training.scroll_offset, 2); - - // Scroll up once. - state.scroll_up(); - assert_eq!(state.training.scroll_offset, 1); - - // Scroll up past zero -- must not underflow. - state.scroll_up(); - state.scroll_up(); - assert_eq!(state.training.scroll_offset, 0); - } - - // -- is_rl_model -------------------------------------------------------- - - #[test] - fn test_is_rl_model() { - assert!(is_rl_model("DQN")); - assert!(is_rl_model("dqn")); - assert!(is_rl_model("PPO")); - assert!(is_rl_model("ContinuousPPO")); - assert!(!is_rl_model("TFT")); - assert!(!is_rl_model("Mamba2")); - assert!(!is_rl_model("XLSTM")); - } - - // -- visible_sub_tabs --------------------------------------------------- - - #[test] - fn test_visible_sub_tabs_rl() { - let tabs = visible_sub_tabs("DQN", false); - assert_eq!( - tabs, - vec![ - DetailSubTab::Overview, - DetailSubTab::Loss, - DetailSubTab::RlDiagnostics, - DetailSubTab::Health, - ] - ); - } - - #[test] - fn test_visible_sub_tabs_supervised() { - let tabs = visible_sub_tabs("TFT", false); - assert_eq!( - tabs, - vec![ - DetailSubTab::Overview, - DetailSubTab::Loss, - DetailSubTab::Metrics, - DetailSubTab::Health, - ] - ); - } - - #[test] - fn test_visible_sub_tabs_hyperopt_rl() { - let tabs = visible_sub_tabs("PPO", true); - assert_eq!( - tabs, - vec![ - DetailSubTab::Overview, - DetailSubTab::Loss, - DetailSubTab::RlDiagnostics, - DetailSubTab::Hyperopt, - DetailSubTab::Health, - ] - ); - } - - #[test] - fn test_visible_sub_tabs_hyperopt_supervised() { - let tabs = visible_sub_tabs("TFT", true); - assert_eq!( - tabs, - vec![ - DetailSubTab::Overview, - DetailSubTab::Loss, - DetailSubTab::Metrics, - DetailSubTab::Hyperopt, - DetailSubTab::Health, - ] - ); - } - - // -- SessionHistory::push ----------------------------------------------- - - #[test] - fn test_session_history_push() { - let mut hist = SessionHistory::default(); - let session = TrainingSession { - epoch_loss: 0.5, - validation_loss: 0.6, - learning_rate: 1e-3, - gradient_norm: 1.2, - batches_per_second: 42.0, - eval_accuracy: 0.85, - eval_f1: 0.82, - q_value_mean: 3.5, - policy_entropy: 1.1, - kl_divergence: 0.02, - advantage_mean: 0.1, - hyperopt_best_objective: 0.45, - ..Default::default() - }; - - hist.push(&session); - hist.push(&session); - - assert_eq!(hist.loss.len(), 2); - assert!((hist.loss[0] - 0.5).abs() < f64::EPSILON); - assert!((hist.accuracy[0] - 0.85).abs() < 1e-6); - assert!((hist.q_value_mean[0] - 3.5).abs() < 1e-6); - } - - // -- select_up / select_down -------------------------------------------- - - fn make_tab_with_sessions(n: usize) -> TrainingTabState { - let mut tab = TrainingTabState::default(); - for i in 0..n { - tab.sessions.push(TrainingSession { - model: format!("DQN_{i}"), - fold: format!("fold_{i}"), - ..Default::default() - }); - } - tab - } - - #[test] - fn test_select_down_from_none() { - let mut tab = make_tab_with_sessions(3); - assert_eq!(tab.selected_index, None); - tab.select_down(); - assert_eq!(tab.selected_index, Some(0)); - } - - #[test] - fn test_select_down_clamps_at_end() { - let mut tab = make_tab_with_sessions(3); - tab.selected_index = Some(2); - tab.select_down(); - assert_eq!(tab.selected_index, Some(2)); - } - - #[test] - fn test_select_up_clamps_at_zero() { - let mut tab = make_tab_with_sessions(3); - tab.selected_index = Some(0); - tab.select_up(); - assert_eq!(tab.selected_index, Some(0)); - } - - #[test] - fn test_select_down_empty_sessions() { - let mut tab = TrainingTabState::default(); - tab.select_down(); - assert_eq!(tab.selected_index, None); - } - - #[test] - fn test_select_up_down_navigation() { - let mut tab = make_tab_with_sessions(5); - tab.selected_index = Some(0); - tab.select_down(); - tab.select_down(); - assert_eq!(tab.selected_index, Some(2)); - tab.select_up(); - assert_eq!(tab.selected_index, Some(1)); - } - - // -- mode transitions --------------------------------------------------- - - #[test] - fn test_enter_detail_with_selection() { - let mut tab = make_tab_with_sessions(3); - tab.selected_index = Some(1); - tab.enter_detail(); - assert_eq!(tab.view_mode, TrainingViewMode::Detail); - assert_eq!(tab.detail_sub_tab, DetailSubTab::Overview); - } - - #[test] - fn test_enter_detail_without_selection_noop() { - let mut tab = make_tab_with_sessions(3); - tab.enter_detail(); - assert_eq!(tab.view_mode, TrainingViewMode::List); - } - - #[test] - fn test_exit_detail() { - let mut tab = make_tab_with_sessions(3); - tab.selected_index = Some(0); - tab.enter_detail(); - tab.scroll_offset = 5; - tab.exit_detail(); - assert_eq!(tab.view_mode, TrainingViewMode::List); - assert_eq!(tab.scroll_offset, 0); - } - - // -- sub-tab navigation ------------------------------------------------- - - #[test] - fn test_next_prev_sub_tab() { - let mut tab = make_tab_with_sessions(1); - tab.selected_index = Some(0); - tab.enter_detail(); - - // DQN -> Overview, Loss, RL, Health - assert_eq!(tab.detail_sub_tab, DetailSubTab::Overview); - tab.next_sub_tab(); - assert_eq!(tab.detail_sub_tab, DetailSubTab::Loss); - tab.next_sub_tab(); - assert_eq!(tab.detail_sub_tab, DetailSubTab::RlDiagnostics); - tab.next_sub_tab(); - assert_eq!(tab.detail_sub_tab, DetailSubTab::Health); - // At the end — should stay - tab.next_sub_tab(); - assert_eq!(tab.detail_sub_tab, DetailSubTab::Health); - - // Go back - tab.prev_sub_tab(); - assert_eq!(tab.detail_sub_tab, DetailSubTab::RlDiagnostics); - tab.prev_sub_tab(); - assert_eq!(tab.detail_sub_tab, DetailSubTab::Loss); - tab.prev_sub_tab(); - assert_eq!(tab.detail_sub_tab, DetailSubTab::Overview); - // At start — should stay - tab.prev_sub_tab(); - assert_eq!(tab.detail_sub_tab, DetailSubTab::Overview); - } - - // -- clamp_selection ---------------------------------------------------- - - #[test] - fn test_clamp_selection_empty() { - let mut tab = TrainingTabState::default(); - tab.selected_index = Some(5); - tab.clamp_selection(); - assert_eq!(tab.selected_index, None); - } - - #[test] - fn test_clamp_selection_out_of_bounds() { - let mut tab = make_tab_with_sessions(3); - tab.selected_index = Some(10); - tab.clamp_selection(); - assert_eq!(tab.selected_index, Some(2)); - } - - #[test] - fn test_clamp_selection_valid() { - let mut tab = make_tab_with_sessions(3); - tab.selected_index = Some(1); - tab.clamp_selection(); - assert_eq!(tab.selected_index, Some(1)); - } - - // -- DetailSubTab::title ------------------------------------------------ - - #[test] - fn test_detail_sub_tab_titles() { - assert_eq!(DetailSubTab::Overview.title(), "Overview"); - assert_eq!(DetailSubTab::Loss.title(), "Loss"); - assert_eq!(DetailSubTab::RlDiagnostics.title(), "RL"); - assert_eq!(DetailSubTab::Metrics.title(), "Metrics"); - assert_eq!(DetailSubTab::Hyperopt.title(), "Hyperopt"); - assert_eq!(DetailSubTab::Health.title(), "Health"); - } -} diff --git a/bin/fxt/src/commands/watch/streams.rs b/bin/fxt/src/commands/watch/streams.rs deleted file mode 100644 index f9c674bb8..000000000 --- a/bin/fxt/src/commands/watch/streams.rs +++ /dev/null @@ -1,805 +0,0 @@ -//! gRPC stream manager for `fxt watch`. -//! -//! Spawns one tokio task per server-streaming RPC plus one one-shot task for -//! initial state (positions + risk metrics). All tasks funnel [`StreamEvent`]s -//! into a single `mpsc` channel that the event loop consumes. -//! -//! Every streaming task runs an infinite retry loop with exponential backoff -//! (1 s initial, 30 s cap) so the dashboard stays alive across transient -//! network failures. - -use anyhow::{Context, Result}; -use tokio::sync::mpsc; -use tokio_stream::StreamExt; -use tonic::metadata::{Ascii, MetadataValue}; -use tonic::transport::Channel; -use tracing::{debug, warn}; - -use crate::proto::monitoring::{ - self as mon, monitoring_service_client::MonitoringServiceClient, -}; -use crate::proto::trading::{self as trd, trading_service_client::TradingServiceClient}; - -use super::state::{GpuInfo, PositionRow, TrainingSession}; - -// --------------------------------------------------------------------------- -// StreamEvent — the single enum consumed by the event loop -// --------------------------------------------------------------------------- - -/// Typed events produced by the individual gRPC stream tasks. -#[derive(Debug, Clone)] -#[allow(dead_code)] // Fields like order_id, severity, symbol reserved for future UI use -pub(super) enum StreamEvent { - /// Training telemetry snapshot (sessions + GPU). - TrainingUpdate { - sessions: Vec, - gpu: GpuInfo, - }, - - /// A single order/fill update. - OrderUpdate { - order_id: String, - symbol: String, - status: String, - filled_qty: f64, - last_fill_price: f64, - timestamp_nanos: i64, - }, - - /// Full positions snapshot (replaces the previous set). - PositionsSnapshot { positions: Vec }, - - /// Periodic risk-metric broadcast. - RiskMetrics { - var: f64, - max_drawdown: f64, - current_drawdown: f64, - sharpe: f64, - }, - - /// Risk circuit-breaker alert. - RiskAlert { - severity: String, - symbol: String, - message: String, - threshold: f64, - current: f64, - }, - - /// Service health heartbeat. - SystemStatus { - service: String, - status: String, - message: String, - }, - - /// Indicates that a named stream has disconnected. - StreamDisconnected { stream_name: String }, -} - -// --------------------------------------------------------------------------- -// Helpers -// --------------------------------------------------------------------------- - -/// Create a lazily-connected gRPC channel to the given URL. -fn connect_channel(url: &str) -> Result { - crate::client::connect_channel_lazy(url) -} - -/// Parse a JWT string into a tonic `MetadataValue` suitable for the -/// `authorization` header. -fn auth_metadata(jwt: &str) -> Result> { - let bearer = format!("Bearer {jwt}"); - bearer - .parse::>() - .context("JWT token contains invalid header characters") -} - -// --------------------------------------------------------------------------- -// Backoff helper -// --------------------------------------------------------------------------- - -/// Capped exponential backoff: doubles each call, capped at 30 s. -fn next_backoff(current: std::time::Duration) -> std::time::Duration { - let doubled = current.saturating_mul(2); - let cap = std::time::Duration::from_secs(30); - if doubled > cap { - cap - } else { - doubled - } -} - -const INITIAL_BACKOFF: std::time::Duration = std::time::Duration::from_secs(1); - -// --------------------------------------------------------------------------- -// Proto -> state conversions -// --------------------------------------------------------------------------- - -fn convert_training_session(s: &mon::TrainingSession) -> TrainingSession { - TrainingSession { - model: s.model.clone(), - fold: s.fold.clone(), - is_hyperopt: s.is_hyperopt, - // Epoch / progress - epoch: s.current_epoch, - epoch_loss: s.epoch_loss, - validation_loss: s.validation_loss, - // Throughput - batches_per_second: s.batches_per_second, - batches_processed: s.batches_processed, - iteration_seconds: s.iteration_seconds, - // Eval metrics - eval_accuracy: s.eval_accuracy, - eval_precision: s.eval_precision, - eval_recall: s.eval_recall, - eval_f1: s.eval_f1, - // Checkpoint - checkpoint_size_bytes: s.checkpoint_size_bytes, - checkpoint_saves: s.checkpoint_saves, - checkpoint_failures: s.checkpoint_failures, - // Health counters - nan_detected: s.nan_detected, - gradient_explosions: s.gradient_explosions, - feature_errors: s.feature_errors, - // Hyperopt - hyperopt_trial_current: s.hyperopt_trial_current, - hyperopt_trial_total: s.hyperopt_trial_total, - hyperopt_best_objective: s.hyperopt_best_objective, - hyperopt_trials_failed: s.hyperopt_trials_failed, - // RL diagnostics - q_value_mean: s.q_value_mean, - q_value_max: s.q_value_max, - policy_entropy: s.policy_entropy, - kl_divergence: s.kl_divergence, - advantage_mean: s.advantage_mean, - replay_buffer_size: s.replay_buffer_size, - // Gradient & training health - gradient_norm: s.gradient_norm, - learning_rate: s.learning_rate, - epoch_duration_seconds: s.epoch_duration_seconds, - // Hyperopt intra-trial - hyperopt_trial_epoch: s.hyperopt_trial_epoch, - hyperopt_trial_best_loss: s.hyperopt_trial_best_loss, - hyperopt_elapsed_seconds: s.hyperopt_elapsed_seconds, - // Local - gpu_percent: 0.0, // per-session GPU not in proto; filled from GpuSnapshot - } -} - -fn convert_gpu(g: &mon::GpuSnapshot) -> GpuInfo { - GpuInfo { - utilization_percent: g.utilization_percent, - memory_used_mb: g.memory_used_mb, - memory_total_mb: g.memory_total_mb, - temperature_celsius: g.temperature_celsius, - } -} - -fn convert_position(p: &trd::Position) -> PositionRow { - let side = if p.quantity >= 0.0 { "Long" } else { "Short" }; - PositionRow { - symbol: p.symbol.clone(), - side: side.to_owned(), - quantity: p.quantity.abs(), - entry_price: p.average_cost, - unrealized_pnl: p.unrealized_pnl, - status: "Open".to_owned(), - } -} - -// --------------------------------------------------------------------------- -// Public entry point -// --------------------------------------------------------------------------- - -/// Spawn background tasks for every dashboard stream and return the unified -/// receiver. -/// -/// Five tasks are spawned: -/// 1. Training metrics stream (monitoring service) -/// 2. Order updates stream (trading service) -/// 3. Risk alerts stream (trading service) -/// 4. System status stream (trading service) -/// 5. Initial state fetch -- one-shot unary RPCs for positions + risk metrics -/// -/// Each streaming task will reconnect automatically on transient failures and -/// emit [`StreamEvent::StreamDisconnected`] when the connection drops. -pub(super) fn spawn_all_streams( - api_gateway_url: &str, - jwt_token: &str, -) -> mpsc::Receiver { - let (tx, rx) = mpsc::channel::(256); - - // 1. Training metrics stream - { - let url = api_gateway_url.to_owned(); - let jwt = jwt_token.to_owned(); - let sender = tx.clone(); - tokio::spawn(async move { - stream_training_loop(&url, &jwt, &sender).await; - }); - } - - // 2. Order updates stream - { - let url = api_gateway_url.to_owned(); - let jwt = jwt_token.to_owned(); - let sender = tx.clone(); - tokio::spawn(async move { - stream_orders_loop(&url, &jwt, &sender).await; - }); - } - - // 3. Risk alerts stream - { - let url = api_gateway_url.to_owned(); - let jwt = jwt_token.to_owned(); - let sender = tx.clone(); - tokio::spawn(async move { - stream_risk_alerts_loop(&url, &jwt, &sender).await; - }); - } - - // 4. System status stream - { - let url = api_gateway_url.to_owned(); - let jwt = jwt_token.to_owned(); - let sender = tx.clone(); - tokio::spawn(async move { - stream_system_status_loop(&url, &jwt, &sender).await; - }); - } - - // 5. One-shot initial state - { - let url = api_gateway_url.to_owned(); - let jwt = jwt_token.to_owned(); - // Last clone — `tx` itself is dropped after the block, leaving only - // the spawned tasks holding senders. - tokio::spawn(async move { - if let Err(e) = fetch_initial_state(&url, &jwt, &tx).await { - warn!(error = %e, "failed to fetch initial state"); - } - }); - } - - rx -} - -// --------------------------------------------------------------------------- -// Training metrics stream -// --------------------------------------------------------------------------- - -async fn stream_training_loop(url: &str, jwt: &str, tx: &mpsc::Sender) -> ! { - let mut backoff = INITIAL_BACKOFF; - loop { - match try_stream_training(url, jwt, tx).await { - Ok(()) => { - // Stream ended cleanly (server closed) -- reset backoff. - backoff = INITIAL_BACKOFF; - } - Err(e) => { - warn!(error = %e, "training metrics stream error"); - } - } - drop( - tx.send(StreamEvent::StreamDisconnected { - stream_name: "training_metrics".to_owned(), - }) - .await, - ); - debug!(backoff_ms = backoff.as_millis(), "reconnecting training stream"); - tokio::time::sleep(backoff).await; - backoff = next_backoff(backoff); - } -} - -async fn try_stream_training( - url: &str, - jwt: &str, - tx: &mpsc::Sender, -) -> Result<()> { - let channel = connect_channel(url)?; - let token = auth_metadata(jwt)?; - - let mut client = MonitoringServiceClient::with_interceptor( - channel, - move |mut req: tonic::Request<()>| { - req.metadata_mut() - .insert("authorization", token.clone()); - Ok(req) - }, - ); - - let request = mon::StreamTrainingMetricsRequest { - model_filter: String::new(), - interval_seconds: 3, - }; - - let response = client - .stream_training_metrics(request) - .await - .context("StreamTrainingMetrics RPC failed")?; - - let mut stream = response.into_inner(); - - while let Some(frame) = stream.next().await { - let snapshot = frame.context("training metrics stream message error")?; - let sessions: Vec = - snapshot.sessions.iter().map(convert_training_session).collect(); - let gpu = snapshot.gpu.as_ref().map(convert_gpu).unwrap_or_default(); - - tx.send(StreamEvent::TrainingUpdate { sessions, gpu }) - .await - .context("event channel closed")?; - } - - Ok(()) -} - -// --------------------------------------------------------------------------- -// Order updates stream -// --------------------------------------------------------------------------- - -async fn stream_orders_loop(url: &str, jwt: &str, tx: &mpsc::Sender) -> ! { - let mut backoff = INITIAL_BACKOFF; - loop { - match try_stream_orders(url, jwt, tx).await { - Ok(()) => { - backoff = INITIAL_BACKOFF; - } - Err(e) => { - warn!(error = %e, "order updates stream error"); - } - } - drop( - tx.send(StreamEvent::StreamDisconnected { - stream_name: "order_updates".to_owned(), - }) - .await, - ); - debug!(backoff_ms = backoff.as_millis(), "reconnecting orders stream"); - tokio::time::sleep(backoff).await; - backoff = next_backoff(backoff); - } -} - -async fn try_stream_orders( - url: &str, - jwt: &str, - tx: &mpsc::Sender, -) -> Result<()> { - let channel = connect_channel(url)?; - let token = auth_metadata(jwt)?; - - let mut client = TradingServiceClient::with_interceptor( - channel, - move |mut req: tonic::Request<()>| { - req.metadata_mut() - .insert("authorization", token.clone()); - Ok(req) - }, - ); - - let request = trd::SubscribeOrderUpdatesRequest { account_id: None }; - - let response = client - .subscribe_order_updates(request) - .await - .context("SubscribeOrderUpdates RPC failed")?; - - let mut stream = response.into_inner(); - - while let Some(frame) = stream.next().await { - let update = frame.context("order updates stream message error")?; - tx.send(StreamEvent::OrderUpdate { - order_id: update.order_id, - symbol: update.symbol, - status: format!("{}", update.status), - filled_qty: update.filled_quantity, - last_fill_price: update.last_fill_price, - timestamp_nanos: update.timestamp_unix_nanos, - }) - .await - .context("event channel closed")?; - } - - Ok(()) -} - -// --------------------------------------------------------------------------- -// Risk alerts stream -// --------------------------------------------------------------------------- - -async fn stream_risk_alerts_loop(url: &str, jwt: &str, tx: &mpsc::Sender) -> ! { - let mut backoff = INITIAL_BACKOFF; - loop { - match try_stream_risk_alerts(url, jwt, tx).await { - Ok(()) => { - backoff = INITIAL_BACKOFF; - } - Err(e) => { - warn!(error = %e, "risk alerts stream error"); - } - } - drop( - tx.send(StreamEvent::StreamDisconnected { - stream_name: "risk_alerts".to_owned(), - }) - .await, - ); - debug!(backoff_ms = backoff.as_millis(), "reconnecting risk alerts stream"); - tokio::time::sleep(backoff).await; - backoff = next_backoff(backoff); - } -} - -async fn try_stream_risk_alerts( - url: &str, - jwt: &str, - tx: &mpsc::Sender, -) -> Result<()> { - let channel = connect_channel(url)?; - let token = auth_metadata(jwt)?; - - let mut client = TradingServiceClient::with_interceptor( - channel, - move |mut req: tonic::Request<()>| { - req.metadata_mut() - .insert("authorization", token.clone()); - Ok(req) - }, - ); - - let request = trd::SubscribeRiskAlertsRequest { - min_severity: Vec::new(), - symbols: Vec::new(), - }; - - let response = client - .subscribe_risk_alerts(request) - .await - .context("SubscribeRiskAlerts RPC failed")?; - - let mut stream = response.into_inner(); - - while let Some(frame) = stream.next().await { - let alert = frame.context("risk alerts stream message error")?; - tx.send(StreamEvent::RiskAlert { - severity: format!("{}", alert.severity), - symbol: alert.symbol, - message: alert.message, - threshold: alert.threshold_value, - current: alert.current_value, - }) - .await - .context("event channel closed")?; - } - - Ok(()) -} - -// --------------------------------------------------------------------------- -// System status stream -// --------------------------------------------------------------------------- - -async fn stream_system_status_loop(url: &str, jwt: &str, tx: &mpsc::Sender) -> ! { - let mut backoff = INITIAL_BACKOFF; - loop { - match try_stream_system_status(url, jwt, tx).await { - Ok(()) => { - backoff = INITIAL_BACKOFF; - } - Err(e) => { - warn!(error = %e, "system status stream error"); - } - } - drop( - tx.send(StreamEvent::StreamDisconnected { - stream_name: "system_status".to_owned(), - }) - .await, - ); - debug!(backoff_ms = backoff.as_millis(), "reconnecting system status stream"); - tokio::time::sleep(backoff).await; - backoff = next_backoff(backoff); - } -} - -async fn try_stream_system_status( - url: &str, - jwt: &str, - tx: &mpsc::Sender, -) -> Result<()> { - let channel = connect_channel(url)?; - let token = auth_metadata(jwt)?; - - let mut client = TradingServiceClient::with_interceptor( - channel, - move |mut req: tonic::Request<()>| { - req.metadata_mut() - .insert("authorization", token.clone()); - Ok(req) - }, - ); - - let request = trd::SubscribeSystemStatusRequest { - service_names: Vec::new(), - }; - - let response = client - .subscribe_system_status(request) - .await - .context("SubscribeSystemStatus RPC failed")?; - - let mut stream = response.into_inner(); - - while let Some(frame) = stream.next().await { - let heartbeat = frame.context("system status stream message error")?; - let status_label = match heartbeat.status { - 1 => "HEALTHY", - 2 => "DEGRADED", - 3 => "UNHEALTHY", - 4 => "CRITICAL", - _ => "UNKNOWN", - }; - tx.send(StreamEvent::SystemStatus { - service: heartbeat.service_name, - status: status_label.to_owned(), - message: heartbeat.message, - }) - .await - .context("event channel closed")?; - } - - Ok(()) -} - -// --------------------------------------------------------------------------- -// Initial state (one-shot unary RPCs) -// --------------------------------------------------------------------------- - -/// Fetches positions and risk metrics via unary RPCs to seed the dashboard -/// with initial state before any streaming updates arrive. -async fn fetch_initial_state( - url: &str, - jwt: &str, - tx: &mpsc::Sender, -) -> Result<()> { - let channel = connect_channel(url)?; - let token = auth_metadata(jwt)?; - - let mut client = TradingServiceClient::with_interceptor( - channel, - move |mut req: tonic::Request<()>| { - req.metadata_mut() - .insert("authorization", token.clone()); - Ok(req) - }, - ); - - // Fetch positions. - let positions_resp = client - .get_positions(trd::GetPositionsRequest { symbol: None }) - .await - .context("GetPositions RPC failed")?; - - let positions: Vec = positions_resp - .into_inner() - .positions - .iter() - .map(convert_position) - .collect(); - - tx.send(StreamEvent::PositionsSnapshot { positions }) - .await - .context("event channel closed")?; - - // Fetch risk metrics. - let risk_resp = client - .get_risk_metrics(trd::GetRiskMetricsRequest { - portfolio_id: None, - start_time_unix_nanos: None, - end_time_unix_nanos: None, - }) - .await - .context("GetRiskMetrics RPC failed")?; - - let rm = risk_resp.into_inner(); - tx.send(StreamEvent::RiskMetrics { - var: rm.value_at_risk, - max_drawdown: rm.max_drawdown, - current_drawdown: rm.current_drawdown, - sharpe: rm.sharpe_ratio, - }) - .await - .context("event channel closed")?; - - debug!("initial state fetched successfully"); - Ok(()) -} - -// --------------------------------------------------------------------------- -// Tests -// --------------------------------------------------------------------------- - -#[cfg(test)] -#[allow(clippy::unwrap_used, clippy::expect_used)] -mod tests { - use super::*; - - #[test] - fn test_backoff_doubles() { - let b = next_backoff(std::time::Duration::from_secs(1)); - assert_eq!(b, std::time::Duration::from_secs(2)); - } - - #[test] - fn test_backoff_caps_at_30s() { - let b = next_backoff(std::time::Duration::from_secs(20)); - assert_eq!(b, std::time::Duration::from_secs(30)); - let b2 = next_backoff(std::time::Duration::from_secs(30)); - assert_eq!(b2, std::time::Duration::from_secs(30)); - } - - #[test] - fn test_auth_metadata_valid() { - let meta = auth_metadata("eyJhbGciOiJIUzI1NiJ9.test.sig"); - assert!(meta.is_ok()); - } - - #[tokio::test] - async fn test_connect_channel_valid_url() { - let ch = connect_channel("http://localhost:50051"); - assert!(ch.is_ok()); - } - - #[test] - fn test_connect_channel_invalid_url() { - // Invalid URLs fail at the `from_shared` step, before any runtime is - // needed, so a sync test is fine. - let ch = connect_channel("not a url at all"); - assert!(ch.is_err()); - } - - #[test] - fn test_convert_gpu() { - let gpu = convert_gpu(&mon::GpuSnapshot { - utilization_percent: 85.0, - memory_used_mb: 4096.0, - memory_total_mb: 8192.0, - temperature_celsius: 72.0, - power_watts: 200.0, - }); - assert!((gpu.utilization_percent - 85.0).abs() < f32::EPSILON); - assert!((gpu.memory_used_mb - 4096.0).abs() < f32::EPSILON); - assert!((gpu.memory_total_mb - 8192.0).abs() < f32::EPSILON); - assert!((gpu.temperature_celsius - 72.0).abs() < f32::EPSILON); - } - - #[test] - fn test_convert_position_long() { - let pos = convert_position(&trd::Position { - symbol: "ES".to_owned(), - quantity: 2.0, - market_price: 5100.0, - market_value: 10200.0, - average_cost: 5000.0, - unrealized_pnl: 200.0, - realized_pnl: 0.0, - }); - assert_eq!(pos.side, "Long"); - assert!((pos.quantity - 2.0).abs() < f64::EPSILON); - assert!((pos.entry_price - 5000.0).abs() < f64::EPSILON); - assert_eq!(pos.status, "Open"); - } - - #[test] - fn test_convert_position_short() { - let pos = convert_position(&trd::Position { - symbol: "NQ".to_owned(), - quantity: -1.0, - market_price: 18000.0, - market_value: -18000.0, - average_cost: 18100.0, - unrealized_pnl: 100.0, - realized_pnl: 0.0, - }); - assert_eq!(pos.side, "Short"); - assert!((pos.quantity - 1.0).abs() < f64::EPSILON); - } - - #[test] - fn test_convert_position_zero_quantity() { - let pos = convert_position(&trd::Position { - symbol: "ZN".to_owned(), - quantity: 0.0, - market_price: 110.0, - market_value: 0.0, - average_cost: 109.0, - unrealized_pnl: 0.0, - realized_pnl: 50.0, - }); - // Zero quantity is technically "Long" (>= 0). - assert_eq!(pos.side, "Long"); - assert!((pos.quantity - 0.0).abs() < f64::EPSILON); - } - - #[test] - fn test_convert_training_session() { - let sess = convert_training_session(&mon::TrainingSession { - model: "DQN".to_owned(), - fold: "fold_0".to_owned(), - is_hyperopt: false, - current_epoch: 5.0, - epoch_loss: 0.123, - validation_loss: 0.456, - learning_rate: 0.001, - batches_per_second: 42.0, - gradient_norm: 1.5, - epoch_duration_seconds: 12.0, - ..Default::default() - }); - assert_eq!(sess.model, "DQN"); - assert_eq!(sess.fold, "fold_0"); - assert!(!sess.is_hyperopt); - assert!((sess.epoch - 5.0).abs() < f32::EPSILON); - assert!((sess.epoch_loss - 0.123).abs() < f32::EPSILON); - assert!((sess.gradient_norm - 1.5).abs() < f32::EPSILON); - } - - #[test] - fn test_stream_event_variants_are_clone() { - let event = StreamEvent::SystemStatus { - service: "gateway".into(), - status: "UP".into(), - message: "ok".into(), - }; - let _cloned = event.clone(); - } - - #[test] - fn test_stream_event_all_variants() { - // Verify all variants can be constructed and cloned. - let events = vec![ - StreamEvent::TrainingUpdate { - sessions: vec![], - gpu: GpuInfo::default(), - }, - StreamEvent::OrderUpdate { - order_id: "o1".into(), - symbol: "ES".into(), - status: "3".into(), - filled_qty: 1.0, - last_fill_price: 5000.0, - timestamp_nanos: 123_456_789, - }, - StreamEvent::PositionsSnapshot { positions: vec![] }, - StreamEvent::RiskMetrics { - var: 0.05, - max_drawdown: 0.10, - current_drawdown: 0.02, - sharpe: 1.5, - }, - StreamEvent::RiskAlert { - severity: "3".into(), - symbol: "NQ".into(), - message: "VaR breach".into(), - threshold: 0.05, - current: 0.07, - }, - StreamEvent::SystemStatus { - service: "trading".into(), - status: "1".into(), - message: "healthy".into(), - }, - StreamEvent::StreamDisconnected { - stream_name: "training_metrics".into(), - }, - ]; - for e in &events { - let _ = e.clone(); - } - assert_eq!(events.len(), 7); - } -} diff --git a/bin/fxt/src/grpc.rs b/bin/fxt/src/grpc.rs new file mode 100644 index 000000000..6fa3bf7be --- /dev/null +++ b/bin/fxt/src/grpc.rs @@ -0,0 +1,43 @@ +//! Unified gRPC client for the Foxhunt API Gateway. +//! +//! All service RPCs route through a single [`FoxhuntClient`] which holds +//! one lazily-connected [`Channel`] to the API Gateway. Individual typed +//! clients (trading, ML, risk, ...) are constructed on demand from the +//! shared channel. + +use anyhow::Result; +use tonic::transport::Channel; + +/// Unified gRPC client -- all service accessors share a single channel. +pub struct FoxhuntClient { + channel: Channel, +} + +impl FoxhuntClient { + /// Create a client that will lazily connect to the API Gateway. + /// + /// The TCP handshake is deferred until the first RPC call (`connect_lazy`). + /// + /// # Errors + /// + /// Returns an error if the URL cannot be parsed as a valid URI. + pub async fn connect(api_url: &str) -> Result { + let mut endpoint = + Channel::from_shared(api_url.to_owned()).map_err(|e| anyhow::anyhow!("{e}"))?; + + // Auto-configure TLS for https endpoints. + if api_url.starts_with("https://") { + endpoint = endpoint + .tls_config(tonic::transport::ClientTlsConfig::new().with_enabled_roots()) + .map_err(|e| anyhow::anyhow!("TLS config error: {e}"))?; + } + + let channel = endpoint.connect_lazy(); + Ok(Self { channel }) + } + + /// Get the underlying channel for creating typed service clients. + pub fn channel(&self) -> Channel { + self.channel.clone() + } +} diff --git a/bin/fxt/src/lib.rs b/bin/fxt/src/lib.rs index 8045264fe..9c638373c 100644 --- a/bin/fxt/src/lib.rs +++ b/bin/fxt/src/lib.rs @@ -1,97 +1,36 @@ #![deny(clippy::unwrap_used, clippy::expect_used)] #![cfg_attr(test, allow(clippy::unwrap_used, clippy::expect_used))] -#![allow(missing_docs)] // Internal implementation details -#![allow(missing_debug_implementations)] // Not all types need Debug -#![allow(unused_crate_dependencies)] // Dev dependencies not used in lib.rs -// FXT CLI domain lints -#![allow(clippy::module_name_repetitions)] // Module-prefixed types provide clarity -#![allow(clippy::integer_division)] // Integer division is intentional in CLI display -#![allow(clippy::cognitive_complexity)] // CLI command handling is complex -#![allow(clippy::manual_let_else)] // if-let pattern preferred in CLI error paths -#![allow(clippy::unnecessary_wraps)] // Result wrapping needed for CLI commands -#![allow(clippy::similar_names)] // CLI variables often have similar names -#![allow(clippy::missing_const_for_fn)] // Const fn not critical for CLI code -#![allow(clippy::too_many_lines)] // CLI command functions can be long -#![allow(clippy::doc_markdown)] // Technical terms in doc comments -#![allow(clippy::must_use_candidate)] // Not all functions need must_use -#![allow(clippy::missing_errors_doc)] // Internal CLI APIs don't need full error docs -#![allow(clippy::default_numeric_fallback)] // Numeric literals are contextually typed -#![allow(clippy::unused_self)] // Self parameter needed for trait consistency -#![allow(clippy::unused_async)] // Async needed for trait implementations -#![allow(clippy::too_many_arguments)] // CLI command functions need many parameters -#![allow(clippy::indexing_slicing)] // Encryption code uses validated slice indices -#![allow(clippy::wildcard_in_or_patterns)] // Wildcard patterns in command matching -#![allow(clippy::redundant_pattern_matching)] // Explicit pattern matching preferred +#![allow(missing_docs)] +#![allow(missing_debug_implementations)] +#![allow(unused_crate_dependencies)] +#![allow(clippy::module_name_repetitions)] +#![allow(clippy::cognitive_complexity)] +#![allow(clippy::unnecessary_wraps)] +#![allow(clippy::similar_names)] +#![allow(clippy::missing_const_for_fn)] +#![allow(clippy::too_many_lines)] +#![allow(clippy::doc_markdown)] +#![allow(clippy::must_use_candidate)] +#![allow(clippy::missing_errors_doc)] +#![allow(clippy::default_numeric_fallback)] +#![allow(clippy::unused_self)] +#![allow(clippy::unused_async)] +#![allow(clippy::too_many_arguments)] +#![allow(clippy::wildcard_in_or_patterns)] +#![allow(clippy::redundant_pattern_matching)] -//! TLI (Terminal Line Interface) - Client for Foxhunt HFT Trading System +//! FXT -- Foxhunt Operations Platform CLI //! -//! This module provides a comprehensive gRPC client infrastructure for connecting -//! to and monitoring core trading services including: -//! -//! ## Core Services -//! - **Trading Service**: Integrated service with all operations (trading, risk, monitoring, config, system status) -//! - **Backtesting Service**: Strategy testing, performance analysis, results management -//! -//! ## Key Features -//! - **Connection Management**: Pooling, health checks, automatic reconnection -//! - **Real-time Streaming**: Market data, order updates, system events -//! - **Error Handling**: Circuit breakers, exponential backoff, comprehensive error types -//! - **Security**: TLS support, authentication, credential management -//! - **Monitoring**: Metrics collection, performance tracking, alerting -//! - **High Availability**: Load balancing, failover, redundancy -//! -//! ## Architecture -//! ```text -//! TLI Client Suite -//! ├── Connection Manager (pooling, health checks) -//! ├── Event Stream Manager (real-time data) -//! ├── Trading Client (ALL operations: trading, risk, monitoring, config, system status) -//! └── Backtesting Client (strategy testing, performance analysis) -//! ``` -//! -//! ## Example Usage -//! ```rust,no_run -//! use fxt::prelude::*; -//! -//! #[tokio::main] -//! async fn main() -> TliResult<()> { -//! // Create client suite with both services -//! let client_suite = TliClientBuilder::new() -//! .with_service_endpoint("trading_service".to_string(), "http://localhost:50051".to_string()) -//! .with_service_endpoint("backtesting_service".to_string(), "http://localhost:50052".to_string()) -//! .with_trading_config(TradingClientConfig::default()) -//! .with_backtesting_config(BacktestingClientConfig::default()) -//! .build() -//! .await?; -//! -//! // Use trading client -//! if let Some(trading_client) = &client_suite.trading_client { -//! // Submit order with integrated risk management -//! let order_request = SubmitOrderRequest { -//! symbol: "AAPL".to_string(), -//! side: OrderSide::Buy as i32, -//! order_type: OrderType::Market as i32, -//! quantity: 100.0, -//! client_order_id: "order_123".to_string(), -//! ..Default::default() -//! }; -//! -//! let response = trading_client.submit_order(order_request).await?; -//! println!("Order submitted: {:?}", response); -//! } -//! -//! // Shutdown -//! client_suite.shutdown().await; -//! Ok(()) -//! } -//! ``` +//! Thin gRPC client that talks exclusively to the API Gateway. +//! All business logic lives in the backend services. -// Suppress false-positive unused extern crate warnings for dependencies used in modules +// Suppress false-positive unused extern crate warnings for dependencies +// that are used in sub-modules or generated code. use chrono as _; use clap as _; use colored as _; use common as _; -use dirs as _; // Used in config module for home directory +use dirs as _; use futures_util as _; use prost as _; use rust_decimal as _; @@ -99,94 +38,24 @@ use serde as _; use serde_json as _; use tabled as _; use thiserror as _; -use toml as _; // Used in config module for TOML parsing +use toml as _; use tonic as _; use tracing_subscriber as _; use uuid as _; // Core modules pub mod auth; -pub mod client; pub mod commands; -pub mod config; // Configuration file support (~/.foxhunt/config.toml) +pub mod config; pub mod error; -pub mod types; +pub mod grpc; +pub mod output; -// Prelude module for convenient imports -pub mod prelude; - -// Placeholder modules - database module removed (should only exist in services) -/// Utility functions and helpers for TLI operations -/// -/// This module contains shared utility functions used across the TLI client, -/// including formatting helpers, validation utilities, and common operations. -pub mod utils {} - -/// Constants and configuration values for TLI -/// -/// This module contains compile-time constants, default values, and configuration -/// parameters used throughout the TLI client application. -pub mod constants {} - -// Terminal UI modules now enabled - -#[cfg(test)] -pub mod tests; - -/// TLI version information +/// FXT version information. pub const VERSION: &str = env!("CARGO_PKG_VERSION"); -/// TLI build information -pub const BUILD_INFO: BuildInfo = BuildInfo { - version: VERSION, - git_hash: match option_env!("GIT_HASH") { - Some(hash) => hash, - None => "unknown", - }, - build_date: match option_env!("BUILD_DATE") { - Some(date) => date, - None => "unknown", - }, - features: &[ - // TLI features are defined by dependencies, not cargo features - "tonic-tls", - "grpc-client", - ], -}; - -/// Build information structure +/// Generated protobuf code for gRPC client interfaces. /// -/// Contains compile-time information about the TLI binary including version, -/// source control metadata, and enabled feature flags for debugging and support. -#[derive(Debug, Clone)] -pub struct BuildInfo { - /// Semantic version string (e.g., "1.0.0") - pub version: &'static str, - /// Git commit hash from build time - pub git_hash: &'static str, - /// ISO 8601 build timestamp - pub build_date: &'static str, - /// List of enabled Cargo feature flags - pub features: &'static [&'static str], -} - -impl std::fmt::Display for BuildInfo { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!( - f, - "TLI v{} ({}), built on {}, features: [{}]", - self.version, - self.git_hash, - self.build_date, - self.features.join(", ") - ) - } -} - -/// Generated protobuf code for gRPC client interfaces -/// -/// This module contains all the auto-generated protobuf types and service clients -/// used for communicating with the Foxhunt HFT system services via gRPC. /// All protos are compiled from the workspace-root `proto/` directory. pub mod proto { /// Trading service protobuf definitions (package `trading`) diff --git a/bin/fxt/src/main.rs b/bin/fxt/src/main.rs index 1c2f756b0..1656ee07b 100644 --- a/bin/fxt/src/main.rs +++ b/bin/fxt/src/main.rs @@ -1,615 +1,202 @@ #![deny(clippy::unwrap_used, clippy::expect_used)] -#![allow(clippy::cognitive_complexity)] // CLI command dispatching is inherently complex -#![allow(clippy::redundant_pattern_matching)] // Explicit pattern matching preferred +#![allow(clippy::cognitive_complexity)] +#![allow(clippy::redundant_pattern_matching)] -//! TLI (Terminal Line Interface) - Client Application for Foxhunt HFT Trading System +//! FXT -- Foxhunt Operations Platform CLI //! -//! Pure client CLI application that connects to trading services: -//! - CLI commands for trading, backtesting, ML training, and tuning -//! - gRPC client connections to Trading and Backtesting services -//! - Authentication and secure token management +//! Thin gRPC client that talks to the API Gateway for all operations: +//! trading, training, tuning, risk, services, cluster, data, and more. -use anyhow::{Context, Result}; +use anyhow::Result; use clap::{Parser, Subcommand}; -use colored::Colorize; -use jsonwebtoken::{decode, Algorithm, DecodingKey, Validation}; -use serde::{Deserialize, Serialize}; -use std::time::{SystemTime, UNIX_EPOCH}; -use fxt::auth::token_manager::FileTokenStorage; -use fxt::{ - commands::{ - agent::{execute_agent_command, AgentArgs}, - auth::{execute_auth_command, AuthCommand}, - backtest_ml::{execute_backtest_ml_command, BacktestMlArgs}, - broker::{execute_broker_command, BrokerArgs}, - model::{execute_model_command, ModelCommand}, - trade::{execute_trade_command, TradeArgs}, - train::{execute_train_command, TrainCommand}, - tune::{execute_tune_command, TuneCommand}, - }, - config::TliConfig, + +use fxt::commands::{ + agent, auth, backtest, broker, cluster, config_cmd, data, mcp, model, risk, service, trade, + train, tune, watch, }; -use tracing::Level; -use tracing_subscriber::FmtSubscriber; +use fxt::grpc::FoxhuntClient; +use fxt::output::OutputFormat; -// Suppress false-positive unused extern crate warnings for dependencies used in modules -use aes_gcm as _; -use argon2 as _; -use async_trait as _; -use base64 as _; -use chrono as _; -use clap as _; -use colored as _; -use comfy_table as _; -use common as _; -use console as _; -use dirs as _; -use futures_util as _; -use getrandom as _; -use hex as _; -#[cfg(feature = "broker-check")] -use ibapi as _; -use indicatif as _; -use keyring as _; -use owo_colors as _; -use prost as _; -use rand as _; -use rpassword as _; -use rust_decimal as _; -use serde as _; -use serde_json as _; -use sha2 as _; -use tabled as _; -use thiserror as _; -use toml as _; -use tonic as _; -use tonic_prost as _; -use uuid as _; -use zeroize as _; - -// Configuration precedence (highest to lowest): -// 1. CLI arguments (--api-gateway-url) -// 2. Environment variables (API_GATEWAY_URL) -// 3. Config file (~/.foxhunt/config.toml) -// 4. Hardcoded defaults - -/// TLI Command Line Interface +/// Foxhunt Operations Platform #[derive(Parser)] -#[clap( - name = "fxt", - version, - about = "Foxhunt Trading System Terminal Interface", - long_about = "Foxhunt Trading System Terminal Interface\n\n\ - Environment Variables:\n\ - API_GATEWAY_URL API Gateway URL (default: https://api.fxhnt.ai)\n\ - TLI_LOG_LEVEL Log level (default: info)\n\ - TLI_TOKEN_STORAGE Token storage backend (default: keyring)\n\n\ - Precedence: CLI args > Environment variables > Config file > Defaults" -)] +#[command(name = "fxt", about = "Foxhunt Operations Platform", version)] struct Cli { - /// Subcommand to execute - #[clap(subcommand)] + /// API Gateway endpoint URL + #[arg( + long = "api-url", + env = "FXT_API_URL", + default_value = "https://api.fxhnt.ai" + )] + api_url: String, + + /// Output as JSON (for CI/LLM integration) + #[arg(long, short)] + json: bool, + + #[command(subcommand)] command: Commands, - - /// API Gateway URL - #[clap( - long, - env = "API_GATEWAY_URL", - default_value = "https://api.fxhnt.ai", - help = "API Gateway URL (env: API_GATEWAY_URL)" - )] - api_gateway_url: String, - - /// Log level (trace, debug, info, warn, error) - #[clap( - long, - env = "TLI_LOG_LEVEL", - default_value = "info", - help = "Log level (env: TLI_LOG_LEVEL)" - )] - log_level: String, - - /// Token storage backend (keyring, file) - #[clap( - long, - env = "TLI_TOKEN_STORAGE", - default_value = "keyring", - help = "Token storage backend (env: TLI_TOKEN_STORAGE)" - )] - token_storage: String, } -/// TLI subcommands #[derive(Subcommand)] enum Commands { - /// Hyperparameter tuning for ML models (DQN, PPO, MAMBA-2, TFT) - #[clap( - long_about = "Start, monitor, and manage hyperparameter tuning jobs.\n\n\ - Supported models: DQN, PPO, MAMBA_2, TFT, TLOB, LIQUID\n\ - Uses Optuna for Bayesian optimization.\n\n\ - Examples:\n\ - fxt tune start --model DQN --trials 100\n\ - fxt tune status --job-id \n\ - fxt tune best --job-id \n\ - fxt tune stop --job-id " - )] - Tune { - #[clap(subcommand)] - tune_cmd: TuneCommand, - }, - - /// Training job management (list, status, stop) - #[clap( - long_about = "Manage ML model training jobs.\n\n\ - Supported operations:\n\ - - List training jobs with filtering\n\ - - Get detailed job status\n\ - - Watch real-time training progress\n\ - - Stop jobs gracefully or forcefully\n\n\ - Examples:\n\ - fxt train list\n\ - fxt train list --status RUNNING --model TFT\n\ - fxt train status train_dqn_es_20251022_1430\n\ - fxt train watch train_tft_nq_20251022_1500\n\ - fxt train stop train_ppo_es_20251022_1600" - )] - Train { - #[clap(subcommand)] - train_cmd: TrainCommand, - }, - - /// Model promotion management (list pending, approve, reject) - #[clap( - long_about = "Manage ML model promotions to production.\n\n\ - Supports operator-in-the-loop workflow for reviewing\n\ - trained models before promoting to live trading.\n\n\ - Examples:\n\ - fxt model list\n\ - fxt model approve \n\ - fxt model reject --reason \"metrics below threshold\"" - )] - Model { - #[clap(subcommand)] - model_cmd: ModelCommand, - }, - - /// Authentication and session management - #[clap(long_about = "Login, logout, and manage authentication tokens.\n\n\ - JWT tokens are stored securely in OS keyring.\n\ - Tokens auto-refresh when expiring (within 60 seconds).\n\n\ - Examples:\n\ - fxt auth login --username trader1\n\ - fxt auth status\n\ - fxt auth refresh\n\ - fxt auth logout")] - Auth { - #[clap(subcommand)] - auth_cmd: AuthCommand, - }, - - /// Trading agent operations (universe selection, asset selection, portfolio allocation) - #[clap( - long_about = "Trading agent operations for automated portfolio management.\n\n\ - Subcommands:\n\ - allocate-portfolio - Allocate capital across selected assets\n\n\ - Examples:\n\ - fxt agent allocate-portfolio --selection-id abc-123 --total-capital 100000\n\ - fxt agent allocate-portfolio --selection-id abc-123 --total-capital 100000 --strategy risk-parity" - )] - Agent { - #[command(flatten)] - agent_args: AgentArgs, - }, - - /// ML trading operations - #[clap(name = "backtest")] - Backtest { - #[command(flatten)] - backtest_args: BacktestMlArgs, - }, - - /// ML trading operations (legacy, use backtest ml instead) - #[clap(name = "trade")] - Trade { - #[command(flatten)] - trade_args: TradeArgs, - }, - - /// Broker connectivity check - #[clap( - long_about = "Validate broker connectivity.\n\n\ - Checks:\n\ - 1. Direct IB Gateway: TCP + ibapi handshake\n\ - 2. Broker Gateway gRPC: HealthCheck + SessionStatus\n\n\ - Examples:\n\ - fxt broker check\n\ - fxt broker check --host 10.0.0.5 --port 4004\n\ - fxt broker check --skip-ibkr" - )] - Broker { - #[command(flatten)] - broker_args: BrokerArgs, - }, - - /// Live streaming dashboard - Watch, -} - -/// JWT token claims structure for validation -#[derive(Debug, Serialize, Deserialize)] -struct Claims { - /// Subject (`user_id`) - sub: String, - /// Expiration time (Unix timestamp in seconds) - exp: u64, - /// Issued at (Unix timestamp in seconds) - iat: u64, - /// JWT ID - jti: String, - /// User roles - roles: Vec, - /// User permissions - permissions: Vec, -} - -/// Load and validate JWT token from OS keyring with automatic refresh -/// -/// This function retrieves the access token from secure storage, validates expiration, -/// and automatically attempts refresh if the token is expired or expiring within 60 seconds. -/// -/// # Arguments -/// * `api_gateway_url` - API Gateway URL for token refresh requests -/// -/// # Returns -/// - `Ok(String)` - Valid access token (possibly refreshed) -/// - `Err(anyhow::Error)` - Token not found, refresh failed, or invalid format -async fn load_jwt_token(api_gateway_url: &str) -> Result { - use fxt::auth::login::LoginClient; - use fxt::auth::token_manager::{AuthTokenManager, TokenStorage}; - - let storage = FileTokenStorage::new().context("Failed to initialize file token storage")?; - - // Read access token directly from keyring - match storage.get_access_token().await? { - Some(token) => { - // Validate token expiry - if let Err(_) = validate_token_expiry(&token).await { - // Token expired - attempt refresh - tracing::info!("Token expired, attempting auto-refresh"); - println!("{}", "Token expiring, refreshing...".yellow()); - - // Check if we have a refresh token - if storage.get_refresh_token().await?.is_none() { - anyhow::bail!( - "No refresh token available. Please login: {}", - "fxt auth login".bright_cyan() - ); - } - - // Store old token for comparison - let old_token = token.clone(); - - // Refresh tokens - let auth_manager = AuthTokenManager::new(storage.clone()); - let channel = fxt::client::connect_channel_lazy(api_gateway_url) - .context("Invalid API Gateway URL")?; - - let login_client = LoginClient::new(channel); - - login_client - .refresh_tokens(&auth_manager) - .await - .context("Failed to refresh tokens")?; - - // Verify new token was stored in keyring - match storage.get_access_token().await? { - Some(new_token) => { - // Verify token was actually updated - if new_token != old_token { - tracing::info!("\u{2713} New access token confirmed in keyring"); - } else { - tracing::error!( - "\u{26a0} Token refresh did not update access token in keyring" - ); - } - - // Verify refresh token is still in keyring - match storage.get_refresh_token().await? { - Some(stored_refresh) => { - tracing::info!("\u{2713} Refresh token confirmed in keyring"); - - // Verify refresh token wasn't accidentally cleared - if stored_refresh.is_empty() { - anyhow::bail!( - "Token refresh succeeded but refresh token is empty in keyring. Please login again: {}", - "fxt auth login".bright_cyan() - ) - } - }, - None => { - anyhow::bail!( - "Token refresh succeeded but refresh token not found in keyring. Please login again: {}", - "fxt auth login".bright_cyan() - ) - }, - } - - println!("{}", "\u{2713} Token refreshed successfully".green()); - Ok(new_token) - }, - None => { - anyhow::bail!( - "Token refresh succeeded but new token not found in keyring. Please login again: {}", - "fxt auth login".bright_cyan() - ) - }, - } - } else { - // Token is still valid - Ok(token) - } - }, - None => { - anyhow::bail!( - "Not authenticated. Please run: {} first", - "fxt auth login".bright_cyan() - ) - }, - } -} - -/// Validate JWT token expiration only (for refresh decision) -/// -/// Checks if the token expires within 60 seconds without full validation. -/// -/// # Arguments -/// * `token` - JWT token string to check -/// -/// # Returns -/// - `Ok(())` - Token not expired and has >60 seconds remaining -/// - `Err(anyhow::Error)` - Token expired or expiring soon -async fn validate_token_expiry(token: &str) -> Result<()> { - // Parse token WITHOUT verification (only check expiry) - let mut validation = Validation::new(Algorithm::HS256); - validation.insecure_disable_signature_validation(); - validation.validate_exp = false; - - let token_data = decode::(token, &DecodingKey::from_secret(b"dummy"), &validation) - .context("Invalid token format")?; - - let now = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs(); - - if token_data.claims.exp <= now + 60 { - anyhow::bail!("Token expired or expiring soon") - } - - Ok(()) + /// Authentication + Auth(auth::AuthCommand), + /// Order management + Trade(trade::TradeCommand), + /// ML training lifecycle + Train(train::TrainCommand), + /// Hyperparameter optimization + Tune(tune::TuneCommand), + /// Model management + Model(model::ModelCommand), + /// Trading agent control + Agent(agent::AgentCommand), + /// Backtesting + Backtest(backtest::BacktestCommand), + /// Broker connectivity + Broker(broker::BrokerCommand), + /// Data pipeline management + Data(data::DataCommand), + /// Service operations + Service(service::ServiceCommand), + /// Cluster operations + Cluster(cluster::ClusterCommand), + /// Risk management + Risk(risk::RiskCommand), + /// System configuration + Config(config_cmd::ConfigCommand), + /// TUI cockpit dashboard + Watch(watch::WatchCommand), + /// MCP server mode + Mcp(mcp::McpCommand), } #[tokio::main] async fn main() -> Result<()> { - // Load config file (~/.foxhunt/config.toml) - let config = TliConfig::load().unwrap_or_default(); - - // Parse CLI args - let mut cli = Cli::parse(); - - // Merge: CLI args override config file (precedence: CLI > Config > Default) - if cli.api_gateway_url == "https://api.fxhnt.ai" { - // Using default, check if config has override - cli.api_gateway_url = config.api_gateway_url.clone(); - } - - if cli.log_level == "info" { - // Using default, check if config has override - cli.log_level = config.log_level.clone(); - } - - // Parse log level - let log_level = match cli.log_level.to_lowercase().as_str() { - "trace" => Level::TRACE, - "debug" => Level::DEBUG, - "info" => Level::INFO, - "warn" => Level::WARN, - "error" => Level::ERROR, - _ => Level::INFO, + let cli = Cli::parse(); + let format = if cli.json { + OutputFormat::Json + } else { + OutputFormat::Human }; - // Initialize tracing with configured log level. - // For the Watch (TUI) command, redirect logs to a file so they don't - // corrupt the ratatui alternate-screen display. - let is_tui = matches!(cli.command, Commands::Watch); - - if is_tui { - let log_dir = dirs::home_dir() - .unwrap_or_default() - .join(".foxhunt"); - std::fs::create_dir_all(&log_dir).ok(); - if let Ok(file) = std::fs::File::create(log_dir.join("watch.log")) { - let subscriber = FmtSubscriber::builder() - .with_max_level(log_level) - .with_writer(std::sync::Mutex::new(file)) - .with_ansi(false) - .finish(); - tracing::subscriber::set_global_default(subscriber).ok(); - } - } else { - let subscriber = FmtSubscriber::builder().with_max_level(log_level).finish(); - if let Err(e) = tracing::subscriber::set_global_default(subscriber) { - eprintln!("Warning: Failed to set global tracing subscriber: {}", e); - } + // Commands that do not need a gRPC connection. + match &cli.command { + Commands::Watch(cmd) => return cmd.execute().await, + Commands::Mcp(cmd) => return cmd.execute().await, + Commands::Auth(_) + | Commands::Trade(_) + | Commands::Train(_) + | Commands::Tune(_) + | Commands::Model(_) + | Commands::Agent(_) + | Commands::Backtest(_) + | Commands::Broker(_) + | Commands::Data(_) + | Commands::Service(_) + | Commands::Cluster(_) + | Commands::Risk(_) + | Commands::Config(_) => {} } - // Route to subcommands - match cli.command { - Commands::Tune { tune_cmd } => { - let jwt_token = load_jwt_token(&cli.api_gateway_url).await?; - execute_tune_command(tune_cmd, &cli.api_gateway_url, &jwt_token).await - } - Commands::Train { train_cmd } => { - let jwt_token = load_jwt_token(&cli.api_gateway_url).await?; - execute_train_command(train_cmd, &cli.api_gateway_url, &jwt_token).await - } - Commands::Model { model_cmd } => { - let jwt_token = load_jwt_token(&cli.api_gateway_url).await?; - execute_model_command(model_cmd, &cli.api_gateway_url, &jwt_token).await - } - Commands::Auth { auth_cmd } => execute_auth_command(auth_cmd).await, - Commands::Agent { agent_args } => { - let jwt_token = load_jwt_token(&cli.api_gateway_url).await?; - execute_agent_command(agent_args, &cli.api_gateway_url, &jwt_token).await - } - Commands::Backtest { backtest_args } => { - execute_backtest_ml_command(backtest_args).await - } - Commands::Trade { trade_args } => { - let jwt_token = load_jwt_token(&cli.api_gateway_url).await?; - execute_trade_command(trade_args, &cli.api_gateway_url, &jwt_token).await - } - Commands::Broker { broker_args } => { - let passed = execute_broker_command(broker_args, &config).await?; - if !passed { - std::process::exit(1); - } - Ok(()) - } - Commands::Watch => { - let jwt_token = load_jwt_token(&cli.api_gateway_url).await?; - fxt::commands::watch::run(&cli.api_gateway_url, &jwt_token).await - } + // Connect to API Gateway (lazy -- no TCP until first RPC). + let client = FoxhuntClient::connect(&cli.api_url).await?; + + match &cli.command { + Commands::Auth(cmd) => cmd.execute(&client, format).await, + Commands::Trade(cmd) => cmd.execute(&client, format).await, + Commands::Train(cmd) => cmd.execute(&client, format).await, + Commands::Tune(cmd) => cmd.execute(&client, format).await, + Commands::Model(cmd) => cmd.execute(&client, format).await, + Commands::Agent(cmd) => cmd.execute(&client, format).await, + Commands::Backtest(cmd) => cmd.execute(&client, format).await, + Commands::Broker(cmd) => cmd.execute(&client, format).await, + Commands::Data(cmd) => cmd.execute(&client, format).await, + Commands::Service(cmd) => cmd.execute(&client, format).await, + Commands::Cluster(cmd) => cmd.execute(&client, format).await, + Commands::Risk(cmd) => cmd.execute(&client, format).await, + Commands::Config(cmd) => cmd.execute(&client, format).await, + // Already handled above -- unreachable. + Commands::Watch(_) | Commands::Mcp(_) => Ok(()), } } + #[cfg(test)] #[allow(clippy::unwrap_used, clippy::expect_used)] mod tests { use super::*; #[test] - fn test_main_function_exists() { - // This test ensures the main function compiles - // Terminal application testing would require mock terminal backend - assert!(true); + fn test_cli_parse_auth_status() { + let cli = Cli::parse_from(["fxt", "auth", "status"]); + assert!(!cli.json); + assert_eq!(cli.api_url, "https://api.fxhnt.ai"); } #[test] - fn test_cli_parsing_tune_command() { - // Test that Cli struct parses tune command correctly - let cli = Cli::parse_from(&[ - "fxt", - "--api-gateway-url", - "http://test.com", - "tune", - "status", - "--job-id", - "550e8400-e29b-41d4-a716-446655440000", + fn test_cli_parse_json_flag() { + let cli = Cli::parse_from(["fxt", "--json", "auth", "status"]); + assert!(cli.json); + } + + #[test] + fn test_cli_parse_api_url_override() { + let cli = Cli::parse_from(["fxt", "--api-url", "http://localhost:9090", "auth", "login"]); + assert_eq!(cli.api_url, "http://localhost:9090"); + } + + #[test] + fn test_cli_parse_trade_submit() { + let cli = Cli::parse_from([ + "fxt", "trade", "submit", "--symbol", "ES.FUT", "--side", "buy", "--qty", "1", ]); - - assert_eq!(cli.api_gateway_url, "http://test.com"); - - match cli.command { - Commands::Tune { .. } => {}, - _ => panic!("Expected Tune command"), - } + assert!(matches!(cli.command, Commands::Trade(_))); } #[test] - fn test_cli_parsing_auth_command() { - let cli = Cli::parse_from(&["fxt", "auth", "status"]); - - match cli.command { - Commands::Auth { .. } => {}, - _ => panic!("Expected Auth command"), - } + fn test_cli_parse_service_list() { + let cli = Cli::parse_from(["fxt", "service", "list"]); + assert!(matches!(cli.command, Commands::Service(_))); } #[test] - fn test_cli_default_values() { - // Test default values when no flags provided - let cli = Cli::parse_from(&["fxt", "auth", "status"]); - - assert_eq!(cli.api_gateway_url, "https://api.fxhnt.ai"); - assert_eq!(cli.log_level, "info"); - assert_eq!(cli.token_storage, "keyring"); + fn test_cli_parse_watch() { + let cli = Cli::parse_from(["fxt", "watch"]); + assert!(matches!(cli.command, Commands::Watch(_))); } #[test] - fn test_cli_custom_values() { - // Test that custom values override defaults - let cli = Cli::parse_from(&[ - "fxt", - "--api-gateway-url", - "http://custom.com:8080", - "--log-level", - "debug", - "--token-storage", - "file", - "auth", - "status", - ]); + fn test_cli_parse_all_subcommands() { + // Verify every top-level subcommand parses without error. + let commands = [ + vec!["fxt", "auth", "status"], + vec!["fxt", "trade", "positions"], + vec!["fxt", "train", "list"], + vec!["fxt", "tune", "status", "some-id"], + vec!["fxt", "model", "list"], + vec!["fxt", "agent", "status"], + vec![ + "fxt", "backtest", "run", "--strategy", "s", "--symbol", "ES", "--from", + "2025-01-01", "--to", "2025-12-31", + ], + vec!["fxt", "broker", "status"], + vec!["fxt", "data", "status"], + vec!["fxt", "service", "health"], + vec!["fxt", "cluster", "status"], + vec!["fxt", "risk", "status"], + vec!["fxt", "config", "env"], + vec!["fxt", "watch"], + vec!["fxt", "mcp"], + ]; - assert_eq!(cli.api_gateway_url, "http://custom.com:8080"); - assert_eq!(cli.log_level, "debug"); - assert_eq!(cli.token_storage, "file"); - } - - #[test] - fn test_tune_command_with_all_args() { - let cli = Cli::parse_from(&[ - "fxt", - "tune", - "start", - "--model", - "DQN", - "--trials", - "100", - "--config", - "test_config.yaml", - "--gpu", - ]); - - match cli.command { - Commands::Tune { .. } => {}, - _ => panic!("Expected Tune command"), - } - } - - #[test] - fn test_auth_login_command_parsing() { - let cli = Cli::parse_from(&["fxt", "auth", "login", "--username", "testuser"]); - - match cli.command { - Commands::Auth { .. } => {}, - _ => panic!("Expected Auth command"), - } - } - - #[test] - fn test_cli_parsing_broker_check() { - let cli = Cli::parse_from(&["fxt", "broker", "check"]); - match cli.command { - Commands::Broker { .. } => {}, - _ => panic!("Expected Broker command"), - } - } - - #[test] - fn test_cli_parsing_broker_check_with_flags() { - let cli = Cli::parse_from(&[ - "fxt", - "broker", - "check", - "--host", - "10.0.0.5", - "--port", - "4001", - "--skip-gateway", - ]); - match cli.command { - Commands::Broker { .. } => {}, - _ => panic!("Expected Broker command"), + for args in &commands { + let result = Cli::try_parse_from(args.iter()); + assert!( + result.is_ok(), + "Failed to parse: {:?} -- {}", + args, + result.err().map_or_else(String::new, |e| e.to_string()) + ); } } } diff --git a/bin/fxt/src/output.rs b/bin/fxt/src/output.rs new file mode 100644 index 000000000..de48a4063 --- /dev/null +++ b/bin/fxt/src/output.rs @@ -0,0 +1,45 @@ +//! Output format abstraction for JSON and human-readable rendering. +//! +//! Every command result type implements [`HumanReadable`] so the CLI can +//! switch between `--json` (machine-readable) and the default tabular output +//! with a single [`render`] call. + +use serde::Serialize; +use std::io::Write; + +/// Output format selected by the global `--json` flag. +#[derive(Debug, Clone, Copy)] +pub enum OutputFormat { + /// Default: pretty-printed tables/text for humans. + Human, + /// Structured JSON (for CI pipelines, LLM tool-use, jq). + Json, +} + +/// Render any serializable result in the chosen format. +/// +/// # Errors +/// +/// Returns an error if JSON serialization or stdout writing fails. +pub fn render(result: &T, format: OutputFormat) -> anyhow::Result<()> { + match format { + OutputFormat::Json => { + let json = serde_json::to_string_pretty(result)?; + let mut stdout = std::io::stdout().lock(); + writeln!(stdout, "{json}")?; + } + OutputFormat::Human => { + result.print_human(); + } + } + Ok(()) +} + +/// Trait for human-readable terminal rendering. +/// +/// Implementors print directly to stdout using `println!`, `colored`, or +/// `tabled` — whatever looks best for the data type. +pub trait HumanReadable { + /// Print a human-friendly representation to stdout. + fn print_human(&self); +} diff --git a/bin/fxt/src/prelude.rs b/bin/fxt/src/prelude.rs deleted file mode 100644 index ec878b3e3..000000000 --- a/bin/fxt/src/prelude.rs +++ /dev/null @@ -1,22 +0,0 @@ -//! Prelude module for convenient imports -//! -//! This module re-exports commonly used types and functions for easy access -//! in TLI client code and examples. - -// Error types -pub use crate::error::{TliError, TliResult}; - -// Client types -pub use crate::client::{ClientFactory, ServiceEndpoints, TliClientBuilder, TliClientSuite}; - -// Client configurations and implementations -pub use crate::client::backtesting_client::{BacktestingClient, BacktestingClientConfig}; -pub use crate::client::connection_manager::{ConnectionConfig, ConnectionManager}; -pub use crate::client::ml_training_client::{MLTrainingClient, MLTrainingClientConfig}; -pub use crate::client::trading_client::{TradingClient, TradingClientConfig}; - -// Proto types - common trading types -pub use crate::proto::trading::{ - CancelOrderRequest, GetPositionsRequest, OrderSide, OrderStatus, OrderType, Position, - SubmitOrderRequest, Trade, -}; diff --git a/bin/fxt/src/tests.rs b/bin/fxt/src/tests.rs deleted file mode 100644 index 40a263d9e..000000000 --- a/bin/fxt/src/tests.rs +++ /dev/null @@ -1,515 +0,0 @@ -//! Unit tests for TLI components -//! -//! This module contains comprehensive unit tests for all TLI functionality -//! including client connections, type conversions, error handling, and -//! configuration management. - -#![allow(dead_code)] - -// use crate::client::{TliClient, ServiceEndpoints}; // Disabled due to compilation issues -use crate::error::TliError; -use crate::types::*; -use proptest::prelude::*; -use std::time::SystemTime; - -mod client_tests { - - #[test] - fn test_tli_basic_functionality() { - // Basic functionality test since ServiceEndpoints is disabled - assert!(true); - } - - #[test] - #[ignore = "ServiceEndpoints disabled due to compilation issues"] - fn test_service_endpoints_environment_override() { - // std::env::set_var("FOXHUNT_TRADING_ENGINE_URL", "http://custom:8080"); - // std::env::set_var("FOXHUNT_RISK_MANAGEMENT_URL", "http://custom:8081"); - - // let endpoints = ServiceEndpoints::default(); - - // assert_eq!(endpoints.trading_engine, "http://custom:8080"); - // assert_eq!(endpoints.risk_management, "http://custom:8081"); - - // Clean up - // std::env::remove_var("FOXHUNT_TRADING_ENGINE_URL"); - // std::env::remove_var("FOXHUNT_RISK_MANAGEMENT_URL"); - } - - #[test] - #[ignore = "TliClient disabled due to compilation issues"] - fn test_client_creation() { - // let client = TliClient::new(); - // assert!(client.trading.is_none()); - // assert!(client.monitoring.is_none()); - // assert!(client.config.is_none()); - // assert!(client.health.is_none()); - } - - #[test] - #[ignore = "ServiceEndpoints and TliClient disabled due to compilation issues"] - fn test_client_with_custom_endpoints() { - // let endpoints = ServiceEndpoints { - // trading_engine: "http://test:1001".to_string(), - // risk_management: "http://test:1002".to_string(), - // ml_signals: "http://test:1003".to_string(), - // market_data: "http://test:1004".to_string(), - // health_check: "http://test:1005".to_string(), - // }; - - // let client = TliClient::with_endpoints(endpoints.clone()); - // assert_eq!(client.endpoints.trading_engine, endpoints.trading_engine); - // assert_eq!(client.endpoints.risk_management, endpoints.risk_management); - } - - #[test] - #[ignore = "TliClient disabled due to compilation issues"] - fn test_service_not_connected_errors() { - // let mut client = TliClient::new(); - - // assert!(matches!(client.trading(), Err(TliError::NotConnected(_)))); - // assert!(matches!(client.monitoring(), Err(TliError::NotConnected(_)))); - // assert!(matches!(client.config(), Err(TliError::NotConnected(_)))); - } -} - -mod types_tests { - use super::*; - - #[test] - fn test_timestamp_conversions() { - let now = SystemTime::now(); - let nanos = system_time_to_unix_nanos(now); - let converted = unix_nanos_to_system_time(nanos); - - // Allow for small timing differences (< 1ms) - let diff = now - .duration_since(converted) - .unwrap_or_else(|_| converted.duration_since(now).unwrap()); - assert!(diff.as_millis() < 1); - } - - #[test] - fn test_current_unix_nanos() { - let timestamp1 = current_unix_nanos(); - std::thread::sleep(std::time::Duration::from_millis(1)); - let timestamp2 = current_unix_nanos(); - - assert!(timestamp2 > timestamp1); - assert!(timestamp2 - timestamp1 > 0); - } - - #[test] - fn test_order_side_conversions() { - // Use TliOrderSide instead of core OrderSide - - assert_eq!(order_side_to_string(TliOrderSide::Buy), "BUY"); - assert_eq!(order_side_to_string(TliOrderSide::Sell), "SELL"); - - assert_eq!(string_to_order_side("BUY").unwrap(), TliOrderSide::Buy); - assert_eq!(string_to_order_side("buy").unwrap(), TliOrderSide::Buy); - assert_eq!(string_to_order_side("SELL").unwrap(), TliOrderSide::Sell); - assert_eq!(string_to_order_side("sell").unwrap(), TliOrderSide::Sell); - - string_to_order_side("INVALID").unwrap_err(); - string_to_order_side("").unwrap_err(); - } - - #[test] - fn test_order_type_conversions() { - use crate::proto::trading::OrderType; - - assert_eq!(order_type_to_string(OrderType::Market), "MARKET"); - assert_eq!(order_type_to_string(OrderType::Limit), "LIMIT"); - assert_eq!(order_type_to_string(OrderType::Stop), "STOP"); - assert_eq!(order_type_to_string(OrderType::StopLimit), "STOP_LIMIT"); - - assert_eq!(string_to_order_type("MARKET").unwrap(), OrderType::Market); - assert_eq!(string_to_order_type("LIMIT").unwrap(), OrderType::Limit); - assert_eq!(string_to_order_type("STOP").unwrap(), OrderType::Stop); - assert_eq!( - string_to_order_type("STOP_LIMIT").unwrap(), - OrderType::StopLimit - ); - - string_to_order_type("INVALID").unwrap_err(); - } - - #[test] - fn test_order_status_conversions() { - use crate::proto::trading::OrderStatus; - - assert_eq!(order_status_to_string(OrderStatus::New), "NEW"); - assert_eq!( - order_status_to_string(OrderStatus::PartiallyFilled), - "PARTIALLY_FILLED" - ); - assert_eq!(order_status_to_string(OrderStatus::Filled), "FILLED"); - assert_eq!(order_status_to_string(OrderStatus::Cancelled), "CANCELLED"); - assert_eq!(order_status_to_string(OrderStatus::Rejected), "REJECTED"); - - assert_eq!(string_to_order_status("NEW").unwrap(), OrderStatus::New); - assert_eq!( - string_to_order_status("FILLED").unwrap(), - OrderStatus::Filled - ); - assert_eq!( - string_to_order_status("CANCELLED").unwrap(), - OrderStatus::Cancelled - ); - - string_to_order_status("INVALID").unwrap_err(); - } - - #[test] - fn test_system_status_conversions() { - assert_eq!(system_status_to_string(TliSystemStatus::Healthy), "HEALTHY"); - assert_eq!(system_status_to_string(TliSystemStatus::Warning), "WARNING"); - assert_eq!( - system_status_to_string(TliSystemStatus::Degraded), - "DEGRADED" - ); - assert_eq!( - system_status_to_string(TliSystemStatus::Critical), - "CRITICAL" - ); - - assert_eq!( - string_to_system_status("HEALTHY").unwrap(), - TliSystemStatus::Healthy - ); - assert_eq!( - string_to_system_status("WARNING").unwrap(), - TliSystemStatus::Warning - ); - assert_eq!( - string_to_system_status("DEGRADED").unwrap(), - TliSystemStatus::Degraded - ); - assert_eq!( - string_to_system_status("CRITICAL").unwrap(), - TliSystemStatus::Critical - ); - - string_to_system_status("INVALID").unwrap_err(); - } - - #[test] - fn test_symbol_validation() { - // Valid symbols - validate_symbol("AAPL").unwrap(); - validate_symbol("BTC.USD").unwrap(); - validate_symbol("EUR-USD").unwrap(); - validate_symbol("SPX_500").unwrap(); - validate_symbol("A").unwrap(); - validate_symbol("123ABC").unwrap(); - - // Invalid symbols - assert!(validate_symbol("").is_err()); - assert!(validate_symbol(&"A".repeat(21)).is_err()); - assert!(validate_symbol("BTC/USD").is_err()); // slash not allowed - assert!(validate_symbol("BTC USD").is_err()); // space not allowed - assert!(validate_symbol("BTC@USD").is_err()); // special chars not allowed - } - - #[test] - fn test_quantity_validation() { - // Valid quantities - validate_quantity(1.0).unwrap(); - validate_quantity(0.0001).unwrap(); - validate_quantity(1000000.0).unwrap(); - - // Invalid quantities - assert!(validate_quantity(0.0).is_err()); - assert!(validate_quantity(-1.0).is_err()); - assert!(validate_quantity(f64::NAN).is_err()); - assert!(validate_quantity(f64::INFINITY).is_err()); - assert!(validate_quantity(f64::NEG_INFINITY).is_err()); - } - - #[test] - fn test_price_validation() { - // Valid prices - validate_price(1.0).unwrap(); - validate_price(0.01).unwrap(); - validate_price(999999.99).unwrap(); - - // Invalid prices - assert!(validate_price(0.0).is_err()); - assert!(validate_price(-1.0).is_err()); - assert!(validate_price(f64::NAN).is_err()); - assert!(validate_price(f64::INFINITY).is_err()); - assert!(validate_price(f64::NEG_INFINITY).is_err()); - } - - #[test] - fn test_create_proto_position() { - let position = create_proto_position("AAPL".to_owned(), 100.0, 150.0, 140.0); - - assert_eq!(position.symbol, "AAPL"); - assert_eq!(position.quantity, 100.0); - assert_eq!(position.market_price, 150.0); - assert_eq!(position.market_value, 15000.0); - assert_eq!(position.average_cost, 140.0); - assert_eq!(position.unrealized_pnl, 1000.0); // (150-140) * 100 - assert_eq!(position.realized_pnl, 0.0); - } - - #[test] - fn test_create_metric() { - use std::collections::HashMap; - - let labels = HashMap::from([ - ("service".to_owned(), "test".to_owned()), - ("environment".to_owned(), "dev".to_owned()), - ]); - - let metric = create_metric( - "test_metric".to_owned(), - 42.5, - "count".to_owned(), - labels.clone(), - ); - - assert_eq!(metric.name, "test_metric"); - assert_eq!(metric.value, 42.5); - assert_eq!(metric.unit, "count"); - assert_eq!(metric.labels, labels); - assert!(metric.timestamp_unix_nanos > 0); - } -} - -mod error_tests { - use super::*; - - #[test] - fn test_error_types() { - let connection_error = TliError::Connection("Connection failed".to_owned()); - let invalid_request_error = TliError::InvalidRequest("Bad request".to_owned()); - let invalid_symbol_error = TliError::InvalidSymbol("Bad symbol".to_owned()); - let not_connected_error = TliError::Connection("Not connected".to_owned()); - - // Test Display implementation - assert!(connection_error.to_string().contains("Connection failed")); - assert!(invalid_request_error.to_string().contains("Bad request")); - assert!(invalid_symbol_error.to_string().contains("Bad symbol")); - assert!(not_connected_error.to_string().contains("Not connected")); - - // Test Debug implementation - assert!(!format!("{:?}", connection_error).is_empty()); - assert!(!format!("{:?}", invalid_request_error).is_empty()); - } - - #[test] - #[ignore = "std::io::Error From conversion not implemented"] - fn test_error_from_conversions() { - // let std_error = std::io::Error::new(std::io::ErrorKind::NotFound, "File not found"); - // let tli_error: TliError = std_error.into(); - - // match tli_error { - // TliError::Connection(msg) => assert!(msg.contains("File not found")), - // _ => panic!("Expected Connection error"), - // } - } -} - -// Property-based tests -proptest! { - #[test] - fn test_timestamp_conversion_property(timestamp in 0_i64..i64::MAX/2) { - let system_time = unix_nanos_to_system_time(timestamp); - let converted = system_time_to_unix_nanos(system_time); - - // Allow for small rounding errors - prop_assert!((converted - timestamp).abs() < 1000); // Within 1 microsecond - } - - #[test] - fn test_symbol_validation_property(symbol in "[A-Za-z0-9._-]{1,20}") { - prop_assert!(validate_symbol(&symbol).is_ok()); - } - - #[test] - fn test_quantity_validation_property(quantity in 0.0001_f64..1000000.0) { - prop_assert!(validate_quantity(quantity).is_ok()); - } - - #[test] - fn test_price_validation_property(price in 0.01_f64..999999.99) { - prop_assert!(validate_price(price).is_ok()); - } - - #[test] - fn test_position_calculation_property( - quantity in -1000.0_f64..1000.0, - market_price in 0.01_f64..10000.0, - average_cost in 0.01_f64..10000.0 - ) { - let position = create_proto_position( - "TEST".to_owned(), - quantity, - market_price, - average_cost, - ); - - prop_assert_eq!(position.quantity, quantity); - prop_assert_eq!(position.market_price, market_price); - prop_assert_eq!(position.average_cost, average_cost); - prop_assert_eq!(position.market_value, quantity * market_price); - - // Use approximate comparison for floating point PnL calculation - let expected_pnl = (market_price - average_cost) * quantity; - let diff = (position.unrealized_pnl - expected_pnl).abs(); - prop_assert!(diff < 0.0001, "PnL difference {} too large", diff); - } -} - -#[cfg(test)] -mod integration_helpers { - - use std::sync::Once; - - static INIT: Once = Once::new(); - - pub(super) fn setup_test_environment() { - INIT.call_once(|| { - // Initialize logging for tests - // Simplified logging setup - let _ = tracing_subscriber::fmt().with_test_writer().try_init(); - - // Set test environment variables - std::env::set_var("RUST_LOG", "info"); - std::env::set_var("TLI_TEST_MODE", "1"); - }); - } - - pub(super) fn cleanup_test_environment() { - // Clean up any test-specific environment variables - std::env::remove_var("TLI_TEST_MODE"); - } -} - -#[cfg(test)] -mod benchmark_helpers { - use super::*; - use std::time::Instant; - - pub(super) fn measure_time(f: F) -> (R, std::time::Duration) - where - F: FnOnce() -> R, - { - let start = Instant::now(); - let result = f(); - let duration = start.elapsed(); - (result, duration) - } - - #[test] - fn test_timestamp_conversion_performance() { - let iterations = 10000; - let start = Instant::now(); - - for i in 0..iterations { - let timestamp = (i as i64) * 1_000_000_000; // Convert to nanoseconds - let system_time = unix_nanos_to_system_time(timestamp); - let _converted = system_time_to_unix_nanos(system_time); - } - - let duration = start.elapsed(); - let avg_duration = duration / iterations; - - // Should be fast - under 1 microsecond per conversion - assert!( - avg_duration.as_nanos() < 1000, - "Timestamp conversion too slow: {:?}", - avg_duration - ); - } - - #[test] - fn test_validation_performance() { - let symbols = vec!["AAPL", "MSFT", "GOOGL", "AMZN", "TSLA"]; - let iterations = 1000; - - let start = Instant::now(); - - for _ in 0..iterations { - for symbol in &symbols { - let _ = validate_symbol(symbol); - let _ = validate_quantity(100.0); - let _ = validate_price(150.0); - } - } - - let duration = start.elapsed(); - let total_validations = iterations * symbols.len() * 3; - let avg_duration = duration / total_validations as u32; - - // Should be very fast - under 100 nanoseconds per validation - assert!( - avg_duration.as_nanos() < 100, - "Validation too slow: {:?}", - avg_duration - ); - } -} - -#[cfg(test)] -mod command_handling_tests { - use crate::types::*; - - #[test] - fn test_order_command_validation() { - // Valid order parameters - validate_symbol("AAPL").unwrap(); - validate_quantity(100.0).unwrap(); - validate_price(150.0).unwrap(); - - // Invalid order parameters - assert!(validate_symbol("").is_err()); - assert!(validate_quantity(0.0).is_err()); - assert!(validate_price(-1.0).is_err()); - } - - #[test] - fn test_order_side_parsing() { - assert_eq!(string_to_order_side("BUY").unwrap(), TliOrderSide::Buy); - assert_eq!(string_to_order_side("buy").unwrap(), TliOrderSide::Buy); - assert_eq!(string_to_order_side("SELL").unwrap(), TliOrderSide::Sell); - assert_eq!(string_to_order_side("sell").unwrap(), TliOrderSide::Sell); - string_to_order_side("INVALID").unwrap_err(); - } -} - -#[cfg(test)] -mod error_display_tests { - use crate::error::TliError; - - #[test] - fn test_error_display() { - let connection_error = TliError::Connection("Connection failed".to_owned()); - let error_str = connection_error.to_string(); - assert!(error_str.contains("Connection failed")); - } - - #[test] - fn test_error_types_comprehensive() { - let errors = vec![ - TliError::Connection("conn".to_owned()), - TliError::InvalidRequest("req".to_owned()), - TliError::InvalidSymbol("sym".to_owned()), - TliError::Config("config".to_owned()), - TliError::Dashboard("dashboard".to_owned()), - TliError::BufferFull("full".to_owned()), - TliError::NotFound("not_found".to_owned()), - TliError::Other("other".to_owned()), - ]; - - for error in errors { - // All errors should have meaningful display - assert!(!error.to_string().is_empty()); - // All errors should have debug output - assert!(!format!("{:?}", error).is_empty()); - } - } -} diff --git a/bin/fxt/src/types.rs b/bin/fxt/src/types.rs deleted file mode 100644 index bd247271e..000000000 --- a/bin/fxt/src/types.rs +++ /dev/null @@ -1,582 +0,0 @@ -//! Type conversions and utilities for TLI gRPC services - -use crate::error::{TliError, TliResult}; -use crate::proto::trading::ServiceStatus; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::time::{SystemTime, UNIX_EPOCH}; -// Simplified imports to avoid core dependency issues -// Removed unused imports: rust_decimal::Decimal, common::Price, common::Quantity, common::Timestamp, common::OrderSide - -// Define basic types locally until core is available - -// Define local types for TLI use (avoiding complex core dependencies) - -/// System status enumeration for TLI services -/// -/// Represents the operational state of trading system services from healthy -/// operation to critical failures requiring immediate attention. -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub enum TliSystemStatus { - /// System is operating normally with all components functional - Healthy, - /// System is operational but showing warning indicators - Warning, - /// System is experiencing reduced functionality or performance - Degraded, - /// System is in critical state requiring immediate intervention - Critical, -} - -/// Order side enumeration for buy/sell operations -/// -/// Represents the direction of a trading order in the financial markets. -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub enum TliOrderSide { - /// Buy order - purchasing an asset - Buy, - /// Sell order - selling an asset - Sell, -} - -/// Metric data structure for monitoring and observability -/// -/// Contains time-series metric data with metadata labels for comprehensive -/// system monitoring and performance tracking. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TliMetric { - /// Name of the metric (e.g., "latency", "throughput") - pub name: String, - /// Numeric value of the metric measurement - pub value: f64, - /// Unit of measurement (e.g., "ms", "ops/sec", "bytes") - pub unit: String, - /// Key-value labels for metric categorization and filtering - pub labels: HashMap, - /// Timestamp when the metric was recorded (Unix nanoseconds) - pub timestamp_unix_nanos: i64, -} - -/// Service status information for system health monitoring -/// -/// Provides comprehensive status information about individual trading -/// system services including operational state and diagnostic data. -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TliServiceStatus { - /// Name of the service (e.g., "`trading_engine`", "`risk_manager`") - pub service_name: String, - /// Current operational status of the service - pub status: TliSystemStatus, - /// Service uptime in seconds since last restart - pub uptime_seconds: f64, - /// Last error message if the service encountered issues - pub last_error: Option, -} - -// Use protobuf types for protocol communication -use crate::proto::trading::{ - OrderStatus as ProtoOrderStatus, OrderType as ProtoOrderType, Position as ProtoPosition, -}; - -/// Convert Unix nanoseconds to `SystemTime` -/// -/// Converts a Unix timestamp in nanoseconds to a Rust `SystemTime` instance. -/// -/// Returns `UNIX_EPOCH` for negative values to handle invalid timestamps gracefully. -/// -/// # Arguments -/// * `nanos` - Unix timestamp in nanoseconds since epoch -/// -/// # Returns -/// `SystemTime` instance representing the timestamp -pub fn unix_nanos_to_system_time(nanos: i64) -> SystemTime { - if nanos < 0 { - UNIX_EPOCH - } else { - UNIX_EPOCH + std::time::Duration::from_nanos(nanos as u64) - } -} - -/// Convert `SystemTime` to Unix nanoseconds -/// -/// Converts a Rust `SystemTime` instance to Unix nanoseconds since epoch. -/// -/// Returns 0 for times before the Unix epoch. -/// -/// # Arguments -/// * `time` - `SystemTime` instance to convert -/// -/// # Returns -/// -/// Unix timestamp in nanoseconds since epoch -pub fn system_time_to_unix_nanos(time: SystemTime) -> i64 { - time.duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_nanos() as i64 -} - -/// Get current timestamp in Unix nanoseconds -/// -/// Returns the current system time as Unix nanoseconds since epoch. -/// -/// Useful for timestamping events and metrics in the trading system. -/// -/// # Returns -/// -/// Current Unix timestamp in nanoseconds -pub fn current_unix_nanos() -> i64 { - system_time_to_unix_nanos(SystemTime::now()) -} - -/// Convert `OrderSide` to string representation -/// -/// Converts a TLI order side enumeration to its standard string representation -/// used in trading protocols and APIs. -/// -/// # Arguments -/// * `side` - The order side to convert -/// -/// # Returns -/// -/// String representation ("BUY" or "SELL") -pub const fn order_side_to_string(side: TliOrderSide) -> &'static str { - match side { - TliOrderSide::Buy => "BUY", - TliOrderSide::Sell => "SELL", - } -} - -/// Convert string to `OrderSide` -/// -/// Parses a string representation of an order side into the TLI enumeration. -/// -/// Case-insensitive parsing supports both "BUY"/"SELL" and "buy"/"sell". -/// -/// # Arguments -/// * `side` - String representation of order side -/// -/// # Returns -/// `TliResult` - Parsed order side or error for invalid input -/// -/// # Errors -/// -/// Returns `TliError::InvalidRequest` for unrecognized order side strings -pub fn string_to_order_side(side: &str) -> TliResult { - match side.to_uppercase().as_str() { - "BUY" => Ok(TliOrderSide::Buy), - "SELL" => Ok(TliOrderSide::Sell), - _ => Err(TliError::InvalidRequest(format!( - "Invalid order side: {}", - side - ))), - } -} -/// Convert protobuf `OrderType` to string representation -/// -/// Converts a protobuf `OrderType` enumeration to its standard string representation -/// used in trading APIs and user interfaces. -/// -/// # Arguments -/// * `order_type` - The protobuf `OrderType` to convert -/// -/// # Returns -/// -/// String representation of the order type -pub const fn order_type_to_string(order_type: ProtoOrderType) -> &'static str { - match order_type { - ProtoOrderType::Market => "MARKET", - ProtoOrderType::Limit => "LIMIT", - ProtoOrderType::Stop => "STOP", - ProtoOrderType::StopLimit => "STOP_LIMIT", - ProtoOrderType::Unspecified => "UNSPECIFIED", - } -} - -/// Convert string to protobuf `OrderType` -/// -/// Parses a string representation into the corresponding protobuf `OrderType`. -/// -/// Used for converting user input and API requests into internal representations. -/// -/// # Arguments -/// * `order_type` - String representation of order type -/// -/// # Returns -/// `TliResult` - Parsed order type or error for invalid input -/// -/// # Errors -/// -/// Returns `TliError::InvalidRequest` for unrecognized order type strings -pub fn string_to_order_type(order_type: &str) -> TliResult { - match order_type { - "MARKET" => Ok(ProtoOrderType::Market), - "LIMIT" => Ok(ProtoOrderType::Limit), - "STOP" => Ok(ProtoOrderType::Stop), - "STOP_LIMIT" => Ok(ProtoOrderType::StopLimit), - _ => Err(TliError::InvalidRequest(format!( - "Invalid order type: {}", - order_type - ))), - } -} -/// Convert protobuf `OrderStatus` to string representation -/// -/// Converts a protobuf `OrderStatus` enumeration to its standard string representation -/// used in trading APIs and order management systems. -/// -/// # Arguments -/// * `status` - The protobuf `OrderStatus` to convert -/// -/// # Returns -/// -/// String representation of the order status -pub const fn order_status_to_string(status: ProtoOrderStatus) -> &'static str { - match status { - ProtoOrderStatus::New => "NEW", - ProtoOrderStatus::PartiallyFilled => "PARTIALLY_FILLED", - ProtoOrderStatus::Filled => "FILLED", - ProtoOrderStatus::Cancelled => "CANCELLED", - ProtoOrderStatus::Rejected => "REJECTED", - ProtoOrderStatus::PendingCancel => "PENDING_CANCEL", - ProtoOrderStatus::Unspecified => "UNSPECIFIED", - } -} - -/// Convert string to protobuf `OrderStatus` -/// -/// Parses a string representation into the corresponding protobuf `OrderStatus`. -/// -/// Used for processing order status updates from trading venues and APIs. -/// -/// # Arguments -/// * `status` - String representation of order status -/// -/// # Returns -/// `TliResult` - Parsed order status or error for invalid input -/// -/// # Errors -/// -/// Returns `TliError::InvalidRequest` for unrecognized order status strings -pub fn string_to_order_status(status: &str) -> TliResult { - match status { - "NEW" => Ok(ProtoOrderStatus::New), - "PARTIALLY_FILLED" => Ok(ProtoOrderStatus::PartiallyFilled), - "FILLED" => Ok(ProtoOrderStatus::Filled), - "CANCELLED" => Ok(ProtoOrderStatus::Cancelled), - "REJECTED" => Ok(ProtoOrderStatus::Rejected), - "PENDING_CANCEL" => Ok(ProtoOrderStatus::PendingCancel), - _ => Err(TliError::InvalidRequest(format!( - "Invalid order status: {}", - status - ))), - } -} -/// Convert TLI `SystemStatus` to string representation -/// -/// Converts a TLI `SystemStatus` enumeration to its standard string representation -/// used in system monitoring and health check APIs. -/// -/// # Arguments -/// * `status` - The TLI `SystemStatus` to convert -/// -/// # Returns -/// -/// String representation of the system status -pub const fn system_status_to_string(status: TliSystemStatus) -> &'static str { - match status { - TliSystemStatus::Healthy => "HEALTHY", - TliSystemStatus::Warning => "WARNING", - TliSystemStatus::Degraded => "DEGRADED", - TliSystemStatus::Critical => "CRITICAL", - } -} - -/// Convert string to TLI `SystemStatus` -/// -/// Parses a string representation into the corresponding TLI `SystemStatus`. -/// -/// Used for processing health check responses and monitoring system states. -/// -/// # Arguments -/// * `status` - String representation of system status -/// -/// # Returns -/// `TliResult` - Parsed system status or error for invalid input -/// -/// # Errors -/// -/// Returns `TliError::InvalidRequest` for unrecognized system status strings -pub fn string_to_system_status(status: &str) -> TliResult { - match status { - "HEALTHY" => Ok(TliSystemStatus::Healthy), - "WARNING" => Ok(TliSystemStatus::Warning), - "DEGRADED" => Ok(TliSystemStatus::Degraded), - "CRITICAL" => Ok(TliSystemStatus::Critical), - _ => Err(TliError::InvalidRequest(format!( - "Invalid system status: {}", - status - ))), - } -} -/// Validate symbol format -/// -/// Validates that a trading symbol meets the required format constraints. -/// -/// Ensures symbols are properly formatted for use in trading operations. -/// -/// # Arguments -/// * `symbol` - The trading symbol to validate -/// -/// # Returns -/// `TliResult<()>` - Ok if valid, error describing validation failure -/// -/// # Errors -/// - `TliError::InvalidSymbol` if symbol is empty, too long, or contains invalid characters -/// -/// # Validation Rules -/// - Symbol must not be empty -/// -/// - Symbol must be 20 characters or less -/// - Symbol must contain only alphanumeric characters and separators (., -, _) -pub fn validate_symbol(symbol: &str) -> TliResult<()> { - if symbol.is_empty() { - return Err(TliError::InvalidSymbol("Symbol cannot be empty".to_owned())); - } - - if symbol.len() > 20 { - return Err(TliError::InvalidSymbol( - "Symbol too long (max 20 characters)".to_owned(), - )); - } - - // Basic symbol validation - alphanumeric plus some common separators - if !symbol - .chars() - .all(|c| c.is_alphanumeric() || c == '.' || c == '-' || c == '_') - { - return Err(TliError::InvalidSymbol( - "Symbol contains invalid characters".to_owned(), - )); - } - - Ok(()) -} - -/// Validate quantity for trading operations -/// -/// Validates that a trading quantity is positive and finite. -/// -/// Prevents invalid order quantities that could cause trading errors. -/// -/// # Arguments -/// * `quantity` - The quantity to validate -/// -/// # Returns -/// `TliResult<()>` - Ok if valid, error describing validation failure -/// -/// # Errors -/// - `TliError::InvalidRequest` if quantity is not positive or not finite -pub fn validate_quantity(quantity: f64) -> TliResult<()> { - if quantity <= 0.0 { - return Err(TliError::InvalidRequest( - "Quantity must be positive".to_owned(), - )); - } - - if !quantity.is_finite() { - return Err(TliError::InvalidRequest( - "Quantity must be finite".to_owned(), - )); - } - - Ok(()) -} - -/// Validate price for trading operations -/// -/// Validates that a trading price is positive and finite. -/// -/// Prevents invalid order prices that could cause trading errors. -/// -/// # Arguments -/// * `price` - The price to validate -/// -/// # Returns -/// `TliResult<()>` - Ok if valid, error describing validation failure -/// -/// # Errors -/// - `TliError::InvalidRequest` if price is not positive or not finite -pub fn validate_price(price: f64) -> TliResult<()> { - if price <= 0.0 { - return Err(TliError::InvalidRequest( - "Price must be positive".to_owned(), - )); - } - - if !price.is_finite() { - return Err(TliError::InvalidRequest("Price must be finite".to_owned())); - } - - Ok(()) -} - -/// Create a metric with current timestamp -/// -/// Creates a new `TliMetric` instance with the current timestamp automatically applied. -/// -/// Useful for recording system metrics and performance measurements. -/// -/// # Arguments -/// * `name` - Name of the metric (e.g., "latency", "throughput") -/// -/// * `value` - Numeric value of the measurement -/// * `unit` - Unit of measurement (e.g., "ms", "ops/sec") -/// -/// * `labels` - Key-value pairs for metric categorization -/// -/// # Returns -/// `TliMetric` instance with current timestamp -pub fn create_metric( - name: String, - value: f64, - unit: String, - labels: HashMap, -) -> TliMetric { - TliMetric { - name, - value, - unit, - labels, - timestamp_unix_nanos: current_unix_nanos(), - } -} - -/// Create a protobuf position from individual fields -/// -/// Creates a `ProtoPosition` instance with calculated market value and unrealized P&L. -/// -/// Used for building position responses and portfolio summaries. -/// -/// # Arguments -/// * `symbol` - Trading symbol for the position -/// -/// * `quantity` - Number of shares/units held -/// * `market_price` - Current market price per unit -/// -/// * `average_cost` - Average cost basis per unit -/// -/// # Returns -/// `ProtoPosition` with calculated market value and unrealized P&L -pub fn create_proto_position( - symbol: String, - quantity: f64, - market_price: f64, - average_cost: f64, -) -> ProtoPosition { - let market_value = quantity * market_price; - let unrealized_pnl = market_value - (quantity * average_cost); - - ProtoPosition { - symbol, - quantity, - market_price, - market_value, - average_cost, - unrealized_pnl, - realized_pnl: 0.0, // This would come from trade history - } -} - -/// Create a service status entry -/// -/// Creates a `ServiceStatus` protobuf message with current timestamp. -/// -/// Used for health check responses and system monitoring. -/// -/// # Arguments -/// * `name` - Name of the service being reported -/// -/// * `status` - Current operational status -/// * `message` - Human-readable status message -/// -/// * `details` - Additional key-value diagnostic information -/// -/// # Returns -/// `ServiceStatus` protobuf message with current timestamp -pub fn create_service_status( - name: String, - status: TliSystemStatus, - message: String, - details: HashMap, -) -> ServiceStatus { - ServiceStatus { - name, - status: status as i32, - message, - last_check_unix_nanos: current_unix_nanos(), - details, - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_timestamp_conversion() { - let now = SystemTime::now(); - let nanos = system_time_to_unix_nanos(now); - let converted = unix_nanos_to_system_time(nanos); - - // Allow for small timing differences - let diff = now - .duration_since(converted) - .unwrap_or_else(|_| converted.duration_since(now).unwrap()); - assert!(diff.as_millis() < 1); - } - - #[test] - fn test_order_side_conversion() { - assert_eq!(order_side_to_string(TliOrderSide::Buy), "BUY"); - assert_eq!(order_side_to_string(TliOrderSide::Sell), "SELL"); - - assert_eq!(string_to_order_side("BUY").unwrap(), TliOrderSide::Buy); - assert_eq!(string_to_order_side("buy").unwrap(), TliOrderSide::Buy); - string_to_order_side("INVALID").unwrap_err(); - } - - #[test] - fn test_symbol_validation() { - validate_symbol("AAPL").unwrap(); - validate_symbol("BTC.USD").unwrap(); - validate_symbol("EUR-USD").unwrap(); - validate_symbol("SPX_500").unwrap(); - - assert!(validate_symbol("").is_err()); - assert!(validate_symbol("A".repeat(21).as_str()).is_err()); - assert!(validate_symbol("BTC/USD").is_err()); // slash not allowed - } - - #[test] - fn test_quantity_validation() { - validate_quantity(1.0).unwrap(); - validate_quantity(0.0001).unwrap(); - - assert!(validate_quantity(0.0).is_err()); - assert!(validate_quantity(-1.0).is_err()); - assert!(validate_quantity(f64::NAN).is_err()); - assert!(validate_quantity(f64::INFINITY).is_err()); - } - - #[test] - fn test_create_position() { - let position = create_proto_position("AAPL".to_owned(), 100.0, 150.0, 140.0); - - assert_eq!(position.symbol, "AAPL"); - assert_eq!(position.quantity, 100.0); - assert_eq!(position.market_price, 150.0); - assert_eq!(position.market_value, 15000.0); - assert_eq!(position.average_cost, 140.0); - assert_eq!(position.unrealized_pnl, 1000.0); // (150-140) * 100 - } -}