diff --git a/cli/main.py b/cli/main.py index b8dd0a381..7ee214d01 100644 --- a/cli/main.py +++ b/cli/main.py @@ -1,60 +1,12 @@ -import datetime -import os import sys -import time -from functools import wraps -from pathlib import Path import typer -from rich.align import Align -from rich.live import Live -from rich.panel import Panel -from cli.announcements import display_announcements, fetch_announcements -from cli.display import ( - ANALYST_ORDER, - AnalystWallTimeTracker, - classify_message_type, - console, - create_layout, - display_complete_report, - message_buffer, - update_analyst_statuses, - update_display, - update_research_team_status, -) -from cli.prefs import load_last_run, sanitize, save_last_run -from cli.prompts import ( - ask_anthropic_effort, - ask_gemini_thinking_config, - ask_glm_region, - ask_minimax_region, - ask_openai_reasoning_effort, - ask_output_language, - ask_qwen_region, - confirm_ollama_endpoint, - detect_asset_type, - ensure_api_key, - get_ticker, - prompt_openai_compatible_url, - resolve_backend_url, - select_analysts, - select_deep_thinking_agent, - select_llm_provider, - select_research_depth, - select_shallow_thinking_agent, -) -from cli.stats_handler import StatsCallbackHandler -from tradingagents.agents.rating import is_review +from cli.display import console +from cli.run import run_analysis from tradingagents.backtest import iter_grid, run_backtest, summarize -from tradingagents.dataflows.symbols import safe_ticker_component from tradingagents.default_config import DEFAULT_CONFIG -from tradingagents.graph.analyst_execution import ( - build_analyst_execution_plan, -) -from tradingagents.graph.trading_graph import TradingAgentsGraph from tradingagents.portfolio import load_portfolio -from tradingagents.reporting import write_report_tree # prompt_toolkit's win32 output module is importable only on Windows (it asserts # the platform at import time), so gate on the platform rather than catching the @@ -75,667 +27,6 @@ app = typer.Typer( ) - - - - -def get_user_selections(): - """Ask for the run's settings, offering the previous run's answers.""" - selections = _prompt_selections(load_last_run()) - save_last_run(selections) - return selections - - -def _prompt_selections(prefs): - """Walk the selection steps. ``prefs`` prefills, the environment skips.""" - # Display ASCII art welcome message - with open(Path(__file__).parent / "static" / "welcome.txt", encoding="utf-8") as f: - welcome_ascii = f.read() - - # Create welcome box content - welcome_content = f"{welcome_ascii}\n" - welcome_content += "[bold green]TradingAgents: Multi-Agents LLM Financial Trading Framework - CLI[/bold green]\n\n" - welcome_content += "[bold]Workflow Steps:[/bold]\n" - welcome_content += "I. Analyst Team → II. Research Team → III. Trader → IV. Risk Management → V. Portfolio Management\n\n" - welcome_content += ( - "[dim]Built by [Tauric Research](https://github.com/TauricResearch)[/dim]" - ) - - # Create and center the welcome box - welcome_box = Panel( - welcome_content, - border_style="green", - padding=(1, 2), - title="Welcome to TradingAgents", - subtitle="Multi-Agents LLM Financial Trading Framework", - ) - console.print(Align.center(welcome_box)) - console.print() - console.print() # Add vertical space before announcements - - # Fetch and display announcements (silent on failure) - announcements = fetch_announcements() - display_announcements(console, announcements) - - # Create a boxed questionnaire for each step - def create_question_box(title, prompt, default=None): - box_content = f"[bold]{title}[/bold]\n" - box_content += f"[dim]{prompt}[/dim]" - if default: - box_content += f"\n[dim]Default: {default}[/dim]" - return Panel(box_content, border_style="blue", padding=(1, 2)) - - def thinking_value_or_prompt(env_var, config_key, label, box_title, box_body, prompt_fn): - """Return the env-configured reasoning/thinking value, or prompt for it. - - When ``env_var`` is set the interactive choice is skipped and the value - the env overlay placed on DEFAULT_CONFIG is used — mirroring the - env-precedence rule applied to the other selection steps. - """ - if os.environ.get(env_var): - value = DEFAULT_CONFIG[config_key] - console.print(f"[green]✓ {label} from environment:[/green] {value}") - return value - console.print(create_question_box(box_title, box_body)) - return prompt_fn() - - # Step 1: Ticker symbol - console.print( - create_question_box( - "Step 1: Ticker Symbol", - "Enter the ticker, with exchange suffix when needed (e.g. SPY, 0700.HK, BTC-USD)", - "SPY", - ) - ) - selected_ticker = get_ticker() - asset_type = detect_asset_type(selected_ticker) - # Only announce when it's not the default stock path, to avoid printing - # "stock" on every run. - if asset_type.value != "stock": - console.print( - f"[green]Detected asset type:[/green] {asset_type.value}" - ) - - # Step 2: Analysis date - default_date = datetime.datetime.now().strftime("%Y-%m-%d") - console.print( - create_question_box( - "Step 2: Analysis Date", - "Enter the analysis date (YYYY-MM-DD)", - default_date, - ) - ) - analysis_date = get_analysis_date() - - # Step 3: Output language (skipped when set via TRADINGAGENTS_OUTPUT_LANGUAGE) - if os.environ.get("TRADINGAGENTS_OUTPUT_LANGUAGE"): - output_language = DEFAULT_CONFIG["output_language"] - console.print( - f"[green]✓ Output language from environment:[/green] {output_language}" - ) - else: - console.print( - create_question_box( - "Step 3: Output Language", - "Select the language for analyst reports and final decision" - ) - ) - output_language = ask_output_language(prefs.get("output_language")) - - # Step 4: Select analysts - console.print( - create_question_box( - "Step 4: Analysts Team", "Select your LLM analyst agents for the analysis" - ) - ) - prefs = sanitize(prefs, asset_type.value) - selected_analysts = select_analysts(asset_type, prefs.get("analysts")) - console.print( - f"[green]Selected analysts:[/green] {', '.join(analyst.value for analyst in selected_analysts)}" - ) - - # Step 5: Research depth (skipped when both round counts are set via env). - # Research depth maps to the debate + risk round counts; when both are - # supplied through TRADINGAGENTS_MAX_DEBATE_ROUNDS / _MAX_RISK_ROUNDS we keep - # the run non-interactive and honor the env values (#977). - depth_from_env = bool(os.environ.get("TRADINGAGENTS_MAX_DEBATE_ROUNDS")) and bool( - os.environ.get("TRADINGAGENTS_MAX_RISK_ROUNDS") - ) - if depth_from_env: - selected_research_depth = DEFAULT_CONFIG["max_debate_rounds"] - console.print( - f"[green]✓ Research depth from environment:[/green] " - f"{DEFAULT_CONFIG['max_debate_rounds']} debate / " - f"{DEFAULT_CONFIG['max_risk_discuss_rounds']} risk rounds" - ) - else: - console.print( - create_question_box( - "Step 5: Research Depth", "Select your research depth level" - ) - ) - selected_research_depth = select_research_depth(prefs.get("research_depth")) - - # Step 6: LLM Provider (skipped when set via TRADINGAGENTS_LLM_PROVIDER). - # The backend URL comes from TRADINGAGENTS_LLM_BACKEND_URL when set, - # otherwise the provider's default endpoint — the same value the menu - # would have picked. - provider_from_env = bool(os.environ.get("TRADINGAGENTS_LLM_PROVIDER")) - if provider_from_env: - selected_llm_provider = DEFAULT_CONFIG["llm_provider"].lower() - backend_url = resolve_backend_url( - selected_llm_provider, env_url=DEFAULT_CONFIG["backend_url"] - ) - console.print(f"[green]✓ LLM provider from environment:[/green] {selected_llm_provider}") - console.print(f"[green]✓ Backend URL:[/green] {backend_url}") - # Still confirm/persist the API key so the run doesn't fail later. - ensure_api_key(selected_llm_provider) - else: - console.print( - create_question_box( - "Step 6: LLM Provider", "Select your LLM provider" - ) - ) - selected_llm_provider, backend_url = select_llm_provider(prefs.get("llm_provider")) - - # Providers with regional endpoints prompt for the region as a secondary - # step so the main dropdown stays clean (mainland China and international - # accounts cannot share API keys). - if selected_llm_provider == "qwen": - selected_llm_provider, backend_url = ask_qwen_region() - elif selected_llm_provider == "minimax": - selected_llm_provider, backend_url = ask_minimax_region() - elif selected_llm_provider == "glm": - selected_llm_provider, backend_url = ask_glm_region() - - # Honor an explicit env backend URL even when the provider was chosen - # interactively, so it isn't overwritten by the menu default (#978). - backend_url = resolve_backend_url( - selected_llm_provider, backend_url, env_url=DEFAULT_CONFIG["backend_url"] - ) - - # The generic OpenAI-compatible endpoint has no default; ask for it if - # neither the menu nor the environment supplied one. - if selected_llm_provider == "openai_compatible" and not backend_url: - remembered_url = (prefs.get("backend_url") - if prefs.get("llm_provider") == selected_llm_provider else None) - backend_url = prompt_openai_compatible_url(remembered_url) - - # For Ollama, surface the resolved endpoint (OLLAMA_BASE_URL vs default) - # before model selection so it's obvious where we're connecting. - if selected_llm_provider == "ollama": - confirm_ollama_endpoint(backend_url) - - # Confirm the provider's API key is present; prompt the user to paste - # one and persist it to .env if it's missing, so the analysis run - # doesn't fail later at the first API call. - ensure_api_key(selected_llm_provider) - - # Step 7: Thinking agents (skipped when either model is set via environment) - if os.environ.get("TRADINGAGENTS_QUICK_THINK_LLM") or os.environ.get("TRADINGAGENTS_DEEP_THINK_LLM"): - selected_shallow_thinker = DEFAULT_CONFIG["quick_think_llm"] - selected_deep_thinker = DEFAULT_CONFIG["deep_think_llm"] - console.print( - f"[green]✓ Thinking agents from environment:[/green] " - f"quick={selected_shallow_thinker}, deep={selected_deep_thinker}" - ) - else: - console.print( - create_question_box( - "Step 7: Thinking Agents", "Select your thinking agents for analysis" - ) - ) - remembered = prefs if prefs.get("llm_provider") == selected_llm_provider else {} - selected_shallow_thinker = select_shallow_thinking_agent( - selected_llm_provider, remembered.get("quick_think_llm") - ) - selected_deep_thinker = select_deep_thinking_agent( - selected_llm_provider, remembered.get("deep_think_llm") - ) - - # Step 8: Provider-specific reasoning/thinking configuration. Each knob is - # settable via its TRADINGAGENTS_* env var; when that var is set (or the - # provider itself came from env) the prompt is skipped and the configured - # value is used — same env-precedence rule as the steps above. None = each - # provider's own default. - thinking_level = None - reasoning_effort = None - anthropic_effort = None - - provider_lower = selected_llm_provider.lower() - if provider_from_env: - thinking_level = DEFAULT_CONFIG["google_thinking_level"] - reasoning_effort = DEFAULT_CONFIG["openai_reasoning_effort"] - anthropic_effort = DEFAULT_CONFIG["anthropic_effort"] - elif provider_lower == "google": - thinking_level = thinking_value_or_prompt( - "TRADINGAGENTS_GOOGLE_THINKING_LEVEL", "google_thinking_level", - "Gemini thinking mode", "Step 8: Thinking Mode", - "Configure Gemini thinking mode", ask_gemini_thinking_config, - ) - elif provider_lower == "openai": - reasoning_effort = thinking_value_or_prompt( - "TRADINGAGENTS_OPENAI_REASONING_EFFORT", "openai_reasoning_effort", - "Reasoning effort", "Step 8: Reasoning Effort", - "Configure OpenAI reasoning effort level", ask_openai_reasoning_effort, - ) - elif provider_lower == "anthropic": - anthropic_effort = thinking_value_or_prompt( - "TRADINGAGENTS_ANTHROPIC_EFFORT", "anthropic_effort", - "Claude effort", "Step 8: Effort Level", - "Configure Claude effort level", ask_anthropic_effort, - ) - - return { - "ticker": selected_ticker, - "asset_type": asset_type.value, - "analysis_date": analysis_date, - "analysts": selected_analysts, - "research_depth": selected_research_depth, - "llm_provider": selected_llm_provider.lower(), - "backend_url": backend_url, - "quick_think_llm": selected_shallow_thinker, - "deep_think_llm": selected_deep_thinker, - "google_thinking_level": thinking_level, - "openai_reasoning_effort": reasoning_effort, - "anthropic_effort": anthropic_effort, - "output_language": output_language, - } - - -def get_analysis_date(): - """Get the analysis date from user input.""" - while True: - date_str = typer.prompt( - "", default=datetime.datetime.now().strftime("%Y-%m-%d") - ) - try: - # Validate date format and ensure it's not in the future - analysis_date = datetime.datetime.strptime(date_str, "%Y-%m-%d") - if analysis_date.date() > datetime.datetime.now().date(): - console.print("[red]Error: Analysis date cannot be in the future[/red]") - continue - return date_str - except ValueError: - console.print( - "[red]Error: Invalid date format. Please use YYYY-MM-DD[/red]" - ) - - - - - - -def _run_directory(config: dict, ticker: str, trade_date: str) -> Path: - """Where this run writes, with the ticker validated as a path component. - - Every other path that interpolates a ticker checks it first; a value of - ".." here would place the run outside the results directory. - """ - return Path(config["results_dir"]) / safe_ticker_component(ticker) / trade_date - - -def _announce_checkpoint_state(graph, ticker: str, trade_date: str) -> None: - """Say whether this run resumed a saved one, where the user can see it. - - The graph logs this, but nothing in the CLI configures logging and the live - view owns the screen, so a resume was invisible. - """ - if getattr(graph, "_resuming", False): - message_buffer.add_message( - "System", f"Resuming the saved run for {ticker} on {trade_date}" - ) - else: - message_buffer.add_message("System", f"Starting fresh for {ticker} on {trade_date}") - - -def _build_run_config(selections: dict, checkpoint: bool | None) -> dict: - """Assemble the run config from interactive selections, honoring env precedence. - - Round counts and checkpoint follow "explicit env/flag wins": an env-applied - value on DEFAULT_CONFIG is preserved unless the user overrode it on the CLI. - """ - config = DEFAULT_CONFIG.copy() - # Research depth sets both round counts, but an explicit env override - # (TRADINGAGENTS_MAX_DEBATE_ROUNDS / _MAX_RISK_ROUNDS) wins over the - # interactive selection — leave the env-applied value in place (#977). - for env_var, key in (("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "max_debate_rounds"), - ("TRADINGAGENTS_MAX_RISK_ROUNDS", "max_risk_discuss_rounds")): - if os.environ.get(env_var): - # The depth prompt still appeared (it is skipped only when both are - # set), so say which half of the answer the environment overrode. - console.print( - f"[green]✓ {key} from environment:[/green] {config[key]} " - f"(set by {env_var}, so the research depth you chose does not apply to it)" - ) - else: - config[key] = selections["research_depth"] - config["quick_think_llm"] = selections["quick_think_llm"] - config["deep_think_llm"] = selections["deep_think_llm"] - config["backend_url"] = selections["backend_url"] - config["llm_provider"] = selections["llm_provider"].lower() - # Provider-specific thinking configuration - config["google_thinking_level"] = selections.get("google_thinking_level") - config["openai_reasoning_effort"] = selections.get("openai_reasoning_effort") - config["anthropic_effort"] = selections.get("anthropic_effort") - config["output_language"] = selections.get("output_language", "English") - # --checkpoint/--no-checkpoint overrides only when explicitly given; omitting - # the flag preserves TRADINGAGENTS_CHECKPOINT_ENABLED / the default (#976). - if checkpoint is not None: - config["checkpoint_enabled"] = checkpoint - return config - - -def run_analysis(checkpoint: bool | None = None, portfolio=None): - # First get all user selections - selections = get_user_selections() - - config = _build_run_config(selections, checkpoint) - - # Create stats callback handler for tracking LLM/tool calls - stats_handler = StatsCallbackHandler() - - # Normalize analyst selection to predefined order (selection is a 'set', order is fixed) - selected_set = {analyst.value for analyst in selections["analysts"]} - selected_analyst_keys = [a for a in ANALYST_ORDER if a in selected_set] - analyst_execution_plan = build_analyst_execution_plan(selected_analyst_keys) - analyst_wall_time_tracker = AnalystWallTimeTracker(analyst_execution_plan) - - # Initialize the graph with callbacks bound to LLMs - graph = TradingAgentsGraph( - selected_analyst_keys, - config=config, - debug=True, - callbacks=[stats_handler], - ) - - # Initialize message buffer with selected analysts - message_buffer.init_for_analysis(selected_analyst_keys) - - # Track start time for elapsed display - start_time = time.time() - - # Create result directory - results_dir = _run_directory(config, selections["ticker"], selections["analysis_date"]) - results_dir.mkdir(parents=True, exist_ok=True) - report_dir = results_dir / "reports" - report_dir.mkdir(parents=True, exist_ok=True) - log_file = results_dir / "message_tool.log" - log_file.touch(exist_ok=True) - - def save_message_decorator(obj, func_name): - func = getattr(obj, func_name) - @wraps(func) - def wrapper(*args, **kwargs): - func(*args, **kwargs) - timestamp, message_type, content = obj.messages[-1] - content = content.replace("\n", " ") # Replace newlines with spaces - with open(log_file, "a", encoding="utf-8") as f: - f.write(f"{timestamp} [{message_type}] {content}\n") - return wrapper - - def save_tool_call_decorator(obj, func_name): - func = getattr(obj, func_name) - @wraps(func) - def wrapper(*args, **kwargs): - func(*args, **kwargs) - timestamp, tool_name, args = obj.tool_calls[-1] - args_str = ", ".join(f"{k}={v}" for k, v in args.items()) - with open(log_file, "a", encoding="utf-8") as f: - f.write(f"{timestamp} [Tool Call] {tool_name}({args_str})\n") - return wrapper - - def save_report_section_decorator(obj, func_name): - func = getattr(obj, func_name) - @wraps(func) - def wrapper(section_name, content): - func(section_name, content) - if section_name in obj.report_sections and obj.report_sections[section_name] is not None: - content = obj.report_sections[section_name] - if content: - file_name = f"{section_name}.md" - text = "\n".join(str(item) for item in content) if isinstance(content, list) else content - with open(report_dir / file_name, "w", encoding="utf-8") as f: - f.write(text) - return wrapper - - message_buffer.add_message = save_message_decorator(message_buffer, "add_message") - message_buffer.add_tool_call = save_tool_call_decorator(message_buffer, "add_tool_call") - message_buffer.update_report_section = save_report_section_decorator(message_buffer, "update_report_section") - - # Now start the display layout - layout = create_layout() - - # The alternate screen keeps a layout taller than the window from redrawing - # by scrolling; the final report prints after this block, on the normal screen. - with Live(layout, refresh_per_second=4, screen=True): - # Initial display - update_display(layout, stats_handler=stats_handler, start_time=start_time) - - # Add initial messages - message_buffer.add_message("System", f"Selected ticker: {selections['ticker']}") - if selections["asset_type"] != "stock": - message_buffer.add_message("System", f"Detected asset type: {selections['asset_type']}") - message_buffer.add_message( - "System", f"Analysis date: {selections['analysis_date']}" - ) - message_buffer.add_message( - "System", - f"Selected analysts: {', '.join(analyst.value for analyst in selections['analysts'])}", - ) - update_display(layout, stats_handler=stats_handler, start_time=start_time) - - # Update agent status to in_progress for the first analyst - first_analyst = analyst_execution_plan.specs[0].agent_node - message_buffer.update_agent_status(first_analyst, "in_progress") - analyst_wall_time_tracker.mark_started(selected_analyst_keys[0]) - update_display(layout, stats_handler=stats_handler, start_time=start_time) - - # Create spinner text - spinner_text = ( - f"Analyzing {selections['ticker']} on {selections['analysis_date']}..." - ) - update_display(layout, spinner_text, stats_handler=stats_handler, start_time=start_time) - - # The same initial state propagate() builds: settled decision log, past - # context and resolved instrument identity. - init_agent_state = graph.create_run_state( - selections["ticker"], selections["analysis_date"], selections["asset_type"], portfolio - ) - # Pass callbacks to graph config for tool execution tracking - # (LLM tracking is handled separately via LLM constructor) - args = graph.propagator.get_graph_args(callbacks=[stats_handler]) - - # Recompile with a checkpointer and inject the thread_id so --checkpoint - # actually saves and resumes on the CLI path (#1249); a no-op when - # checkpointing is disabled. Torn down in the finally below. - checkpoint_tid = graph.begin_checkpoint( - selections["ticker"], selections["analysis_date"], selections["asset_type"], portfolio - ) - if checkpoint_tid is not None: - args.setdefault("config", {}).setdefault("configurable", {})["thread_id"] = checkpoint_tid - _announce_checkpoint_state(graph, selections["ticker"], selections["analysis_date"]) - - # Stream the analysis. On resume, feed None so LangGraph continues the - # interrupted run instead of re-appending the initial state (#1249); the - # try/finally tears the checkpointer down even if the stream raises. - trace = [] - try: - for chunk in graph.graph.stream(graph.checkpoint_input(init_agent_state), **args): - # Process all messages in chunk, deduplicating by message ID - for message in chunk.get("messages", []): - msg_id = getattr(message, "id", None) - if msg_id is not None: - if msg_id in message_buffer._processed_message_ids: - continue - message_buffer._processed_message_ids.add(msg_id) - - msg_type, content = classify_message_type(message) - if content and content.strip(): - message_buffer.add_message(msg_type, content) - - if hasattr(message, "tool_calls") and message.tool_calls: - for tool_call in message.tool_calls: - if isinstance(tool_call, dict): - message_buffer.add_tool_call(tool_call["name"], tool_call["args"]) - else: - message_buffer.add_tool_call(tool_call.name, tool_call.args) - - # Update analyst statuses based on report state (runs on every chunk) - update_analyst_statuses( - message_buffer, - chunk, - wall_time_tracker=analyst_wall_time_tracker, - ) - - # Research Team - Handle Investment Debate State - if chunk.get("investment_debate_state"): - debate_state = chunk["investment_debate_state"] - bull_hist = debate_state.get("bull_history", "").strip() - bear_hist = debate_state.get("bear_history", "").strip() - judge = debate_state.get("judge_decision", "").strip() - - # Only update status when there's actual content - if bull_hist or bear_hist: - update_research_team_status("in_progress") - if bull_hist: - message_buffer.update_report_section( - "investment_plan", f"### Bull Researcher Analysis\n{bull_hist}" - ) - if bear_hist: - message_buffer.update_report_section( - "investment_plan", f"### Bear Researcher Analysis\n{bear_hist}" - ) - if judge: - message_buffer.update_report_section( - "investment_plan", f"### Research Manager Decision\n{judge}" - ) - update_research_team_status("completed") - message_buffer.update_agent_status("Trader", "in_progress") - - # Trading Team - if chunk.get("trader_investment_plan"): - message_buffer.update_report_section( - "trader_investment_plan", chunk["trader_investment_plan"] - ) - if message_buffer.agent_status.get("Trader") != "completed": - message_buffer.update_agent_status("Trader", "completed") - message_buffer.update_agent_status("Aggressive Analyst", "in_progress") - - # Risk Management Team - Handle Risk Debate State - if chunk.get("risk_debate_state"): - risk_state = chunk["risk_debate_state"] - agg_hist = risk_state.get("aggressive_history", "").strip() - con_hist = risk_state.get("conservative_history", "").strip() - neu_hist = risk_state.get("neutral_history", "").strip() - judge = risk_state.get("judge_decision", "").strip() - - if agg_hist: - if message_buffer.agent_status.get("Aggressive Analyst") != "completed": - message_buffer.update_agent_status("Aggressive Analyst", "in_progress") - message_buffer.update_report_section( - "final_trade_decision", f"### Aggressive Analyst Analysis\n{agg_hist}" - ) - if con_hist: - if message_buffer.agent_status.get("Conservative Analyst") != "completed": - message_buffer.update_agent_status("Conservative Analyst", "in_progress") - message_buffer.update_report_section( - "final_trade_decision", f"### Conservative Analyst Analysis\n{con_hist}" - ) - if neu_hist: - if message_buffer.agent_status.get("Neutral Analyst") != "completed": - message_buffer.update_agent_status("Neutral Analyst", "in_progress") - message_buffer.update_report_section( - "final_trade_decision", f"### Neutral Analyst Analysis\n{neu_hist}" - ) - if judge and message_buffer.agent_status.get("Portfolio Manager") != "completed": - message_buffer.update_agent_status("Portfolio Manager", "in_progress") - message_buffer.update_report_section( - "final_trade_decision", f"### Portfolio Manager Decision\n{judge}" - ) - message_buffer.update_agent_status("Aggressive Analyst", "completed") - message_buffer.update_agent_status("Conservative Analyst", "completed") - message_buffer.update_agent_status("Neutral Analyst", "completed") - message_buffer.update_agent_status("Portfolio Manager", "completed") - - # Update the display - update_display(layout, stats_handler=stats_handler, start_time=start_time) - - trace.append(chunk) - - # Streamed chunks are per-node deltas, not full state. Merge them - # so every report field populated across the run is present. - final_state = {} - for chunk in trace: - final_state.update(chunk) - - # Clean run: log the decision, then drop this run's checkpoint so a - # later run starts fresh. A mid-stream failure skips both, keeping - # the checkpoint for resume. - graph.record_decision(selections["ticker"], selections["analysis_date"], final_state) - graph.clear_checkpoint_on_success( - selections["ticker"], selections["analysis_date"], selections["asset_type"], portfolio - ) - finally: - # Always restore the plain uncheckpointed graph, even on failure. - graph.end_checkpoint() - - # Update all agent statuses to completed - for agent in message_buffer.agent_status: - message_buffer.update_agent_status(agent, "completed") - - message_buffer.add_message( - "System", f"Completed analysis for {selections['analysis_date']}" - ) - message_buffer.add_message("System", analyst_wall_time_tracker.format_summary()) - - # Update final report sections - for section in message_buffer.report_sections: - if section in final_state: - message_buffer.update_report_section(section, final_state[section]) - - update_display(layout, stats_handler=stats_handler, start_time=start_time) - - # Post-analysis prompts (outside Live context for clean interaction) - console.print("\n[bold cyan]Analysis Complete![/bold cyan]\n") - - # A decision nobody can read is not a position. Say so here rather than - # leaving the run to look like a normal result. - if is_review(graph.process_signal(final_state.get("final_trade_decision", ""))): - console.print( - "[yellow]No rating could be read from the final decision, so this run " - "is recorded for review rather than as a position. Re-run, or read the " - "decision text below and judge it yourself.[/yellow]\n" - ) - console.print(f"[dim]{analyst_wall_time_tracker.format_summary()}[/dim]") - - # Prompt to save report - save_choice = typer.prompt("Save report?", default="Y").strip().upper() - if save_choice in ("Y", "YES", ""): - timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") - # Under results_dir, not the working directory: in Docker the working - # directory is inside the container and the report goes with it, while - # results_dir is the mounted volume the rest of the run already writes to. - default_path = (Path(config["results_dir"]) / "reports" - / f"{safe_ticker_component(selections['ticker'])}_{timestamp}") - save_path_str = typer.prompt( - "Save path (press Enter for default)", - default=str(default_path) - ).strip() - save_path = Path(save_path_str) - try: - report_file = write_report_tree(final_state, selections["ticker"], save_path) - console.print(f"\n[green]✓ Report saved to:[/green] {save_path.resolve()}") - console.print(f" [dim]Complete report:[/dim] {report_file.name}") - except Exception as e: - console.print(f"[red]Error saving report: {e}[/red]") - - # Prompt to display full report - display_choice = typer.prompt("\nDisplay full report on screen?", default="Y").strip().upper() - if display_choice in ("Y", "YES", ""): - display_complete_report(final_state) - - @app.callback(invoke_without_command=True) def analyze( ctx: typer.Context, @@ -766,7 +57,6 @@ def analyze( console.print(f"[yellow]Cleared {n} checkpoint(s).[/yellow]") portfolio_context = None if portfolio: - from tradingagents.portfolio import load_portfolio try: portfolio_context = load_portfolio(portfolio) except ValueError as exc: diff --git a/cli/run.py b/cli/run.py new file mode 100644 index 000000000..3018d125e --- /dev/null +++ b/cli/run.py @@ -0,0 +1,406 @@ +"""Running one analysis from the CLI: build the graph, stream it into the live view, save the report.""" + +import datetime +import os +import time +from functools import wraps +from pathlib import Path + +import typer +from rich.live import Live + +from cli.display import ( + ANALYST_ORDER, + AnalystWallTimeTracker, + classify_message_type, + console, + create_layout, + display_complete_report, + message_buffer, + update_analyst_statuses, + update_display, + update_research_team_status, +) +from cli.selections import get_user_selections +from cli.stats_handler import StatsCallbackHandler +from tradingagents.agents.rating import is_review +from tradingagents.dataflows.symbols import safe_ticker_component +from tradingagents.default_config import DEFAULT_CONFIG +from tradingagents.graph.analyst_execution import ( + build_analyst_execution_plan, +) +from tradingagents.graph.trading_graph import TradingAgentsGraph +from tradingagents.reporting import write_report_tree + + +def _run_directory(config: dict, ticker: str, trade_date: str) -> Path: + """Where this run writes, with the ticker validated as a path component. + + Every other path that interpolates a ticker checks it first; a value of + ".." here would place the run outside the results directory. + """ + return Path(config["results_dir"]) / safe_ticker_component(ticker) / trade_date + + +def _announce_checkpoint_state(graph, ticker: str, trade_date: str) -> None: + """Say whether this run resumed a saved one, where the user can see it. + + The graph logs this, but nothing in the CLI configures logging and the live + view owns the screen, so a resume was invisible. + """ + if getattr(graph, "_resuming", False): + message_buffer.add_message( + "System", f"Resuming the saved run for {ticker} on {trade_date}" + ) + else: + message_buffer.add_message("System", f"Starting fresh for {ticker} on {trade_date}") + + +def _build_run_config(selections: dict, checkpoint: bool | None) -> dict: + """Assemble the run config from interactive selections, honoring env precedence. + + Round counts and checkpoint follow "explicit env/flag wins": an env-applied + value on DEFAULT_CONFIG is preserved unless the user overrode it on the CLI. + """ + config = DEFAULT_CONFIG.copy() + # Research depth sets both round counts, but an explicit env override + # (TRADINGAGENTS_MAX_DEBATE_ROUNDS / _MAX_RISK_ROUNDS) wins over the + # interactive selection — leave the env-applied value in place (#977). + for env_var, key in (("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "max_debate_rounds"), + ("TRADINGAGENTS_MAX_RISK_ROUNDS", "max_risk_discuss_rounds")): + if os.environ.get(env_var): + # The depth prompt still appeared (it is skipped only when both are + # set), so say which half of the answer the environment overrode. + console.print( + f"[green]✓ {key} from environment:[/green] {config[key]} " + f"(set by {env_var}, so the research depth you chose does not apply to it)" + ) + else: + config[key] = selections["research_depth"] + config["quick_think_llm"] = selections["quick_think_llm"] + config["deep_think_llm"] = selections["deep_think_llm"] + config["backend_url"] = selections["backend_url"] + config["llm_provider"] = selections["llm_provider"].lower() + # Provider-specific thinking configuration + config["google_thinking_level"] = selections.get("google_thinking_level") + config["openai_reasoning_effort"] = selections.get("openai_reasoning_effort") + config["anthropic_effort"] = selections.get("anthropic_effort") + config["output_language"] = selections.get("output_language", "English") + # --checkpoint/--no-checkpoint overrides only when explicitly given; omitting + # the flag preserves TRADINGAGENTS_CHECKPOINT_ENABLED / the default (#976). + if checkpoint is not None: + config["checkpoint_enabled"] = checkpoint + return config + + +def run_analysis(checkpoint: bool | None = None, portfolio=None): + # First get all user selections + selections = get_user_selections() + + config = _build_run_config(selections, checkpoint) + + # Create stats callback handler for tracking LLM/tool calls + stats_handler = StatsCallbackHandler() + + # Normalize analyst selection to predefined order (selection is a 'set', order is fixed) + selected_set = {analyst.value for analyst in selections["analysts"]} + selected_analyst_keys = [a for a in ANALYST_ORDER if a in selected_set] + analyst_execution_plan = build_analyst_execution_plan(selected_analyst_keys) + analyst_wall_time_tracker = AnalystWallTimeTracker(analyst_execution_plan) + + # Initialize the graph with callbacks bound to LLMs + graph = TradingAgentsGraph( + selected_analyst_keys, + config=config, + debug=True, + callbacks=[stats_handler], + ) + + # Initialize message buffer with selected analysts + message_buffer.init_for_analysis(selected_analyst_keys) + + # Track start time for elapsed display + start_time = time.time() + + # Create result directory + results_dir = _run_directory(config, selections["ticker"], selections["analysis_date"]) + results_dir.mkdir(parents=True, exist_ok=True) + report_dir = results_dir / "reports" + report_dir.mkdir(parents=True, exist_ok=True) + log_file = results_dir / "message_tool.log" + log_file.touch(exist_ok=True) + + def save_message_decorator(obj, func_name): + func = getattr(obj, func_name) + + @wraps(func) + def wrapper(*args, **kwargs): + func(*args, **kwargs) + timestamp, message_type, content = obj.messages[-1] + content = content.replace("\n", " ") # Replace newlines with spaces + with open(log_file, "a", encoding="utf-8") as f: + f.write(f"{timestamp} [{message_type}] {content}\n") + return wrapper + + def save_tool_call_decorator(obj, func_name): + func = getattr(obj, func_name) + + @wraps(func) + def wrapper(*args, **kwargs): + func(*args, **kwargs) + timestamp, tool_name, args = obj.tool_calls[-1] + args_str = ", ".join(f"{k}={v}" for k, v in args.items()) + with open(log_file, "a", encoding="utf-8") as f: + f.write(f"{timestamp} [Tool Call] {tool_name}({args_str})\n") + return wrapper + + def save_report_section_decorator(obj, func_name): + func = getattr(obj, func_name) + + @wraps(func) + def wrapper(section_name, content): + func(section_name, content) + if section_name in obj.report_sections and obj.report_sections[section_name] is not None: + content = obj.report_sections[section_name] + if content: + file_name = f"{section_name}.md" + text = "\n".join(str(item) for item in content) if isinstance(content, list) else content + with open(report_dir / file_name, "w", encoding="utf-8") as f: + f.write(text) + return wrapper + + message_buffer.add_message = save_message_decorator(message_buffer, "add_message") + message_buffer.add_tool_call = save_tool_call_decorator(message_buffer, "add_tool_call") + message_buffer.update_report_section = save_report_section_decorator(message_buffer, "update_report_section") + + # Now start the display layout + layout = create_layout() + + # The alternate screen keeps a layout taller than the window from redrawing + # by scrolling; the final report prints after this block, on the normal screen. + with Live(layout, refresh_per_second=4, screen=True): + # Initial display + update_display(layout, stats_handler=stats_handler, start_time=start_time) + + # Add initial messages + message_buffer.add_message("System", f"Selected ticker: {selections['ticker']}") + if selections["asset_type"] != "stock": + message_buffer.add_message("System", f"Detected asset type: {selections['asset_type']}") + message_buffer.add_message( + "System", f"Analysis date: {selections['analysis_date']}" + ) + message_buffer.add_message( + "System", + f"Selected analysts: {', '.join(analyst.value for analyst in selections['analysts'])}", + ) + update_display(layout, stats_handler=stats_handler, start_time=start_time) + + # Update agent status to in_progress for the first analyst + first_analyst = analyst_execution_plan.specs[0].agent_node + message_buffer.update_agent_status(first_analyst, "in_progress") + analyst_wall_time_tracker.mark_started(selected_analyst_keys[0]) + update_display(layout, stats_handler=stats_handler, start_time=start_time) + + # Create spinner text + spinner_text = ( + f"Analyzing {selections['ticker']} on {selections['analysis_date']}..." + ) + update_display(layout, spinner_text, stats_handler=stats_handler, start_time=start_time) + + # The same initial state propagate() builds: settled decision log, past + # context and resolved instrument identity. + init_agent_state = graph.create_run_state( + selections["ticker"], selections["analysis_date"], selections["asset_type"], portfolio + ) + # Pass callbacks to graph config for tool execution tracking + # (LLM tracking is handled separately via LLM constructor) + args = graph.propagator.get_graph_args(callbacks=[stats_handler]) + + # Recompile with a checkpointer and inject the thread_id so --checkpoint + # actually saves and resumes on the CLI path (#1249); a no-op when + # checkpointing is disabled. Torn down in the finally below. + checkpoint_tid = graph.begin_checkpoint( + selections["ticker"], selections["analysis_date"], selections["asset_type"], portfolio + ) + if checkpoint_tid is not None: + args.setdefault("config", {}).setdefault("configurable", {})["thread_id"] = checkpoint_tid + _announce_checkpoint_state(graph, selections["ticker"], selections["analysis_date"]) + + # Stream the analysis. On resume, feed None so LangGraph continues the + # interrupted run instead of re-appending the initial state (#1249); the + # try/finally tears the checkpointer down even if the stream raises. + trace = [] + try: + for chunk in graph.graph.stream(graph.checkpoint_input(init_agent_state), **args): + # Process all messages in chunk, deduplicating by message ID + for message in chunk.get("messages", []): + msg_id = getattr(message, "id", None) + if msg_id is not None: + if msg_id in message_buffer._processed_message_ids: + continue + message_buffer._processed_message_ids.add(msg_id) + + msg_type, content = classify_message_type(message) + if content and content.strip(): + message_buffer.add_message(msg_type, content) + + if hasattr(message, "tool_calls") and message.tool_calls: + for tool_call in message.tool_calls: + if isinstance(tool_call, dict): + message_buffer.add_tool_call(tool_call["name"], tool_call["args"]) + else: + message_buffer.add_tool_call(tool_call.name, tool_call.args) + + # Update analyst statuses based on report state (runs on every chunk) + update_analyst_statuses( + message_buffer, + chunk, + wall_time_tracker=analyst_wall_time_tracker, + ) + + # Research Team - Handle Investment Debate State + if chunk.get("investment_debate_state"): + debate_state = chunk["investment_debate_state"] + bull_hist = debate_state.get("bull_history", "").strip() + bear_hist = debate_state.get("bear_history", "").strip() + judge = debate_state.get("judge_decision", "").strip() + + # Only update status when there's actual content + if bull_hist or bear_hist: + update_research_team_status("in_progress") + if bull_hist: + message_buffer.update_report_section( + "investment_plan", f"### Bull Researcher Analysis\n{bull_hist}" + ) + if bear_hist: + message_buffer.update_report_section( + "investment_plan", f"### Bear Researcher Analysis\n{bear_hist}" + ) + if judge: + message_buffer.update_report_section( + "investment_plan", f"### Research Manager Decision\n{judge}" + ) + update_research_team_status("completed") + message_buffer.update_agent_status("Trader", "in_progress") + + # Trading Team + if chunk.get("trader_investment_plan"): + message_buffer.update_report_section( + "trader_investment_plan", chunk["trader_investment_plan"] + ) + if message_buffer.agent_status.get("Trader") != "completed": + message_buffer.update_agent_status("Trader", "completed") + message_buffer.update_agent_status("Aggressive Analyst", "in_progress") + + # Risk Management Team - Handle Risk Debate State + if chunk.get("risk_debate_state"): + risk_state = chunk["risk_debate_state"] + agg_hist = risk_state.get("aggressive_history", "").strip() + con_hist = risk_state.get("conservative_history", "").strip() + neu_hist = risk_state.get("neutral_history", "").strip() + judge = risk_state.get("judge_decision", "").strip() + + if agg_hist: + if message_buffer.agent_status.get("Aggressive Analyst") != "completed": + message_buffer.update_agent_status("Aggressive Analyst", "in_progress") + message_buffer.update_report_section( + "final_trade_decision", f"### Aggressive Analyst Analysis\n{agg_hist}" + ) + if con_hist: + if message_buffer.agent_status.get("Conservative Analyst") != "completed": + message_buffer.update_agent_status("Conservative Analyst", "in_progress") + message_buffer.update_report_section( + "final_trade_decision", f"### Conservative Analyst Analysis\n{con_hist}" + ) + if neu_hist: + if message_buffer.agent_status.get("Neutral Analyst") != "completed": + message_buffer.update_agent_status("Neutral Analyst", "in_progress") + message_buffer.update_report_section( + "final_trade_decision", f"### Neutral Analyst Analysis\n{neu_hist}" + ) + if judge and message_buffer.agent_status.get("Portfolio Manager") != "completed": + message_buffer.update_agent_status("Portfolio Manager", "in_progress") + message_buffer.update_report_section( + "final_trade_decision", f"### Portfolio Manager Decision\n{judge}" + ) + message_buffer.update_agent_status("Aggressive Analyst", "completed") + message_buffer.update_agent_status("Conservative Analyst", "completed") + message_buffer.update_agent_status("Neutral Analyst", "completed") + message_buffer.update_agent_status("Portfolio Manager", "completed") + + # Update the display + update_display(layout, stats_handler=stats_handler, start_time=start_time) + + trace.append(chunk) + + # Streamed chunks are per-node deltas, not full state. Merge them + # so every report field populated across the run is present. + final_state = {} + for chunk in trace: + final_state.update(chunk) + + # Clean run: log the decision, then drop this run's checkpoint so a + # later run starts fresh. A mid-stream failure skips both, keeping + # the checkpoint for resume. + graph.record_decision(selections["ticker"], selections["analysis_date"], final_state) + graph.clear_checkpoint_on_success( + selections["ticker"], selections["analysis_date"], selections["asset_type"], portfolio + ) + finally: + # Always restore the plain uncheckpointed graph, even on failure. + graph.end_checkpoint() + + # Update all agent statuses to completed + for agent in message_buffer.agent_status: + message_buffer.update_agent_status(agent, "completed") + + message_buffer.add_message( + "System", f"Completed analysis for {selections['analysis_date']}" + ) + message_buffer.add_message("System", analyst_wall_time_tracker.format_summary()) + + # Update final report sections + for section in message_buffer.report_sections: + if section in final_state: + message_buffer.update_report_section(section, final_state[section]) + + update_display(layout, stats_handler=stats_handler, start_time=start_time) + + # Post-analysis prompts (outside Live context for clean interaction) + console.print("\n[bold cyan]Analysis Complete![/bold cyan]\n") + + # A decision nobody can read is not a position. Say so here rather than + # leaving the run to look like a normal result. + if is_review(graph.process_signal(final_state.get("final_trade_decision", ""))): + console.print( + "[yellow]No rating could be read from the final decision, so this run " + "is recorded for review rather than as a position. Re-run, or read the " + "decision text below and judge it yourself.[/yellow]\n" + ) + console.print(f"[dim]{analyst_wall_time_tracker.format_summary()}[/dim]") + + # Prompt to save report + save_choice = typer.prompt("Save report?", default="Y").strip().upper() + if save_choice in ("Y", "YES", ""): + timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") + # Under results_dir, not the working directory: in Docker the working + # directory is inside the container and the report goes with it, while + # results_dir is the mounted volume the rest of the run already writes to. + default_path = (Path(config["results_dir"]) / "reports" + / f"{safe_ticker_component(selections['ticker'])}_{timestamp}") + save_path_str = typer.prompt( + "Save path (press Enter for default)", + default=str(default_path) + ).strip() + save_path = Path(save_path_str) + try: + report_file = write_report_tree(final_state, selections["ticker"], save_path) + console.print(f"\n[green]✓ Report saved to:[/green] {save_path.resolve()}") + console.print(f" [dim]Complete report:[/dim] {report_file.name}") + except Exception as e: + console.print(f"[red]Error saving report: {e}[/red]") + + # Prompt to display full report + display_choice = typer.prompt("\nDisplay full report on screen?", default="Y").strip().upper() + if display_choice in ("Y", "YES", ""): + display_complete_report(final_state) diff --git a/cli/selections.py b/cli/selections.py new file mode 100644 index 000000000..ec263ac56 --- /dev/null +++ b/cli/selections.py @@ -0,0 +1,319 @@ +"""The interactive choices for a run: ticker, date, analysts, depth, provider and models.""" + +import datetime +import os +from pathlib import Path + +import typer +from rich.align import Align +from rich.panel import Panel + +from cli.announcements import display_announcements, fetch_announcements +from cli.display import ( + console, +) +from cli.prefs import load_last_run, sanitize, save_last_run +from cli.prompts import ( + ask_anthropic_effort, + ask_gemini_thinking_config, + ask_glm_region, + ask_minimax_region, + ask_openai_reasoning_effort, + ask_output_language, + ask_qwen_region, + confirm_ollama_endpoint, + detect_asset_type, + ensure_api_key, + get_ticker, + prompt_openai_compatible_url, + resolve_backend_url, + select_analysts, + select_deep_thinking_agent, + select_llm_provider, + select_research_depth, + select_shallow_thinking_agent, +) +from tradingagents.default_config import DEFAULT_CONFIG + + +def get_user_selections(): + """Ask for the run's settings, offering the previous run's answers.""" + selections = _prompt_selections(load_last_run()) + save_last_run(selections) + return selections + + +def _prompt_selections(prefs): + """Walk the selection steps. ``prefs`` prefills, the environment skips.""" + # Display ASCII art welcome message + with open(Path(__file__).parent / "static" / "welcome.txt", encoding="utf-8") as f: + welcome_ascii = f.read() + + # Create welcome box content + welcome_content = f"{welcome_ascii}\n" + welcome_content += "[bold green]TradingAgents: Multi-Agents LLM Financial Trading Framework - CLI[/bold green]\n\n" + welcome_content += "[bold]Workflow Steps:[/bold]\n" + welcome_content += "I. Analyst Team → II. Research Team → III. Trader → IV. Risk Management → V. Portfolio Management\n\n" + welcome_content += ( + "[dim]Built by [Tauric Research](https://github.com/TauricResearch)[/dim]" + ) + + # Create and center the welcome box + welcome_box = Panel( + welcome_content, + border_style="green", + padding=(1, 2), + title="Welcome to TradingAgents", + subtitle="Multi-Agents LLM Financial Trading Framework", + ) + console.print(Align.center(welcome_box)) + console.print() + console.print() # Add vertical space before announcements + + # Fetch and display announcements (silent on failure) + announcements = fetch_announcements() + display_announcements(console, announcements) + + # Create a boxed questionnaire for each step + def create_question_box(title, prompt, default=None): + box_content = f"[bold]{title}[/bold]\n" + box_content += f"[dim]{prompt}[/dim]" + if default: + box_content += f"\n[dim]Default: {default}[/dim]" + return Panel(box_content, border_style="blue", padding=(1, 2)) + + def thinking_value_or_prompt(env_var, config_key, label, box_title, box_body, prompt_fn): + """Return the env-configured reasoning/thinking value, or prompt for it. + + When ``env_var`` is set the interactive choice is skipped and the value + the env overlay placed on DEFAULT_CONFIG is used — mirroring the + env-precedence rule applied to the other selection steps. + """ + if os.environ.get(env_var): + value = DEFAULT_CONFIG[config_key] + console.print(f"[green]✓ {label} from environment:[/green] {value}") + return value + console.print(create_question_box(box_title, box_body)) + return prompt_fn() + + # Step 1: Ticker symbol + console.print( + create_question_box( + "Step 1: Ticker Symbol", + "Enter the ticker, with exchange suffix when needed (e.g. SPY, 0700.HK, BTC-USD)", + "SPY", + ) + ) + selected_ticker = get_ticker() + asset_type = detect_asset_type(selected_ticker) + # Only announce when it's not the default stock path, to avoid printing + # "stock" on every run. + if asset_type.value != "stock": + console.print( + f"[green]Detected asset type:[/green] {asset_type.value}" + ) + + # Step 2: Analysis date + default_date = datetime.datetime.now().strftime("%Y-%m-%d") + console.print( + create_question_box( + "Step 2: Analysis Date", + "Enter the analysis date (YYYY-MM-DD)", + default_date, + ) + ) + analysis_date = get_analysis_date() + + # Step 3: Output language (skipped when set via TRADINGAGENTS_OUTPUT_LANGUAGE) + if os.environ.get("TRADINGAGENTS_OUTPUT_LANGUAGE"): + output_language = DEFAULT_CONFIG["output_language"] + console.print( + f"[green]✓ Output language from environment:[/green] {output_language}" + ) + else: + console.print( + create_question_box( + "Step 3: Output Language", + "Select the language for analyst reports and final decision" + ) + ) + output_language = ask_output_language(prefs.get("output_language")) + + # Step 4: Select analysts + console.print( + create_question_box( + "Step 4: Analysts Team", "Select your LLM analyst agents for the analysis" + ) + ) + prefs = sanitize(prefs, asset_type.value) + selected_analysts = select_analysts(asset_type, prefs.get("analysts")) + console.print( + f"[green]Selected analysts:[/green] {', '.join(analyst.value for analyst in selected_analysts)}" + ) + + # Step 5: Research depth (skipped when both round counts are set via env). + # Research depth maps to the debate + risk round counts; when both are + # supplied through TRADINGAGENTS_MAX_DEBATE_ROUNDS / _MAX_RISK_ROUNDS we keep + # the run non-interactive and honor the env values (#977). + depth_from_env = bool(os.environ.get("TRADINGAGENTS_MAX_DEBATE_ROUNDS")) and bool( + os.environ.get("TRADINGAGENTS_MAX_RISK_ROUNDS") + ) + if depth_from_env: + selected_research_depth = DEFAULT_CONFIG["max_debate_rounds"] + console.print( + f"[green]✓ Research depth from environment:[/green] " + f"{DEFAULT_CONFIG['max_debate_rounds']} debate / " + f"{DEFAULT_CONFIG['max_risk_discuss_rounds']} risk rounds" + ) + else: + console.print( + create_question_box( + "Step 5: Research Depth", "Select your research depth level" + ) + ) + selected_research_depth = select_research_depth(prefs.get("research_depth")) + + # Step 6: LLM Provider (skipped when set via TRADINGAGENTS_LLM_PROVIDER). + # The backend URL comes from TRADINGAGENTS_LLM_BACKEND_URL when set, + # otherwise the provider's default endpoint — the same value the menu + # would have picked. + provider_from_env = bool(os.environ.get("TRADINGAGENTS_LLM_PROVIDER")) + if provider_from_env: + selected_llm_provider = DEFAULT_CONFIG["llm_provider"].lower() + backend_url = resolve_backend_url( + selected_llm_provider, env_url=DEFAULT_CONFIG["backend_url"] + ) + console.print(f"[green]✓ LLM provider from environment:[/green] {selected_llm_provider}") + console.print(f"[green]✓ Backend URL:[/green] {backend_url}") + # Still confirm/persist the API key so the run doesn't fail later. + ensure_api_key(selected_llm_provider) + else: + console.print( + create_question_box( + "Step 6: LLM Provider", "Select your LLM provider" + ) + ) + selected_llm_provider, backend_url = select_llm_provider(prefs.get("llm_provider")) + + # Providers with regional endpoints prompt for the region as a secondary + # step so the main dropdown stays clean (mainland China and international + # accounts cannot share API keys). + if selected_llm_provider == "qwen": + selected_llm_provider, backend_url = ask_qwen_region() + elif selected_llm_provider == "minimax": + selected_llm_provider, backend_url = ask_minimax_region() + elif selected_llm_provider == "glm": + selected_llm_provider, backend_url = ask_glm_region() + + # Honor an explicit env backend URL even when the provider was chosen + # interactively, so it isn't overwritten by the menu default (#978). + backend_url = resolve_backend_url( + selected_llm_provider, backend_url, env_url=DEFAULT_CONFIG["backend_url"] + ) + + # The generic OpenAI-compatible endpoint has no default; ask for it if + # neither the menu nor the environment supplied one. + if selected_llm_provider == "openai_compatible" and not backend_url: + remembered_url = (prefs.get("backend_url") + if prefs.get("llm_provider") == selected_llm_provider else None) + backend_url = prompt_openai_compatible_url(remembered_url) + + # For Ollama, surface the resolved endpoint (OLLAMA_BASE_URL vs default) + # before model selection so it's obvious where we're connecting. + if selected_llm_provider == "ollama": + confirm_ollama_endpoint(backend_url) + + # Confirm the provider's API key is present; prompt the user to paste + # one and persist it to .env if it's missing, so the analysis run + # doesn't fail later at the first API call. + ensure_api_key(selected_llm_provider) + + # Step 7: Thinking agents (skipped when either model is set via environment) + if os.environ.get("TRADINGAGENTS_QUICK_THINK_LLM") or os.environ.get("TRADINGAGENTS_DEEP_THINK_LLM"): + selected_shallow_thinker = DEFAULT_CONFIG["quick_think_llm"] + selected_deep_thinker = DEFAULT_CONFIG["deep_think_llm"] + console.print( + f"[green]✓ Thinking agents from environment:[/green] " + f"quick={selected_shallow_thinker}, deep={selected_deep_thinker}" + ) + else: + console.print( + create_question_box( + "Step 7: Thinking Agents", "Select your thinking agents for analysis" + ) + ) + remembered = prefs if prefs.get("llm_provider") == selected_llm_provider else {} + selected_shallow_thinker = select_shallow_thinking_agent( + selected_llm_provider, remembered.get("quick_think_llm") + ) + selected_deep_thinker = select_deep_thinking_agent( + selected_llm_provider, remembered.get("deep_think_llm") + ) + + # Step 8: Provider-specific reasoning/thinking configuration. Each knob is + # settable via its TRADINGAGENTS_* env var; when that var is set (or the + # provider itself came from env) the prompt is skipped and the configured + # value is used — same env-precedence rule as the steps above. None = each + # provider's own default. + thinking_level = None + reasoning_effort = None + anthropic_effort = None + + provider_lower = selected_llm_provider.lower() + if provider_from_env: + thinking_level = DEFAULT_CONFIG["google_thinking_level"] + reasoning_effort = DEFAULT_CONFIG["openai_reasoning_effort"] + anthropic_effort = DEFAULT_CONFIG["anthropic_effort"] + elif provider_lower == "google": + thinking_level = thinking_value_or_prompt( + "TRADINGAGENTS_GOOGLE_THINKING_LEVEL", "google_thinking_level", + "Gemini thinking mode", "Step 8: Thinking Mode", + "Configure Gemini thinking mode", ask_gemini_thinking_config, + ) + elif provider_lower == "openai": + reasoning_effort = thinking_value_or_prompt( + "TRADINGAGENTS_OPENAI_REASONING_EFFORT", "openai_reasoning_effort", + "Reasoning effort", "Step 8: Reasoning Effort", + "Configure OpenAI reasoning effort level", ask_openai_reasoning_effort, + ) + elif provider_lower == "anthropic": + anthropic_effort = thinking_value_or_prompt( + "TRADINGAGENTS_ANTHROPIC_EFFORT", "anthropic_effort", + "Claude effort", "Step 8: Effort Level", + "Configure Claude effort level", ask_anthropic_effort, + ) + + return { + "ticker": selected_ticker, + "asset_type": asset_type.value, + "analysis_date": analysis_date, + "analysts": selected_analysts, + "research_depth": selected_research_depth, + "llm_provider": selected_llm_provider.lower(), + "backend_url": backend_url, + "quick_think_llm": selected_shallow_thinker, + "deep_think_llm": selected_deep_thinker, + "google_thinking_level": thinking_level, + "openai_reasoning_effort": reasoning_effort, + "anthropic_effort": anthropic_effort, + "output_language": output_language, + } + + +def get_analysis_date(): + """Get the analysis date from user input.""" + while True: + date_str = typer.prompt( + "", default=datetime.datetime.now().strftime("%Y-%m-%d") + ) + try: + # Validate date format and ensure it's not in the future + analysis_date = datetime.datetime.strptime(date_str, "%Y-%m-%d") + if analysis_date.date() > datetime.datetime.now().date(): + console.print("[red]Error: Analysis date cannot be in the future[/red]") + continue + return date_str + except ValueError: + console.print( + "[red]Error: Invalid date format. Please use YYYY-MM-DD[/red]" + ) diff --git a/tests/test_cli_config_precedence.py b/tests/test_cli_config_precedence.py index 54ff41da9..cb94cff8a 100644 --- a/tests/test_cli_config_precedence.py +++ b/tests/test_cli_config_precedence.py @@ -10,6 +10,7 @@ from unittest import mock import pytest import cli.main as m +import cli.run as cli_run # Minimal selections dict shaped like get_user_selections()'s return value. SELECTIONS = { @@ -28,7 +29,7 @@ SELECTIONS = { def test_research_depth_sets_both_rounds_without_env(monkeypatch): for var in ("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "TRADINGAGENTS_MAX_RISK_ROUNDS"): monkeypatch.delenv(var, raising=False) - cfg = m._build_run_config(SELECTIONS, checkpoint=None) + cfg = cli_run._build_run_config(SELECTIONS, checkpoint=None) assert cfg["max_debate_rounds"] == 5 assert cfg["max_risk_discuss_rounds"] == 5 @@ -37,9 +38,9 @@ def test_env_round_counts_win_over_selection(monkeypatch): monkeypatch.setenv("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "2") monkeypatch.setenv("TRADINGAGENTS_MAX_RISK_ROUNDS", "4") # DEFAULT_CONFIG already reflects the env (applied at import); emulate that. - patched = dict(m.DEFAULT_CONFIG, max_debate_rounds=2, max_risk_discuss_rounds=4) - with mock.patch.object(m, "DEFAULT_CONFIG", patched): - cfg = m._build_run_config(SELECTIONS, checkpoint=None) + patched = dict(cli_run.DEFAULT_CONFIG, max_debate_rounds=2, max_risk_discuss_rounds=4) + with mock.patch.object(cli_run, "DEFAULT_CONFIG", patched): + cfg = cli_run._build_run_config(SELECTIONS, checkpoint=None) assert cfg["max_debate_rounds"] == 2 # env value, not research_depth=5 assert cfg["max_risk_discuss_rounds"] == 4 @@ -47,25 +48,25 @@ def test_env_round_counts_win_over_selection(monkeypatch): def test_partial_env_only_overrides_that_count(monkeypatch): monkeypatch.setenv("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "2") monkeypatch.delenv("TRADINGAGENTS_MAX_RISK_ROUNDS", raising=False) - patched = dict(m.DEFAULT_CONFIG, max_debate_rounds=2) - with mock.patch.object(m, "DEFAULT_CONFIG", patched): - cfg = m._build_run_config(SELECTIONS, checkpoint=None) + patched = dict(cli_run.DEFAULT_CONFIG, max_debate_rounds=2) + with mock.patch.object(cli_run, "DEFAULT_CONFIG", patched): + cfg = cli_run._build_run_config(SELECTIONS, checkpoint=None) assert cfg["max_debate_rounds"] == 2 # env wins assert cfg["max_risk_discuss_rounds"] == 5 # falls through to research_depth def test_checkpoint_none_preserves_env_default(): - patched = dict(m.DEFAULT_CONFIG, checkpoint_enabled=True) # e.g. env-enabled - with mock.patch.object(m, "DEFAULT_CONFIG", patched): - cfg = m._build_run_config(SELECTIONS, checkpoint=None) + patched = dict(cli_run.DEFAULT_CONFIG, checkpoint_enabled=True) # e.g. env-enabled + with mock.patch.object(cli_run, "DEFAULT_CONFIG", patched): + cfg = cli_run._build_run_config(SELECTIONS, checkpoint=None) assert cfg["checkpoint_enabled"] is True # not clobbered back to False @pytest.mark.parametrize("flag", [True, False]) def test_checkpoint_flag_overrides_env(flag): - patched = dict(m.DEFAULT_CONFIG, checkpoint_enabled=not flag) - with mock.patch.object(m, "DEFAULT_CONFIG", patched): - cfg = m._build_run_config(SELECTIONS, checkpoint=flag) + patched = dict(cli_run.DEFAULT_CONFIG, checkpoint_enabled=not flag) + with mock.patch.object(cli_run, "DEFAULT_CONFIG", patched): + cfg = cli_run._build_run_config(SELECTIONS, checkpoint=flag) assert cfg["checkpoint_enabled"] is flag @@ -89,14 +90,13 @@ def test_glm_resolves_to_the_endpoint_its_key_belongs_to(): def test_a_half_set_round_count_says_which_value_won(capsys, monkeypatch): """With only one of the two round-count variables set, the depth prompt is still shown but half the answer is discarded; the user was never told.""" - import cli.main as m monkeypatch.setenv("TRADINGAGENTS_MAX_DEBATE_ROUNDS", "1") monkeypatch.delenv("TRADINGAGENTS_MAX_RISK_ROUNDS", raising=False) printed = [] monkeypatch.setattr(m.console, "print", lambda *a, **k: printed.append(str(a[0]) if a else "")) - config = m._build_run_config({ + config = cli_run._build_run_config({ "ticker": "NVDA", "analysis_date": "2026-09-01", "asset_type": "stock", "analysts": [], "research_depth": 5, "llm_provider": "openai", "quick_think_llm": "gpt-5.6-luna", "deep_think_llm": "gpt-5.6", diff --git a/tests/test_cli_decision_log.py b/tests/test_cli_decision_log.py index f9f291bd7..be695b095 100644 --- a/tests/test_cli_decision_log.py +++ b/tests/test_cli_decision_log.py @@ -11,6 +11,7 @@ from __future__ import annotations import pytest +import cli.run as cli_run from tradingagents.decision_log import TradingMemoryLog from tradingagents.graph.trading_graph import TradingAgentsGraph @@ -145,20 +146,20 @@ def _run_cli(monkeypatch, tmp_path, fake): from cli.models import AnalystType buffer = _FakeBuffer() - monkeypatch.setattr(m, "TradingAgentsGraph", lambda *a, **k: fake) - monkeypatch.setattr(m, "message_buffer", buffer) - monkeypatch.setattr(m, "create_layout", lambda: None) - monkeypatch.setattr(m, "update_display", lambda *a, **k: None) - monkeypatch.setattr(m, "Live", _NullLive) - monkeypatch.setattr(m, "get_user_selections", lambda: { + monkeypatch.setattr(cli_run, "TradingAgentsGraph", lambda *a, **k: fake) + monkeypatch.setattr(cli_run, "message_buffer", buffer) + monkeypatch.setattr(cli_run, "create_layout", lambda: None) + monkeypatch.setattr(cli_run, "update_display", lambda *a, **k: None) + monkeypatch.setattr(cli_run, "Live", _NullLive) + monkeypatch.setattr(cli_run, "get_user_selections", lambda: { "ticker": "NVDA", "analysis_date": "2026-01-10", "analysts": [AnalystType.MARKET], "asset_type": "stock", }) - monkeypatch.setattr(m, "_build_run_config", lambda selections, checkpoint: { + monkeypatch.setattr(cli_run, "_build_run_config", lambda selections, checkpoint: { "data_cache_dir": str(tmp_path / "cache"), "results_dir": str(tmp_path / "results"), }) monkeypatch.setattr(m.typer, "prompt", lambda *a, **k: "N") - m.run_analysis() + cli_run.run_analysis() return buffer diff --git a/tests/test_cli_env_skip.py b/tests/test_cli_env_skip.py index d50bfc057..812f016fb 100644 --- a/tests/test_cli_env_skip.py +++ b/tests/test_cli_env_skip.py @@ -11,6 +11,8 @@ from unittest import mock import pytest +import cli.selections as cli_selections + @pytest.mark.unit class TestProviderDefaultUrl(unittest.TestCase): @@ -33,7 +35,6 @@ class TestProviderDefaultUrl(unittest.TestCase): @pytest.mark.unit class TestCliSkipsPromptsFromEnv(unittest.TestCase): def test_env_config_skips_llm_prompts(self): - import cli.main as m env = { "TRADINGAGENTS_LLM_PROVIDER": "openai", @@ -42,7 +43,7 @@ class TestCliSkipsPromptsFromEnv(unittest.TestCase): "TRADINGAGENTS_LLM_BACKEND_URL": "https://opencode.ai/zen/go/v1", "TRADINGAGENTS_OUTPUT_LANGUAGE": "Japanese", } - fake_cfg = dict(m.DEFAULT_CONFIG) + fake_cfg = dict(cli_selections.DEFAULT_CONFIG) fake_cfg.update({ "llm_provider": "openai", "backend_url": "https://opencode.ai/zen/go/v1", @@ -52,19 +53,19 @@ class TestCliSkipsPromptsFromEnv(unittest.TestCase): }) with mock.patch.dict(os.environ, env, clear=False), \ - mock.patch.object(m, "DEFAULT_CONFIG", fake_cfg), \ - mock.patch.object(m, "fetch_announcements", return_value=None), \ - mock.patch.object(m, "display_announcements"), \ - mock.patch.object(m, "get_ticker", return_value="AAPL"), \ - mock.patch.object(m, "get_analysis_date", return_value="2026-05-29"), \ - mock.patch.object(m, "select_analysts", return_value=[]), \ - mock.patch.object(m, "select_research_depth", return_value=1), \ - mock.patch.object(m, "ensure_api_key") as ensure_key, \ - mock.patch.object(m, "select_llm_provider") as prompt_provider, \ - mock.patch.object(m, "ask_output_language") as prompt_lang, \ - mock.patch.object(m, "select_shallow_thinking_agent") as prompt_quick, \ - mock.patch.object(m, "select_deep_thinking_agent") as prompt_deep: - sel = m.get_user_selections() + mock.patch.object(cli_selections, "DEFAULT_CONFIG", fake_cfg), \ + mock.patch.object(cli_selections, "fetch_announcements", return_value=None), \ + mock.patch.object(cli_selections, "display_announcements"), \ + mock.patch.object(cli_selections, "get_ticker", return_value="AAPL"), \ + mock.patch.object(cli_selections, "get_analysis_date", return_value="2026-05-29"), \ + mock.patch.object(cli_selections, "select_analysts", return_value=[]), \ + mock.patch.object(cli_selections, "select_research_depth", return_value=1), \ + mock.patch.object(cli_selections, "ensure_api_key") as ensure_key, \ + mock.patch.object(cli_selections, "select_llm_provider") as prompt_provider, \ + mock.patch.object(cli_selections, "ask_output_language") as prompt_lang, \ + mock.patch.object(cli_selections, "select_shallow_thinking_agent") as prompt_quick, \ + mock.patch.object(cli_selections, "select_deep_thinking_agent") as prompt_deep: + sel = cli_selections.get_user_selections() # None of the LLM selection prompts should have been shown. prompt_provider.assert_not_called() @@ -85,30 +86,29 @@ class TestCliSkipsPromptsFromEnv(unittest.TestCase): @pytest.mark.unit class TestResearchDepthSkippedFromEnv(unittest.TestCase): def test_both_round_envs_skip_depth_prompt(self): - import cli.main as m env = { "TRADINGAGENTS_MAX_DEBATE_ROUNDS": "2", "TRADINGAGENTS_MAX_RISK_ROUNDS": "4", } - fake_cfg = dict(m.DEFAULT_CONFIG) + fake_cfg = dict(cli_selections.DEFAULT_CONFIG) fake_cfg.update({"max_debate_rounds": 2, "max_risk_discuss_rounds": 4}) with mock.patch.dict(os.environ, env, clear=False), \ - mock.patch.object(m, "DEFAULT_CONFIG", fake_cfg), \ - mock.patch.object(m, "fetch_announcements", return_value=None), \ - mock.patch.object(m, "display_announcements"), \ - mock.patch.object(m, "get_ticker", return_value="AAPL"), \ - mock.patch.object(m, "get_analysis_date", return_value="2026-05-29"), \ - mock.patch.object(m, "select_analysts", return_value=[]), \ - mock.patch.object(m, "select_research_depth") as prompt_depth, \ - mock.patch.object(m, "ensure_api_key"), \ - mock.patch.object(m, "select_llm_provider", return_value=("openai", None)), \ - mock.patch.object(m, "ask_output_language", return_value="English"), \ - mock.patch.object(m, "select_shallow_thinking_agent", return_value="gpt-5.4-mini"), \ - mock.patch.object(m, "select_deep_thinking_agent", return_value="gpt-5.5"), \ - mock.patch.object(m, "ask_openai_reasoning_effort", return_value=None): - sel = m.get_user_selections() + mock.patch.object(cli_selections, "DEFAULT_CONFIG", fake_cfg), \ + mock.patch.object(cli_selections, "fetch_announcements", return_value=None), \ + mock.patch.object(cli_selections, "display_announcements"), \ + mock.patch.object(cli_selections, "get_ticker", return_value="AAPL"), \ + mock.patch.object(cli_selections, "get_analysis_date", return_value="2026-05-29"), \ + mock.patch.object(cli_selections, "select_analysts", return_value=[]), \ + mock.patch.object(cli_selections, "select_research_depth") as prompt_depth, \ + mock.patch.object(cli_selections, "ensure_api_key"), \ + mock.patch.object(cli_selections, "select_llm_provider", return_value=("openai", None)), \ + mock.patch.object(cli_selections, "ask_output_language", return_value="English"), \ + mock.patch.object(cli_selections, "select_shallow_thinking_agent", return_value="gpt-5.4-mini"), \ + mock.patch.object(cli_selections, "select_deep_thinking_agent", return_value="gpt-5.5"), \ + mock.patch.object(cli_selections, "ask_openai_reasoning_effort", return_value=None): + sel = cli_selections.get_user_selections() # The research-depth prompt is skipped; the value comes from the env config. prompt_depth.assert_not_called() @@ -118,27 +118,26 @@ class TestResearchDepthSkippedFromEnv(unittest.TestCase): @pytest.mark.unit class TestReasoningEffortSkippedFromEnv(unittest.TestCase): def test_effort_env_skips_step8_prompt(self): - import cli.main as m env = {"TRADINGAGENTS_OPENAI_REASONING_EFFORT": "high"} - fake_cfg = dict(m.DEFAULT_CONFIG) + fake_cfg = dict(cli_selections.DEFAULT_CONFIG) fake_cfg.update({"openai_reasoning_effort": "high"}) with mock.patch.dict(os.environ, env, clear=False), \ - mock.patch.object(m, "DEFAULT_CONFIG", fake_cfg), \ - mock.patch.object(m, "fetch_announcements", return_value=None), \ - mock.patch.object(m, "display_announcements"), \ - mock.patch.object(m, "get_ticker", return_value="AAPL"), \ - mock.patch.object(m, "get_analysis_date", return_value="2026-05-29"), \ - mock.patch.object(m, "select_analysts", return_value=[]), \ - mock.patch.object(m, "select_research_depth", return_value=1), \ - mock.patch.object(m, "ensure_api_key"), \ - mock.patch.object(m, "select_llm_provider", return_value=("openai", None)), \ - mock.patch.object(m, "ask_output_language", return_value="English"), \ - mock.patch.object(m, "select_shallow_thinking_agent", return_value="gpt-5.4-mini"), \ - mock.patch.object(m, "select_deep_thinking_agent", return_value="gpt-5.5"), \ - mock.patch.object(m, "ask_openai_reasoning_effort") as prompt_effort: - sel = m.get_user_selections() + mock.patch.object(cli_selections, "DEFAULT_CONFIG", fake_cfg), \ + mock.patch.object(cli_selections, "fetch_announcements", return_value=None), \ + mock.patch.object(cli_selections, "display_announcements"), \ + mock.patch.object(cli_selections, "get_ticker", return_value="AAPL"), \ + mock.patch.object(cli_selections, "get_analysis_date", return_value="2026-05-29"), \ + mock.patch.object(cli_selections, "select_analysts", return_value=[]), \ + mock.patch.object(cli_selections, "select_research_depth", return_value=1), \ + mock.patch.object(cli_selections, "ensure_api_key"), \ + mock.patch.object(cli_selections, "select_llm_provider", return_value=("openai", None)), \ + mock.patch.object(cli_selections, "ask_output_language", return_value="English"), \ + mock.patch.object(cli_selections, "select_shallow_thinking_agent", return_value="gpt-5.4-mini"), \ + mock.patch.object(cli_selections, "select_deep_thinking_agent", return_value="gpt-5.5"), \ + mock.patch.object(cli_selections, "ask_openai_reasoning_effort") as prompt_effort: + sel = cli_selections.get_user_selections() # The reasoning-effort prompt is skipped; the value comes from env config. prompt_effort.assert_not_called() diff --git a/tests/test_cli_prefs.py b/tests/test_cli_prefs.py index 7271a011f..a833c52e1 100644 --- a/tests/test_cli_prefs.py +++ b/tests/test_cli_prefs.py @@ -13,6 +13,7 @@ from unittest import mock import pytest +import cli.selections as cli_selections from cli.models import AnalystType from cli.prefs import load_last_run, sanitize, save_last_run @@ -100,26 +101,26 @@ def _answer_every_prompt(monkeypatch): """Drive the real selection flow, answering each prompt with a fixed value.""" import cli.main as m - monkeypatch.setattr(m, "fetch_announcements", lambda: []) - monkeypatch.setattr(m, "display_announcements", lambda *a: None) - monkeypatch.setattr(m, "get_ticker", lambda: "NVDA") - monkeypatch.setattr(m, "get_analysis_date", lambda: "2026-09-01") - monkeypatch.setattr(m, "ask_output_language", lambda default=None: "English") - monkeypatch.setattr(m, "select_analysts", lambda asset_type, default=None: [AnalystType.MARKET]) - monkeypatch.setattr(m, "select_research_depth", lambda default=None: 3) - monkeypatch.setattr(m, "select_llm_provider", lambda default=None: ("openai", None)) - monkeypatch.setattr(m, "select_shallow_thinking_agent", lambda p, default=None: "gpt-5.6-mini") - monkeypatch.setattr(m, "select_deep_thinking_agent", lambda p, default=None: "gpt-5.6") - monkeypatch.setattr(m, "ask_openai_reasoning_effort", lambda: "medium") + monkeypatch.setattr(cli_selections, "fetch_announcements", lambda: []) + monkeypatch.setattr(cli_selections, "display_announcements", lambda *a: None) + monkeypatch.setattr(cli_selections, "get_ticker", lambda: "NVDA") + monkeypatch.setattr(cli_selections, "get_analysis_date", lambda: "2026-09-01") + monkeypatch.setattr(cli_selections, "ask_output_language", lambda default=None: "English") + monkeypatch.setattr(cli_selections, "select_analysts", lambda asset_type, default=None: [AnalystType.MARKET]) + monkeypatch.setattr(cli_selections, "select_research_depth", lambda default=None: 3) + monkeypatch.setattr(cli_selections, "select_llm_provider", lambda default=None: ("openai", None)) + monkeypatch.setattr(cli_selections, "select_shallow_thinking_agent", lambda p, default=None: "gpt-5.6-mini") + monkeypatch.setattr(cli_selections, "select_deep_thinking_agent", lambda p, default=None: "gpt-5.6") + monkeypatch.setattr(cli_selections, "ask_openai_reasoning_effort", lambda: "medium") return m @pytest.mark.unit def test_selections_are_remembered_after_a_run(monkeypatch): """Drives the real flow: a stubbed selections dict would hide a key mismatch.""" - m = _answer_every_prompt(monkeypatch) + _answer_every_prompt(monkeypatch) - m.get_user_selections() + cli_selections.get_user_selections() remembered = load_last_run() assert remembered["analysts"] == ["market"] @@ -147,23 +148,22 @@ def test_a_custom_language_is_remembered_without_breaking_the_next_run(): def test_a_remembered_endpoint_is_offered_back(monkeypatch): """Users of a local or custom endpoint retyped the URL every run: it was remembered and validated, then never read.""" - import cli.main as m save_last_run({"llm_provider": "openai_compatible", "backend_url": "http://localhost:1234/v1"}) offered = {} - monkeypatch.setattr(m, "select_llm_provider", lambda default=None: ("openai_compatible", None)) - monkeypatch.setattr(m, "prompt_openai_compatible_url", + monkeypatch.setattr(cli_selections, "select_llm_provider", lambda default=None: ("openai_compatible", None)) + monkeypatch.setattr(cli_selections, "prompt_openai_compatible_url", lambda default=None: offered.setdefault("default", default) or "http://x/v1") - monkeypatch.setattr(m, "fetch_announcements", lambda: []) - monkeypatch.setattr(m, "display_announcements", lambda *a: None) - monkeypatch.setattr(m, "get_ticker", lambda: "NVDA") - monkeypatch.setattr(m, "get_analysis_date", lambda: "2026-09-01") - monkeypatch.setattr(m, "ask_output_language", lambda default=None: "English") - monkeypatch.setattr(m, "select_analysts", lambda asset_type, default=None: [AnalystType.MARKET]) - monkeypatch.setattr(m, "select_research_depth", lambda default=None: 1) - monkeypatch.setattr(m, "select_shallow_thinking_agent", lambda p, default=None: "local-model") - monkeypatch.setattr(m, "select_deep_thinking_agent", lambda p, default=None: "local-model") + monkeypatch.setattr(cli_selections, "fetch_announcements", lambda: []) + monkeypatch.setattr(cli_selections, "display_announcements", lambda *a: None) + monkeypatch.setattr(cli_selections, "get_ticker", lambda: "NVDA") + monkeypatch.setattr(cli_selections, "get_analysis_date", lambda: "2026-09-01") + monkeypatch.setattr(cli_selections, "ask_output_language", lambda default=None: "English") + monkeypatch.setattr(cli_selections, "select_analysts", lambda asset_type, default=None: [AnalystType.MARKET]) + monkeypatch.setattr(cli_selections, "select_research_depth", lambda default=None: 1) + monkeypatch.setattr(cli_selections, "select_shallow_thinking_agent", lambda p, default=None: "local-model") + monkeypatch.setattr(cli_selections, "select_deep_thinking_agent", lambda p, default=None: "local-model") - m.get_user_selections() + cli_selections.get_user_selections() assert offered["default"] == "http://localhost:1234/v1" diff --git a/tests/test_cli_symbol_handling.py b/tests/test_cli_symbol_handling.py index 0e0baaefc..5530476e8 100644 --- a/tests/test_cli_symbol_handling.py +++ b/tests/test_cli_symbol_handling.py @@ -5,6 +5,7 @@ stock), #982 (BTC-USDT accepted but unpriceable on Yahoo). """ import pytest +import cli.run as cli_run from cli.models import AssetType from cli.prompts import detect_asset_type, is_valid_ticker_input, normalize_ticker_symbol from tradingagents.dataflows.symbols import normalize_symbol @@ -66,10 +67,9 @@ def test_cli_normalize_delegates_to_data_layer(): def test_the_run_directory_cannot_escape_the_results_directory(tmp_path, monkeypatch): """Every other path that interpolates a ticker validates it first; the CLI's own results tree did not, so a ticker of '..' wrote a level up.""" - import cli.main as m with pytest.raises(ValueError): - m._run_directory({"results_dir": str(tmp_path)}, "..", "2026-09-01") + cli_run._run_directory({"results_dir": str(tmp_path)}, "..", "2026-09-01") - ok = m._run_directory({"results_dir": str(tmp_path)}, "NVDA", "2026-09-01") + ok = cli_run._run_directory({"results_dir": str(tmp_path)}, "NVDA", "2026-09-01") assert str(ok).startswith(str(tmp_path)) diff --git a/tests/test_ollama_base_url.py b/tests/test_ollama_base_url.py index 2d7502115..6cac1f55e 100644 --- a/tests/test_ollama_base_url.py +++ b/tests/test_ollama_base_url.py @@ -23,16 +23,17 @@ def _resync_reloaded_modules(): """Restore module state after this file's importlib.reload() calls. Several tests below reload ``cli.prompts`` to re-evaluate OLLAMA_BASE_URL. - That leaves ``cli.main``'s star-imported names (e.g. get_ticker) bound to - the pre-reload module objects, which breaks identity checks in unrelated - tests that happen to run afterward. Re-sync once on teardown so the reload - doesn't leak across test modules. + That leaves the modules importing from it (cli.selections, then cli.run and + cli.main) bound to the pre-reload functions, which breaks identity checks in + unrelated tests that run afterward. Re-sync them in import order on teardown. """ yield import cli.main import cli.prompts - importlib.reload(cli.prompts) - importlib.reload(cli.main) + import cli.run + import cli.selections + for module in (cli.prompts, cli.selections, cli.run, cli.main): + importlib.reload(module) # ---- openai_client side: registry-driven base_url resolution -------------- diff --git a/tests/test_rating_integrity.py b/tests/test_rating_integrity.py index 78625f9f7..1ee0a4c87 100644 --- a/tests/test_rating_integrity.py +++ b/tests/test_rating_integrity.py @@ -11,6 +11,7 @@ from __future__ import annotations import pytest +import cli.run as cli_run from tradingagents.agents.rating import RATING_REVIEW, extract_rating, parse_rating INVERTED = ("The aggressive analyst pushed hard for a Buy on the AI backlog, but the " @@ -145,23 +146,23 @@ def test_the_cli_says_when_a_run_produced_no_usable_rating(monkeypatch, tmp_path fake = _Graph() fake.graph = fake fake.propagator = fake - monkeypatch.setattr(m, "TradingAgentsGraph", lambda *a, **k: fake) - monkeypatch.setattr(m, "create_layout", lambda: None) - monkeypatch.setattr(m, "update_display", lambda *a, **k: None) - monkeypatch.setattr(m, "Live", type("L", (), {"__init__": lambda s, *a, **k: None, + monkeypatch.setattr(cli_run, "TradingAgentsGraph", lambda *a, **k: fake) + monkeypatch.setattr(cli_run, "create_layout", lambda: None) + monkeypatch.setattr(cli_run, "update_display", lambda *a, **k: None) + monkeypatch.setattr(cli_run, "Live", type("L", (), {"__init__": lambda s, *a, **k: None, "__enter__": lambda s: s, "__exit__": lambda s, *a: False})) monkeypatch.setattr(m.console, "print", lambda *a, **k: printed.append(" ".join(str(x) for x in a))) - monkeypatch.setattr(m, "display_complete_report", lambda *a, **k: None) + monkeypatch.setattr(cli_run, "display_complete_report", lambda *a, **k: None) monkeypatch.setattr(m.typer, "prompt", lambda *a, **k: "N") - monkeypatch.setattr(m, "get_user_selections", lambda: { + monkeypatch.setattr(cli_run, "get_user_selections", lambda: { "ticker": "NVDA", "analysis_date": "2026-01-10", "analysts": [AnalystType.MARKET], "asset_type": "stock", }) - monkeypatch.setattr(m, "_build_run_config", lambda s, c: { + monkeypatch.setattr(cli_run, "_build_run_config", lambda s, c: { "data_cache_dir": str(tmp_path / "c"), "results_dir": str(tmp_path / "r")}) - m.run_analysis() + cli_run.run_analysis() assert any("review" in line.lower() for line in printed), printed[-5:] diff --git a/tests/test_ticker_symbol_handling.py b/tests/test_ticker_symbol_handling.py index dc206d472..2a07e2bef 100644 --- a/tests/test_ticker_symbol_handling.py +++ b/tests/test_ticker_symbol_handling.py @@ -17,12 +17,11 @@ class TickerSymbolHandlingTests(unittest.TestCase): self.assertIn("exchange suffix", context) def test_single_get_ticker_no_shadow(self): - # Regression: cli/main.py had a duplicate get_ticker with an empty - # questionary prompt (rendered as a bare "?") that shadowed the - # descriptive one in cli/prompts. Keep a single canonical definition. - import cli.main + # A second get_ticker with an empty prompt (a bare "?") once shadowed + # the descriptive one; the selection flow must use the one in prompts. import cli.prompts - self.assertIs(cli.main.get_ticker, cli.prompts.get_ticker) + import cli.selections + self.assertIs(cli.selections.get_ticker, cli.prompts.get_ticker) if __name__ == "__main__":