diff --git a/rdagent/components/backtesting/results_db.py b/rdagent/components/backtesting/results_db.py index 3c5105fd..b745c299 100644 --- a/rdagent/components/backtesting/results_db.py +++ b/rdagent/components/backtesting/results_db.py @@ -166,7 +166,7 @@ class ResultsDatabase: self.conn.commit() return c.lastrowid - def add_loop(self, loop_idx: int, success: int, fail: int, best_ic: float = None, status: str = "completed") -> int: + def add_loop(self, loop_idx: int, success: int, fail: int, best_ic: float | None = None, status: str = "completed") -> int: c = self.conn.cursor() rate = success / (success + fail) if (success + fail) > 0 else 0 c.execute("""INSERT INTO loop_results (loop_index, factors_success, factors_fail, success_rate, best_ic, status) diff --git a/rdagent/components/backtesting/vbt_backtest.py b/rdagent/components/backtesting/vbt_backtest.py index 6c125a86..a5c20dc3 100644 --- a/rdagent/components/backtesting/vbt_backtest.py +++ b/rdagent/components/backtesting/vbt_backtest.py @@ -67,7 +67,6 @@ def _cross_check_with_vbt( close: pd.Series, position: pd.Series, txn_cost: float, - manual_total_return: float, freq: str, ) -> float | None: """Run a vectorbt simulation and return its total_return for comparison.""" @@ -264,7 +263,6 @@ def backtest_signal( close=close, position=position, txn_cost=txn_cost, - manual_total_return=total_return, freq=freq, ) diff --git a/rdagent/core/utils.py b/rdagent/core/utils.py index cc35c017..0602dbb4 100644 --- a/rdagent/core/utils.py +++ b/rdagent/core/utils.py @@ -83,10 +83,24 @@ def import_class(class_path: str) -> Any: Returns ------- class of `class_path` + + Raises + ------ + ImportError + If module or class cannot be found. """ - module_path, class_name = class_path.rsplit(".", 1) - module = importlib.import_module(module_path) - return getattr(module, class_name) + try: + module_path, class_name = class_path.rsplit(".", 1) + except ValueError: + raise ImportError(f"Invalid class path: {class_path!r}") + try: + module = importlib.import_module(module_path) + except ModuleNotFoundError as e: + raise ImportError(f"Module not found: {module_path!r}") from e + try: + return getattr(module, class_name) + except AttributeError as e: + raise ImportError(f"Class not found: {class_name!r} in {module_path!r}") from e class CacheSeedGen: