refactor(cli): move the live view to cli/display.py and the prompts to cli/prompts.py

- display.py holds the message buffer, layout, status tables and report panels, the analyst wall-time tracker (CLI-only, from graph/analyst_execution) and the one Console
- utils.py is renamed prompts.py, which is what it holds; its analyst list is ANALYST_CHOICES, apart from display's ANALYST_ORDER
- get_initial_analyst_node, a one-line helper with one caller, is inlined
- the wall-time tracker tests sit with the other display tests, and tests import cli.prompts as prompts
This commit is contained in:
Yijia-Xiao
2026-09-24 05:00:36 +00:00
parent 969861e8be
commit 56bd98f690
19 changed files with 827 additions and 803 deletions
+659
View File
@@ -0,0 +1,659 @@
"""The live view of a run: message log, agent status, report panels and timings."""
import datetime
import time
from collections import deque
from time import monotonic
from rich import box
from rich.console import Console
from rich.layout import Layout
from rich.markdown import Markdown
from rich.panel import Panel
from rich.rule import Rule
from rich.spinner import Spinner
from rich.table import Table
from rich.text import Text
from tradingagents.graph.analyst_execution import (
ANALYST_NODE_SPECS,
AnalystExecutionPlan,
)
# Create a deque to store recent messages with a maximum length
console = Console()
class MessageBuffer:
# Fixed teams that always run (not user-selectable)
FIXED_AGENTS = {
"Research Team": ["Bull Researcher", "Bear Researcher", "Research Manager"],
"Trading Team": ["Trader"],
"Risk Management": ["Aggressive Analyst", "Neutral Analyst", "Conservative Analyst"],
"Portfolio Management": ["Portfolio Manager"],
}
# Analyst name mapping
ANALYST_MAPPING = {
"market": "Market Analyst",
"social": "Sentiment Analyst",
"news": "News Analyst",
"fundamentals": "Fundamentals Analyst",
}
# Report section mapping: section -> (analyst_key for filtering, finalizing_agent)
# analyst_key: which analyst selection controls this section (None = always included)
# finalizing_agent: which agent must be "completed" for this report to count as done
REPORT_SECTIONS = {
"market_report": ("market", "Market Analyst"),
"sentiment_report": ("social", "Sentiment Analyst"),
"news_report": ("news", "News Analyst"),
"fundamentals_report": ("fundamentals", "Fundamentals Analyst"),
"investment_plan": (None, "Research Manager"),
"trader_investment_plan": (None, "Trader"),
"final_trade_decision": (None, "Portfolio Manager"),
}
def __init__(self, max_length=100):
self.messages = deque(maxlen=max_length)
self.tool_calls = deque(maxlen=max_length)
self.current_report = None
self.agent_status = {}
self.report_sections = {}
self.selected_analysts = []
self._processed_message_ids = set()
def init_for_analysis(self, selected_analysts):
"""Initialize agent status and report sections based on selected analysts.
Args:
selected_analysts: List of analyst type strings (e.g., ["market", "news"])
"""
self.selected_analysts = [a.lower() for a in selected_analysts]
# Build agent_status dynamically
self.agent_status = {}
# Add selected analysts
for analyst_key in self.selected_analysts:
if analyst_key in self.ANALYST_MAPPING:
self.agent_status[self.ANALYST_MAPPING[analyst_key]] = "pending"
# Add fixed teams
for team_agents in self.FIXED_AGENTS.values():
for agent in team_agents:
self.agent_status[agent] = "pending"
# Build report_sections dynamically
self.report_sections = {}
for section, (analyst_key, _) in self.REPORT_SECTIONS.items():
if analyst_key is None or analyst_key in self.selected_analysts:
self.report_sections[section] = None
# Reset other state
self.current_report = None
self.messages.clear()
self.tool_calls.clear()
self._processed_message_ids.clear()
def get_completed_reports_count(self):
"""Count reports that are finalized (their finalizing agent is completed).
A report is considered complete when:
1. The report section has content (not None), AND
2. The agent responsible for finalizing that report has status "completed"
This prevents interim updates (like debate rounds) from counting as completed.
"""
count = 0
for section in self.report_sections:
if section not in self.REPORT_SECTIONS:
continue
_, finalizing_agent = self.REPORT_SECTIONS[section]
# Report is complete if it has content AND its finalizing agent is done
has_content = self.report_sections.get(section) is not None
agent_done = self.agent_status.get(finalizing_agent) == "completed"
if has_content and agent_done:
count += 1
return count
def add_message(self, message_type, content):
timestamp = datetime.datetime.now().strftime("%H:%M:%S")
self.messages.append((timestamp, message_type, content))
def add_tool_call(self, tool_name, args):
timestamp = datetime.datetime.now().strftime("%H:%M:%S")
self.tool_calls.append((timestamp, tool_name, args))
def update_agent_status(self, agent, status):
if agent in self.agent_status:
self.agent_status[agent] = status
def update_report_section(self, section_name, content):
if section_name in self.report_sections:
self.report_sections[section_name] = content
self._update_current_report()
def _update_current_report(self):
# For the panel display, only show the most recently updated section
latest_section = None
latest_content = None
# Find the most recently updated section
for section, content in self.report_sections.items():
if content is not None:
latest_section = section
latest_content = content
if latest_section and latest_content:
# Format the current section for display
section_titles = {
"market_report": "Market Analysis",
"sentiment_report": "Social Sentiment",
"news_report": "News Analysis",
"fundamentals_report": "Fundamentals Analysis",
"investment_plan": "Research Team Decision",
"trader_investment_plan": "Trading Team Plan",
"final_trade_decision": "Portfolio Management Decision",
}
self.current_report = (
f"### {section_titles[latest_section]}\n{latest_content}"
)
message_buffer = MessageBuffer()
def create_layout():
layout = Layout()
layout.split_column(
Layout(name="header", size=3),
Layout(name="main"),
Layout(name="footer", size=3),
)
layout["main"].split_column(
Layout(name="upper", ratio=3), Layout(name="analysis", ratio=5)
)
layout["upper"].split_row(
Layout(name="progress", ratio=2), Layout(name="messages", ratio=3)
)
return layout
def format_tokens(n):
"""Format token count for display."""
if n >= 1000:
return f"{n/1000:.1f}k"
return str(n)
def update_display(layout, spinner_text=None, stats_handler=None, start_time=None):
# Header with welcome message
layout["header"].update(
Panel(
"[bold green]Welcome to TradingAgents CLI[/bold green]\n"
"[dim]© [Tauric Research](https://github.com/TauricResearch)[/dim]",
title="Welcome to TradingAgents",
border_style="green",
padding=(1, 2),
expand=True,
)
)
# Progress panel showing agent status
progress_table = Table(
show_header=True,
header_style="bold magenta",
show_footer=False,
box=box.SIMPLE_HEAD, # Use simple header with horizontal lines
title=None, # Remove the redundant Progress title
padding=(0, 2), # Add horizontal padding
expand=True, # Make table expand to fill available space
)
progress_table.add_column("Team", style="cyan", justify="center", width=20)
progress_table.add_column("Agent", style="green", justify="center", width=20)
progress_table.add_column("Status", style="yellow", justify="center", width=20)
# Group agents by team - filter to only include agents in agent_status
all_teams = {
"Analyst Team": [
"Market Analyst",
"Sentiment Analyst",
"News Analyst",
"Fundamentals Analyst",
],
"Research Team": ["Bull Researcher", "Bear Researcher", "Research Manager"],
"Trading Team": ["Trader"],
"Risk Management": ["Aggressive Analyst", "Neutral Analyst", "Conservative Analyst"],
"Portfolio Management": ["Portfolio Manager"],
}
# Filter teams to only include agents that are in agent_status
teams = {}
for team, agents in all_teams.items():
active_agents = [a for a in agents if a in message_buffer.agent_status]
if active_agents:
teams[team] = active_agents
for team, agents in teams.items():
# Add first agent with team name
first_agent = agents[0]
status = message_buffer.agent_status.get(first_agent, "pending")
if status == "in_progress":
spinner = Spinner(
"dots", text="[blue]in_progress[/blue]", style="bold cyan"
)
status_cell = spinner
else:
status_color = {
"pending": "yellow",
"completed": "green",
"error": "red",
}.get(status, "white")
status_cell = f"[{status_color}]{status}[/{status_color}]"
progress_table.add_row(team, first_agent, status_cell)
# Add remaining agents in team
for agent in agents[1:]:
status = message_buffer.agent_status.get(agent, "pending")
if status == "in_progress":
spinner = Spinner(
"dots", text="[blue]in_progress[/blue]", style="bold cyan"
)
status_cell = spinner
else:
status_color = {
"pending": "yellow",
"completed": "green",
"error": "red",
}.get(status, "white")
status_cell = f"[{status_color}]{status}[/{status_color}]"
progress_table.add_row("", agent, status_cell)
# Add horizontal line after each team
progress_table.add_row("─" * 20, "─" * 20, "─" * 20, style="dim")
layout["progress"].update(
Panel(progress_table, title="Progress", border_style="cyan", padding=(1, 2))
)
# Messages panel showing recent messages and tool calls
messages_table = Table(
show_header=True,
header_style="bold magenta",
show_footer=False,
expand=True, # Make table expand to fill available space
box=box.MINIMAL, # Use minimal box style for a lighter look
show_lines=True, # Keep horizontal lines
padding=(0, 1), # Add some padding between columns
)
messages_table.add_column("Time", style="cyan", width=8, justify="center")
messages_table.add_column("Type", style="green", width=10, justify="center")
messages_table.add_column(
"Content", style="white", no_wrap=False, ratio=1
) # Make content column expand
# Combine tool calls and messages
all_messages = []
# Add tool calls
for timestamp, tool_name, args in message_buffer.tool_calls:
formatted_args = format_tool_args(args)
all_messages.append((timestamp, "Tool", f"{tool_name}: {formatted_args}"))
# Add regular messages
for timestamp, msg_type, content in message_buffer.messages:
content_str = str(content) if content else ""
if len(content_str) > 200:
content_str = content_str[:197] + "..."
all_messages.append((timestamp, msg_type, content_str))
# Sort by timestamp descending (newest first)
all_messages.sort(key=lambda x: x[0], reverse=True)
# Calculate how many messages we can show based on available space
max_messages = 12
# Get the first N messages (newest ones)
recent_messages = all_messages[:max_messages]
# Add messages to table (already in newest-first order)
for timestamp, msg_type, content in recent_messages:
# Format content with word wrapping
wrapped_content = Text(content, overflow="fold")
messages_table.add_row(timestamp, msg_type, wrapped_content)
layout["messages"].update(
Panel(
messages_table,
title="Messages & Tools",
border_style="blue",
padding=(1, 2),
)
)
# Analysis panel showing current report
if message_buffer.current_report:
layout["analysis"].update(
Panel(
Markdown(message_buffer.current_report),
title="Current Report",
border_style="green",
padding=(1, 2),
)
)
else:
layout["analysis"].update(
Panel(
"[italic]Waiting for analysis report...[/italic]",
title="Current Report",
border_style="green",
padding=(1, 2),
)
)
# Footer with statistics
# Agent progress - derived from agent_status dict
agents_completed = sum(
1 for status in message_buffer.agent_status.values() if status == "completed"
)
agents_total = len(message_buffer.agent_status)
# Report progress - based on agent completion (not just content existence)
reports_completed = message_buffer.get_completed_reports_count()
reports_total = len(message_buffer.report_sections)
# Build stats parts
stats_parts = [f"Agents: {agents_completed}/{agents_total}"]
# LLM and tool stats from callback handler
if stats_handler:
stats = stats_handler.get_stats()
stats_parts.append(f"LLM: {stats['llm_calls']}")
stats_parts.append(f"Tools: {stats['tool_calls']}")
# Token display with graceful fallback
if stats["tokens_in"] > 0 or stats["tokens_out"] > 0:
tokens_str = f"Tokens: {format_tokens(stats['tokens_in'])}\u2191 {format_tokens(stats['tokens_out'])}\u2193"
else:
tokens_str = "Tokens: --"
stats_parts.append(tokens_str)
stats_parts.append(f"Reports: {reports_completed}/{reports_total}")
# Elapsed time
if start_time:
elapsed = time.time() - start_time
elapsed_str = f"\u23f1 {int(elapsed // 60):02d}:{int(elapsed % 60):02d}"
stats_parts.append(elapsed_str)
stats_table = Table(show_header=False, box=None, padding=(0, 2), expand=True)
stats_table.add_column("Stats", justify="center")
stats_table.add_row(" | ".join(stats_parts))
layout["footer"].update(Panel(stats_table, border_style="grey50"))
def display_complete_report(final_state):
"""Display the complete analysis report sequentially (avoids truncation)."""
console.print()
console.print(Rule("Complete Analysis Report", style="bold green"))
# I. Analyst Team Reports
analysts = []
if final_state.get("market_report"):
analysts.append(("Market Analyst", final_state["market_report"]))
if final_state.get("sentiment_report"):
analysts.append(("Sentiment Analyst", final_state["sentiment_report"]))
if final_state.get("news_report"):
analysts.append(("News Analyst", final_state["news_report"]))
if final_state.get("fundamentals_report"):
analysts.append(("Fundamentals Analyst", final_state["fundamentals_report"]))
if analysts:
console.print(Panel("[bold]I. Analyst Team Reports[/bold]", border_style="cyan"))
for title, content in analysts:
console.print(Panel(Markdown(content), title=title, border_style="blue", padding=(1, 2)))
# II. Research Team Reports
if final_state.get("investment_debate_state"):
debate = final_state["investment_debate_state"]
research = []
if debate.get("bull_history"):
research.append(("Bull Researcher", debate["bull_history"]))
if debate.get("bear_history"):
research.append(("Bear Researcher", debate["bear_history"]))
if debate.get("judge_decision"):
research.append(("Research Manager", debate["judge_decision"]))
if research:
console.print(Panel("[bold]II. Research Team Decision[/bold]", border_style="magenta"))
for title, content in research:
console.print(Panel(Markdown(content), title=title, border_style="blue", padding=(1, 2)))
# III. Trading Team
if final_state.get("trader_investment_plan"):
console.print(Panel("[bold]III. Trading Team Plan[/bold]", border_style="yellow"))
console.print(Panel(Markdown(final_state["trader_investment_plan"]), title="Trader", border_style="blue", padding=(1, 2)))
# IV. Risk Management Team
if final_state.get("risk_debate_state"):
risk = final_state["risk_debate_state"]
risk_reports = []
if risk.get("aggressive_history"):
risk_reports.append(("Aggressive Analyst", risk["aggressive_history"]))
if risk.get("conservative_history"):
risk_reports.append(("Conservative Analyst", risk["conservative_history"]))
if risk.get("neutral_history"):
risk_reports.append(("Neutral Analyst", risk["neutral_history"]))
if risk_reports:
console.print(Panel("[bold]IV. Risk Management Team Decision[/bold]", border_style="red"))
for title, content in risk_reports:
console.print(Panel(Markdown(content), title=title, border_style="blue", padding=(1, 2)))
# V. Portfolio Manager Decision
if risk.get("judge_decision"):
console.print(Panel("[bold]V. Portfolio Manager Decision[/bold]", border_style="green"))
console.print(Panel(Markdown(risk["judge_decision"]), title="Portfolio Manager", border_style="blue", padding=(1, 2)))
def update_research_team_status(status):
"""Update status for research team members (not Trader)."""
research_team = ["Bull Researcher", "Bear Researcher", "Research Manager"]
for agent in research_team:
message_buffer.update_agent_status(agent, status)
# Ordered list of analysts for status transitions
ANALYST_ORDER = ["market", "social", "news", "fundamentals"]
ANALYST_AGENT_NAMES = {
"market": "Market Analyst",
"social": "Sentiment Analyst",
"news": "News Analyst",
"fundamentals": "Fundamentals Analyst",
}
ANALYST_REPORT_MAP = {
"market": "market_report",
"social": "sentiment_report",
"news": "news_report",
"fundamentals": "fundamentals_report",
}
def update_analyst_statuses(message_buffer, chunk, wall_time_tracker=None):
"""Update analyst statuses based on accumulated report state.
Logic:
- Store new report content from the current chunk if present
- Check accumulated report_sections (not just current chunk) for status
- Analysts with reports = completed
- First analyst without report = in_progress
- Remaining analysts without reports = pending
- When all analysts done, set Bull Researcher to in_progress
"""
selected = message_buffer.selected_analysts
found_active = False
if wall_time_tracker is not None:
sync_analyst_tracker_from_chunk(wall_time_tracker, chunk)
for analyst_key in ANALYST_ORDER:
if analyst_key not in selected:
continue
agent_name = ANALYST_AGENT_NAMES[analyst_key]
report_key = ANALYST_REPORT_MAP[analyst_key]
# Capture new report content from current chunk
if chunk.get(report_key):
message_buffer.update_report_section(report_key, chunk[report_key])
# Determine status from accumulated sections, not just current chunk
has_report = bool(message_buffer.report_sections.get(report_key))
if has_report:
message_buffer.update_agent_status(agent_name, "completed")
elif not found_active:
message_buffer.update_agent_status(agent_name, "in_progress")
found_active = True
else:
message_buffer.update_agent_status(agent_name, "pending")
# When all analysts complete, transition research team to in_progress
if (
not found_active
and selected
and message_buffer.agent_status.get("Bull Researcher") == "pending"
):
message_buffer.update_agent_status("Bull Researcher", "in_progress")
def extract_content_string(content):
"""Extract string content from various message formats.
Returns None if no meaningful text content is found.
"""
def is_empty(val):
"""Whether a value carries nothing to show.
Text is judged by whether anything was written, not by what it would
mean as Python: a report saying "0" or "None" is a message the run
produced, and reading it as a falsy literal dropped it from the display.
"""
if isinstance(val, str):
return not val.strip()
return val is None or not bool(val)
if is_empty(content):
return None
if isinstance(content, str):
return content.strip()
if isinstance(content, dict):
text = content.get('text', '')
return text.strip() if not is_empty(text) else None
if isinstance(content, list):
text_parts = [
item.get('text', '').strip() if isinstance(item, dict) and item.get('type') == 'text'
else (item.strip() if isinstance(item, str) else '')
for item in content
]
result = ' '.join(t for t in text_parts if t and not is_empty(t))
return result if result else None
return str(content).strip() if not is_empty(content) else None
def classify_message_type(message) -> tuple[str, str | None]:
"""Classify LangChain message into display type and extract content.
Returns:
(type, content) - type is one of: User, Agent, Data, Control
- content is extracted string or None
"""
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
content = extract_content_string(getattr(message, 'content', None))
if isinstance(message, HumanMessage):
if content and content.strip() == "Continue":
return ("Control", content)
return ("User", content)
if isinstance(message, ToolMessage):
return ("Data", content)
if isinstance(message, AIMessage):
return ("Agent", content)
# Fallback for unknown types
return ("System", content)
def format_tool_args(args, max_length=80) -> str:
"""Format tool arguments for terminal display."""
result = str(args)
if len(result) > max_length:
return result[:max_length - 3] + "..."
return result
class AnalystWallTimeTracker:
def __init__(self, plan: AnalystExecutionPlan):
self.plan = plan
self._started_at: dict[str, float] = {}
self._wall_times: dict[str, float] = {}
def mark_started(self, analyst_key: str, started_at: float | None = None) -> None:
if analyst_key not in ANALYST_NODE_SPECS:
raise ValueError(f"unknown analyst key: {analyst_key}")
self._started_at.setdefault(analyst_key, monotonic() if started_at is None else started_at)
def mark_completed(
self,
analyst_key: str,
completed_at: float | None = None,
) -> None:
if analyst_key not in ANALYST_NODE_SPECS:
raise ValueError(f"unknown analyst key: {analyst_key}")
if analyst_key in self._wall_times:
return
started_at = self._started_at.get(analyst_key)
if started_at is None:
return
finished_at = monotonic() if completed_at is None else completed_at
self._wall_times[analyst_key] = max(0.0, finished_at - started_at)
def format_summary(self) -> str:
parts = []
for spec in self.plan.specs:
duration = self._wall_times.get(spec.key)
if duration is not None:
label = spec.agent_node.removesuffix(" Analyst")
parts.append(f"{label} {duration:.2f}s")
if not parts:
return "Analyst wall time: pending"
return "Analyst wall time: " + " | ".join(parts)
def sync_analyst_tracker_from_chunk(
tracker: AnalystWallTimeTracker,
chunk: dict[str, str],
now: float | None = None,
) -> None:
current_time = monotonic() if now is None else now
active_found = False
for spec in tracker.plan.specs:
has_report = bool(chunk.get(spec.report_key))
if has_report:
tracker.mark_started(spec.key, started_at=current_time)
tracker.mark_completed(spec.key, completed_at=current_time)
continue
if not active_found:
tracker.mark_started(spec.key, started_at=current_time)
active_found = True
+15 -581
View File
@@ -2,27 +2,29 @@ import datetime
import os import os
import sys import sys
import time import time
from collections import deque
from functools import wraps from functools import wraps
from pathlib import Path from pathlib import Path
import typer import typer
from rich import box
from rich.align import Align from rich.align import Align
from rich.console import Console
from rich.layout import Layout
from rich.live import Live from rich.live import Live
from rich.markdown import Markdown
from rich.panel import Panel from rich.panel import Panel
from rich.rule import Rule
from rich.spinner import Spinner
from rich.table import Table
from rich.text import Text
from cli.announcements import display_announcements, fetch_announcements 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.prefs import load_last_run, sanitize, save_last_run
from cli.stats_handler import StatsCallbackHandler from cli.prompts import (
from cli.utils import (
ask_anthropic_effort, ask_anthropic_effort,
ask_gemini_thinking_config, ask_gemini_thinking_config,
ask_glm_region, ask_glm_region,
@@ -42,22 +44,18 @@ from cli.utils import (
select_research_depth, select_research_depth,
select_shallow_thinking_agent, select_shallow_thinking_agent,
) )
from cli.stats_handler import StatsCallbackHandler
from tradingagents.agents.rating import is_review from tradingagents.agents.rating import is_review
from tradingagents.backtest import iter_grid, run_backtest, summarize from tradingagents.backtest import iter_grid, run_backtest, summarize
from tradingagents.dataflows.symbols import safe_ticker_component from tradingagents.dataflows.symbols import safe_ticker_component
from tradingagents.default_config import DEFAULT_CONFIG from tradingagents.default_config import DEFAULT_CONFIG
from tradingagents.graph.analyst_execution import ( from tradingagents.graph.analyst_execution import (
AnalystWallTimeTracker,
build_analyst_execution_plan, build_analyst_execution_plan,
get_initial_analyst_node,
sync_analyst_tracker_from_chunk,
) )
from tradingagents.graph.trading_graph import TradingAgentsGraph from tradingagents.graph.trading_graph import TradingAgentsGraph
from tradingagents.portfolio import load_portfolio from tradingagents.portfolio import load_portfolio
from tradingagents.reporting import write_report_tree from tradingagents.reporting import write_report_tree
console = Console()
# prompt_toolkit's win32 output module is importable only on Windows (it asserts # 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 # the platform at import time), so gate on the platform rather than catching the
# failure — that way a genuinely broken prompt_toolkit on Windows still surfaces # failure — that way a genuinely broken prompt_toolkit on Windows still surfaces
@@ -77,375 +75,8 @@ app = typer.Typer(
) )
# Create a deque to store recent messages with a maximum length
class MessageBuffer:
# Fixed teams that always run (not user-selectable)
FIXED_AGENTS = {
"Research Team": ["Bull Researcher", "Bear Researcher", "Research Manager"],
"Trading Team": ["Trader"],
"Risk Management": ["Aggressive Analyst", "Neutral Analyst", "Conservative Analyst"],
"Portfolio Management": ["Portfolio Manager"],
}
# Analyst name mapping
ANALYST_MAPPING = {
"market": "Market Analyst",
"social": "Sentiment Analyst",
"news": "News Analyst",
"fundamentals": "Fundamentals Analyst",
}
# Report section mapping: section -> (analyst_key for filtering, finalizing_agent)
# analyst_key: which analyst selection controls this section (None = always included)
# finalizing_agent: which agent must be "completed" for this report to count as done
REPORT_SECTIONS = {
"market_report": ("market", "Market Analyst"),
"sentiment_report": ("social", "Sentiment Analyst"),
"news_report": ("news", "News Analyst"),
"fundamentals_report": ("fundamentals", "Fundamentals Analyst"),
"investment_plan": (None, "Research Manager"),
"trader_investment_plan": (None, "Trader"),
"final_trade_decision": (None, "Portfolio Manager"),
}
def __init__(self, max_length=100):
self.messages = deque(maxlen=max_length)
self.tool_calls = deque(maxlen=max_length)
self.current_report = None
self.agent_status = {}
self.report_sections = {}
self.selected_analysts = []
self._processed_message_ids = set()
def init_for_analysis(self, selected_analysts):
"""Initialize agent status and report sections based on selected analysts.
Args:
selected_analysts: List of analyst type strings (e.g., ["market", "news"])
"""
self.selected_analysts = [a.lower() for a in selected_analysts]
# Build agent_status dynamically
self.agent_status = {}
# Add selected analysts
for analyst_key in self.selected_analysts:
if analyst_key in self.ANALYST_MAPPING:
self.agent_status[self.ANALYST_MAPPING[analyst_key]] = "pending"
# Add fixed teams
for team_agents in self.FIXED_AGENTS.values():
for agent in team_agents:
self.agent_status[agent] = "pending"
# Build report_sections dynamically
self.report_sections = {}
for section, (analyst_key, _) in self.REPORT_SECTIONS.items():
if analyst_key is None or analyst_key in self.selected_analysts:
self.report_sections[section] = None
# Reset other state
self.current_report = None
self.messages.clear()
self.tool_calls.clear()
self._processed_message_ids.clear()
def get_completed_reports_count(self):
"""Count reports that are finalized (their finalizing agent is completed).
A report is considered complete when:
1. The report section has content (not None), AND
2. The agent responsible for finalizing that report has status "completed"
This prevents interim updates (like debate rounds) from counting as completed.
"""
count = 0
for section in self.report_sections:
if section not in self.REPORT_SECTIONS:
continue
_, finalizing_agent = self.REPORT_SECTIONS[section]
# Report is complete if it has content AND its finalizing agent is done
has_content = self.report_sections.get(section) is not None
agent_done = self.agent_status.get(finalizing_agent) == "completed"
if has_content and agent_done:
count += 1
return count
def add_message(self, message_type, content):
timestamp = datetime.datetime.now().strftime("%H:%M:%S")
self.messages.append((timestamp, message_type, content))
def add_tool_call(self, tool_name, args):
timestamp = datetime.datetime.now().strftime("%H:%M:%S")
self.tool_calls.append((timestamp, tool_name, args))
def update_agent_status(self, agent, status):
if agent in self.agent_status:
self.agent_status[agent] = status
def update_report_section(self, section_name, content):
if section_name in self.report_sections:
self.report_sections[section_name] = content
self._update_current_report()
def _update_current_report(self):
# For the panel display, only show the most recently updated section
latest_section = None
latest_content = None
# Find the most recently updated section
for section, content in self.report_sections.items():
if content is not None:
latest_section = section
latest_content = content
if latest_section and latest_content:
# Format the current section for display
section_titles = {
"market_report": "Market Analysis",
"sentiment_report": "Social Sentiment",
"news_report": "News Analysis",
"fundamentals_report": "Fundamentals Analysis",
"investment_plan": "Research Team Decision",
"trader_investment_plan": "Trading Team Plan",
"final_trade_decision": "Portfolio Management Decision",
}
self.current_report = (
f"### {section_titles[latest_section]}\n{latest_content}"
)
message_buffer = MessageBuffer()
def create_layout():
layout = Layout()
layout.split_column(
Layout(name="header", size=3),
Layout(name="main"),
Layout(name="footer", size=3),
)
layout["main"].split_column(
Layout(name="upper", ratio=3), Layout(name="analysis", ratio=5)
)
layout["upper"].split_row(
Layout(name="progress", ratio=2), Layout(name="messages", ratio=3)
)
return layout
def format_tokens(n):
"""Format token count for display."""
if n >= 1000:
return f"{n/1000:.1f}k"
return str(n)
def update_display(layout, spinner_text=None, stats_handler=None, start_time=None):
# Header with welcome message
layout["header"].update(
Panel(
"[bold green]Welcome to TradingAgents CLI[/bold green]\n"
"[dim]© [Tauric Research](https://github.com/TauricResearch)[/dim]",
title="Welcome to TradingAgents",
border_style="green",
padding=(1, 2),
expand=True,
)
)
# Progress panel showing agent status
progress_table = Table(
show_header=True,
header_style="bold magenta",
show_footer=False,
box=box.SIMPLE_HEAD, # Use simple header with horizontal lines
title=None, # Remove the redundant Progress title
padding=(0, 2), # Add horizontal padding
expand=True, # Make table expand to fill available space
)
progress_table.add_column("Team", style="cyan", justify="center", width=20)
progress_table.add_column("Agent", style="green", justify="center", width=20)
progress_table.add_column("Status", style="yellow", justify="center", width=20)
# Group agents by team - filter to only include agents in agent_status
all_teams = {
"Analyst Team": [
"Market Analyst",
"Sentiment Analyst",
"News Analyst",
"Fundamentals Analyst",
],
"Research Team": ["Bull Researcher", "Bear Researcher", "Research Manager"],
"Trading Team": ["Trader"],
"Risk Management": ["Aggressive Analyst", "Neutral Analyst", "Conservative Analyst"],
"Portfolio Management": ["Portfolio Manager"],
}
# Filter teams to only include agents that are in agent_status
teams = {}
for team, agents in all_teams.items():
active_agents = [a for a in agents if a in message_buffer.agent_status]
if active_agents:
teams[team] = active_agents
for team, agents in teams.items():
# Add first agent with team name
first_agent = agents[0]
status = message_buffer.agent_status.get(first_agent, "pending")
if status == "in_progress":
spinner = Spinner(
"dots", text="[blue]in_progress[/blue]", style="bold cyan"
)
status_cell = spinner
else:
status_color = {
"pending": "yellow",
"completed": "green",
"error": "red",
}.get(status, "white")
status_cell = f"[{status_color}]{status}[/{status_color}]"
progress_table.add_row(team, first_agent, status_cell)
# Add remaining agents in team
for agent in agents[1:]:
status = message_buffer.agent_status.get(agent, "pending")
if status == "in_progress":
spinner = Spinner(
"dots", text="[blue]in_progress[/blue]", style="bold cyan"
)
status_cell = spinner
else:
status_color = {
"pending": "yellow",
"completed": "green",
"error": "red",
}.get(status, "white")
status_cell = f"[{status_color}]{status}[/{status_color}]"
progress_table.add_row("", agent, status_cell)
# Add horizontal line after each team
progress_table.add_row("─" * 20, "─" * 20, "─" * 20, style="dim")
layout["progress"].update(
Panel(progress_table, title="Progress", border_style="cyan", padding=(1, 2))
)
# Messages panel showing recent messages and tool calls
messages_table = Table(
show_header=True,
header_style="bold magenta",
show_footer=False,
expand=True, # Make table expand to fill available space
box=box.MINIMAL, # Use minimal box style for a lighter look
show_lines=True, # Keep horizontal lines
padding=(0, 1), # Add some padding between columns
)
messages_table.add_column("Time", style="cyan", width=8, justify="center")
messages_table.add_column("Type", style="green", width=10, justify="center")
messages_table.add_column(
"Content", style="white", no_wrap=False, ratio=1
) # Make content column expand
# Combine tool calls and messages
all_messages = []
# Add tool calls
for timestamp, tool_name, args in message_buffer.tool_calls:
formatted_args = format_tool_args(args)
all_messages.append((timestamp, "Tool", f"{tool_name}: {formatted_args}"))
# Add regular messages
for timestamp, msg_type, content in message_buffer.messages:
content_str = str(content) if content else ""
if len(content_str) > 200:
content_str = content_str[:197] + "..."
all_messages.append((timestamp, msg_type, content_str))
# Sort by timestamp descending (newest first)
all_messages.sort(key=lambda x: x[0], reverse=True)
# Calculate how many messages we can show based on available space
max_messages = 12
# Get the first N messages (newest ones)
recent_messages = all_messages[:max_messages]
# Add messages to table (already in newest-first order)
for timestamp, msg_type, content in recent_messages:
# Format content with word wrapping
wrapped_content = Text(content, overflow="fold")
messages_table.add_row(timestamp, msg_type, wrapped_content)
layout["messages"].update(
Panel(
messages_table,
title="Messages & Tools",
border_style="blue",
padding=(1, 2),
)
)
# Analysis panel showing current report
if message_buffer.current_report:
layout["analysis"].update(
Panel(
Markdown(message_buffer.current_report),
title="Current Report",
border_style="green",
padding=(1, 2),
)
)
else:
layout["analysis"].update(
Panel(
"[italic]Waiting for analysis report...[/italic]",
title="Current Report",
border_style="green",
padding=(1, 2),
)
)
# Footer with statistics
# Agent progress - derived from agent_status dict
agents_completed = sum(
1 for status in message_buffer.agent_status.values() if status == "completed"
)
agents_total = len(message_buffer.agent_status)
# Report progress - based on agent completion (not just content existence)
reports_completed = message_buffer.get_completed_reports_count()
reports_total = len(message_buffer.report_sections)
# Build stats parts
stats_parts = [f"Agents: {agents_completed}/{agents_total}"]
# LLM and tool stats from callback handler
if stats_handler:
stats = stats_handler.get_stats()
stats_parts.append(f"LLM: {stats['llm_calls']}")
stats_parts.append(f"Tools: {stats['tool_calls']}")
# Token display with graceful fallback
if stats["tokens_in"] > 0 or stats["tokens_out"] > 0:
tokens_str = f"Tokens: {format_tokens(stats['tokens_in'])}\u2191 {format_tokens(stats['tokens_out'])}\u2193"
else:
tokens_str = "Tokens: --"
stats_parts.append(tokens_str)
stats_parts.append(f"Reports: {reports_completed}/{reports_total}")
# Elapsed time
if start_time:
elapsed = time.time() - start_time
elapsed_str = f"\u23f1 {int(elapsed // 60):02d}:{int(elapsed % 60):02d}"
stats_parts.append(elapsed_str)
stats_table = Table(show_header=False, box=None, padding=(0, 2), expand=True)
stats_table.add_column("Stats", justify="center")
stats_table.add_row(" | ".join(stats_parts))
layout["footer"].update(Panel(stats_table, border_style="grey50"))
def get_user_selections(): def get_user_selections():
@@ -731,207 +362,10 @@ def get_analysis_date():
) )
def display_complete_report(final_state):
"""Display the complete analysis report sequentially (avoids truncation)."""
console.print()
console.print(Rule("Complete Analysis Report", style="bold green"))
# I. Analyst Team Reports
analysts = []
if final_state.get("market_report"):
analysts.append(("Market Analyst", final_state["market_report"]))
if final_state.get("sentiment_report"):
analysts.append(("Sentiment Analyst", final_state["sentiment_report"]))
if final_state.get("news_report"):
analysts.append(("News Analyst", final_state["news_report"]))
if final_state.get("fundamentals_report"):
analysts.append(("Fundamentals Analyst", final_state["fundamentals_report"]))
if analysts:
console.print(Panel("[bold]I. Analyst Team Reports[/bold]", border_style="cyan"))
for title, content in analysts:
console.print(Panel(Markdown(content), title=title, border_style="blue", padding=(1, 2)))
# II. Research Team Reports
if final_state.get("investment_debate_state"):
debate = final_state["investment_debate_state"]
research = []
if debate.get("bull_history"):
research.append(("Bull Researcher", debate["bull_history"]))
if debate.get("bear_history"):
research.append(("Bear Researcher", debate["bear_history"]))
if debate.get("judge_decision"):
research.append(("Research Manager", debate["judge_decision"]))
if research:
console.print(Panel("[bold]II. Research Team Decision[/bold]", border_style="magenta"))
for title, content in research:
console.print(Panel(Markdown(content), title=title, border_style="blue", padding=(1, 2)))
# III. Trading Team
if final_state.get("trader_investment_plan"):
console.print(Panel("[bold]III. Trading Team Plan[/bold]", border_style="yellow"))
console.print(Panel(Markdown(final_state["trader_investment_plan"]), title="Trader", border_style="blue", padding=(1, 2)))
# IV. Risk Management Team
if final_state.get("risk_debate_state"):
risk = final_state["risk_debate_state"]
risk_reports = []
if risk.get("aggressive_history"):
risk_reports.append(("Aggressive Analyst", risk["aggressive_history"]))
if risk.get("conservative_history"):
risk_reports.append(("Conservative Analyst", risk["conservative_history"]))
if risk.get("neutral_history"):
risk_reports.append(("Neutral Analyst", risk["neutral_history"]))
if risk_reports:
console.print(Panel("[bold]IV. Risk Management Team Decision[/bold]", border_style="red"))
for title, content in risk_reports:
console.print(Panel(Markdown(content), title=title, border_style="blue", padding=(1, 2)))
# V. Portfolio Manager Decision
if risk.get("judge_decision"):
console.print(Panel("[bold]V. Portfolio Manager Decision[/bold]", border_style="green"))
console.print(Panel(Markdown(risk["judge_decision"]), title="Portfolio Manager", border_style="blue", padding=(1, 2)))
def update_research_team_status(status):
"""Update status for research team members (not Trader)."""
research_team = ["Bull Researcher", "Bear Researcher", "Research Manager"]
for agent in research_team:
message_buffer.update_agent_status(agent, status)
# Ordered list of analysts for status transitions
ANALYST_ORDER = ["market", "social", "news", "fundamentals"]
ANALYST_AGENT_NAMES = {
"market": "Market Analyst",
"social": "Sentiment Analyst",
"news": "News Analyst",
"fundamentals": "Fundamentals Analyst",
}
ANALYST_REPORT_MAP = {
"market": "market_report",
"social": "sentiment_report",
"news": "news_report",
"fundamentals": "fundamentals_report",
}
def update_analyst_statuses(message_buffer, chunk, wall_time_tracker=None):
"""Update analyst statuses based on accumulated report state.
Logic:
- Store new report content from the current chunk if present
- Check accumulated report_sections (not just current chunk) for status
- Analysts with reports = completed
- First analyst without report = in_progress
- Remaining analysts without reports = pending
- When all analysts done, set Bull Researcher to in_progress
"""
selected = message_buffer.selected_analysts
found_active = False
if wall_time_tracker is not None:
sync_analyst_tracker_from_chunk(wall_time_tracker, chunk)
for analyst_key in ANALYST_ORDER:
if analyst_key not in selected:
continue
agent_name = ANALYST_AGENT_NAMES[analyst_key]
report_key = ANALYST_REPORT_MAP[analyst_key]
# Capture new report content from current chunk
if chunk.get(report_key):
message_buffer.update_report_section(report_key, chunk[report_key])
# Determine status from accumulated sections, not just current chunk
has_report = bool(message_buffer.report_sections.get(report_key))
if has_report:
message_buffer.update_agent_status(agent_name, "completed")
elif not found_active:
message_buffer.update_agent_status(agent_name, "in_progress")
found_active = True
else:
message_buffer.update_agent_status(agent_name, "pending")
# When all analysts complete, transition research team to in_progress
if (
not found_active
and selected
and message_buffer.agent_status.get("Bull Researcher") == "pending"
):
message_buffer.update_agent_status("Bull Researcher", "in_progress")
def extract_content_string(content):
"""Extract string content from various message formats.
Returns None if no meaningful text content is found.
"""
def is_empty(val):
"""Whether a value carries nothing to show.
Text is judged by whether anything was written, not by what it would
mean as Python: a report saying "0" or "None" is a message the run
produced, and reading it as a falsy literal dropped it from the display.
"""
if isinstance(val, str):
return not val.strip()
return val is None or not bool(val)
if is_empty(content):
return None
if isinstance(content, str):
return content.strip()
if isinstance(content, dict):
text = content.get('text', '')
return text.strip() if not is_empty(text) else None
if isinstance(content, list):
text_parts = [
item.get('text', '').strip() if isinstance(item, dict) and item.get('type') == 'text'
else (item.strip() if isinstance(item, str) else '')
for item in content
]
result = ' '.join(t for t in text_parts if t and not is_empty(t))
return result if result else None
return str(content).strip() if not is_empty(content) else None
def classify_message_type(message) -> tuple[str, str | None]:
"""Classify LangChain message into display type and extract content.
Returns:
(type, content) - type is one of: User, Agent, Data, Control
- content is extracted string or None
"""
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
content = extract_content_string(getattr(message, 'content', None))
if isinstance(message, HumanMessage):
if content and content.strip() == "Continue":
return ("Control", content)
return ("User", content)
if isinstance(message, ToolMessage):
return ("Data", content)
if isinstance(message, AIMessage):
return ("Agent", content)
# Fallback for unknown types
return ("System", content)
def format_tool_args(args, max_length=80) -> str:
"""Format tool arguments for terminal display."""
result = str(args)
if len(result) > max_length:
return result[:max_length - 3] + "..."
return result
def _run_directory(config: dict, ticker: str, trade_date: str) -> Path: def _run_directory(config: dict, ticker: str, trade_date: str) -> Path:
"""Where this run writes, with the ticker validated as a path component. """Where this run writes, with the ticker validated as a path component.
@@ -1092,7 +526,7 @@ def run_analysis(checkpoint: bool | None = None, portfolio=None):
update_display(layout, stats_handler=stats_handler, start_time=start_time) update_display(layout, stats_handler=stats_handler, start_time=start_time)
# Update agent status to in_progress for the first analyst # Update agent status to in_progress for the first analyst
first_analyst = get_initial_analyst_node(analyst_execution_plan) first_analyst = analyst_execution_plan.specs[0].agent_node
message_buffer.update_agent_status(first_analyst, "in_progress") message_buffer.update_agent_status(first_analyst, "in_progress")
analyst_wall_time_tracker.mark_started(selected_analyst_keys[0]) analyst_wall_time_tracker.mark_started(selected_analyst_keys[0])
update_display(layout, stats_handler=stats_handler, start_time=start_time) update_display(layout, stats_handler=stats_handler, start_time=start_time)
+1 -1
View File
@@ -21,7 +21,7 @@ import os
from pathlib import Path from pathlib import Path
from cli.models import AnalystType, AssetType from cli.models import AnalystType, AssetType
from cli.utils import _llm_provider_table, filter_analysts_for_asset_type from cli.prompts import _llm_provider_table, filter_analysts_for_asset_type
from tradingagents.llm_clients.model_catalog import get_model_options from tradingagents.llm_clients.model_catalog import get_model_options
_PREFS_PATH = Path(os.path.expanduser("~")) / ".tradingagents" / "cli_prefs.json" _PREFS_PATH = Path(os.path.expanduser("~")) / ".tradingagents" / "cli_prefs.json"
+4 -6
View File
@@ -3,17 +3,15 @@ from pathlib import Path
import questionary import questionary
from dotenv import find_dotenv, set_key from dotenv import find_dotenv, set_key
from rich.console import Console
from cli.display import console
from cli.models import AnalystType, AssetType from cli.models import AnalystType, AssetType
from tradingagents.llm_clients.api_key_env import get_api_key_env from tradingagents.llm_clients.api_key_env import get_api_key_env
from tradingagents.llm_clients.model_catalog import get_model_options from tradingagents.llm_clients.model_catalog import get_model_options
console = Console()
TICKER_INPUT_EXAMPLES = "SPY, 0700.HK, BTC-USD" TICKER_INPUT_EXAMPLES = "SPY, 0700.HK, BTC-USD"
ANALYST_ORDER = [ ANALYST_CHOICES = [
("Market Analyst", AnalystType.MARKET), ("Market Analyst", AnalystType.MARKET),
("Sentiment Analyst", AnalystType.SOCIAL), ("Sentiment Analyst", AnalystType.SOCIAL),
("News Analyst", AnalystType.NEWS), ("News Analyst", AnalystType.NEWS),
@@ -110,14 +108,14 @@ def select_analysts(asset_type: AssetType = AssetType.STOCK, default=None) -> li
``default`` pre-checks the previous run's analysts; the prompt still shows. ``default`` pre-checks the previous run's analysts; the prompt still shows.
""" """
available_analysts = filter_analysts_for_asset_type( available_analysts = filter_analysts_for_asset_type(
[value for _, value in ANALYST_ORDER], [value for _, value in ANALYST_CHOICES],
asset_type, asset_type,
) )
choices = questionary.checkbox( choices = questionary.checkbox(
"Select Your [Analysts Team]:", "Select Your [Analysts Team]:",
choices=[ choices=[
questionary.Choice(display, value=value, checked=value.value in (default or [])) questionary.Choice(display, value=value, checked=value.value in (default or []))
for display, value in ANALYST_ORDER for display, value in ANALYST_CHOICES
if value in available_analysts if value in available_analysts
], ],
instruction="\n- Press Space to select/unselect analysts\n- Press 'a' to select/unselect all\n- Press Enter when done", instruction="\n- Press Space to select/unselect analysts\n- Press 'a' to select/unselect all\n- Press Enter when done",
-57
View File
@@ -1,10 +1,7 @@
import unittest import unittest
from tradingagents.graph.analyst_execution import ( from tradingagents.graph.analyst_execution import (
AnalystWallTimeTracker,
build_analyst_execution_plan, build_analyst_execution_plan,
get_initial_analyst_node,
sync_analyst_tracker_from_chunk,
) )
@@ -21,14 +18,6 @@ class AnalystExecutionPlanTests(unittest.TestCase):
with self.assertRaises(ValueError): with self.assertRaises(ValueError):
build_analyst_execution_plan(["market", "macro"]) build_analyst_execution_plan(["market", "macro"])
def test_get_initial_analyst_node_uses_plan_metadata(self):
plan = build_analyst_execution_plan(["fundamentals", "news"])
self.assertEqual(
get_initial_analyst_node(plan),
"Fundamentals Analyst",
)
def test_social_key_displays_as_sentiment_analyst(self): def test_social_key_displays_as_sentiment_analyst(self):
# The wire key stays "social" for saved-config back-compat, but the # The wire key stays "social" for saved-config back-compat, but the
# user-visible agent_node label must match the v0.2.5 rename so the # user-visible agent_node label must match the v0.2.5 rename so the
@@ -39,49 +28,3 @@ class AnalystExecutionPlanTests(unittest.TestCase):
self.assertEqual(spec.key, "social") self.assertEqual(spec.key, "social")
self.assertEqual(spec.agent_node, "Sentiment Analyst") self.assertEqual(spec.agent_node, "Sentiment Analyst")
self.assertEqual(spec.report_key, "sentiment_report") self.assertEqual(spec.report_key, "sentiment_report")
class AnalystWallTimeTrackerTests(unittest.TestCase):
def test_records_wall_time_when_analyst_completes(self):
plan = build_analyst_execution_plan(["market", "news"])
tracker = AnalystWallTimeTracker(plan)
tracker.mark_started("market", started_at=10.0)
tracker.mark_completed("market", completed_at=13.5)
self.assertEqual(tracker.format_summary(), "Analyst wall time: Market 3.50s")
def test_formats_summary_in_plan_order(self):
plan = build_analyst_execution_plan(["news", "market"])
tracker = AnalystWallTimeTracker(plan)
tracker.mark_started("market", started_at=20.0)
tracker.mark_completed("market", completed_at=22.25)
tracker.mark_started("news", started_at=10.0)
tracker.mark_completed("news", completed_at=14.0)
self.assertEqual(
tracker.format_summary(),
"Analyst wall time: News 4.00s | Market 2.25s",
)
def test_syncs_wall_time_from_sequential_chunks(self):
plan = build_analyst_execution_plan(["market", "news"])
tracker = AnalystWallTimeTracker(plan)
sync_analyst_tracker_from_chunk(tracker, {}, now=10.0)
self.assertEqual(tracker.format_summary(), "Analyst wall time: pending")
sync_analyst_tracker_from_chunk(
tracker,
{"market_report": "done"},
now=13.0,
)
self.assertEqual(tracker.format_summary(), "Analyst wall time: Market 3.00s")
sync_analyst_tracker_from_chunk(
tracker,
{"market_report": "done", "news_report": "done"},
now=18.0,
)
self.assertEqual(tracker.format_summary(), "Analyst wall time: Market 3.00s | News 5.00s")
+32 -32
View File
@@ -15,7 +15,7 @@ from tradingagents.llm_clients.api_key_env import PROVIDER_API_KEY_ENV, get_api_
def test_every_select_llm_provider_choice_has_an_entry(): def test_every_select_llm_provider_choice_has_an_entry():
"""select_llm_provider() must not present a provider the mapping doesn't know about.""" """select_llm_provider() must not present a provider the mapping doesn't know about."""
# Mirrors the dropdown order in cli/utils.select_llm_provider so the two # Mirrors the dropdown order in cli/prompts.select_llm_provider so the two
# stay in lockstep. Region-specific keys (qwen-cn / minimax-cn / glm-cn) # stay in lockstep. Region-specific keys (qwen-cn / minimax-cn / glm-cn)
# are reached via the secondary region prompt, so they must also be present. # are reached via the secondary region prompt, so they must also be present.
expected = { expected = {
@@ -67,44 +67,44 @@ def test_case_insensitive_lookup():
@pytest.fixture @pytest.fixture
def cli_utils(monkeypatch): def prompts(monkeypatch):
"""Import cli.utils with a fresh environment so module-level state is consistent.""" """Import cli.prompts with a fresh environment so module-level state is consistent."""
import importlib import importlib
import cli.utils as cli_utils_module import cli.prompts as prompts_module
return importlib.reload(cli_utils_module) return importlib.reload(prompts_module)
def test_ensure_api_key_returns_existing(monkeypatch, cli_utils): def test_ensure_api_key_returns_existing(monkeypatch, prompts):
monkeypatch.setenv("OPENAI_API_KEY", "sk-already-set") monkeypatch.setenv("OPENAI_API_KEY", "sk-already-set")
result = cli_utils.ensure_api_key("openai") result = prompts.ensure_api_key("openai")
assert result == "sk-already-set" assert result == "sk-already-set"
def test_ensure_api_key_no_op_for_ollama(monkeypatch, cli_utils): def test_ensure_api_key_no_op_for_ollama(monkeypatch, prompts):
# Even with no env var set, ollama should not prompt and should return None. # Even with no env var set, ollama should not prompt and should return None.
monkeypatch.delenv("OPENAI_API_KEY", raising=False) monkeypatch.delenv("OPENAI_API_KEY", raising=False)
with patch.object(cli_utils, "questionary") as mock_q: with patch.object(prompts, "questionary") as mock_q:
result = cli_utils.ensure_api_key("ollama") result = prompts.ensure_api_key("ollama")
assert result is None assert result is None
mock_q.password.assert_not_called() mock_q.password.assert_not_called()
def test_ensure_api_key_unknown_provider_no_prompt(monkeypatch, cli_utils): def test_ensure_api_key_unknown_provider_no_prompt(monkeypatch, prompts):
with patch.object(cli_utils, "questionary") as mock_q: with patch.object(prompts, "questionary") as mock_q:
result = cli_utils.ensure_api_key("totally-fake-provider") result = prompts.ensure_api_key("totally-fake-provider")
assert result is None assert result is None
mock_q.password.assert_not_called() mock_q.password.assert_not_called()
def test_ensure_api_key_prompts_and_writes_to_env(monkeypatch, tmp_path, cli_utils): def test_ensure_api_key_prompts_and_writes_to_env(monkeypatch, tmp_path, prompts):
"""When key is missing, user-pasted value must be written to .env AND os.environ.""" """When key is missing, user-pasted value must be written to .env AND os.environ."""
monkeypatch.delenv("DEEPSEEK_API_KEY", raising=False) monkeypatch.delenv("DEEPSEEK_API_KEY", raising=False)
monkeypatch.chdir(tmp_path) monkeypatch.chdir(tmp_path)
fake_prompt = type("P", (), {"ask": staticmethod(lambda: "sk-deepseek-test")})() fake_prompt = type("P", (), {"ask": staticmethod(lambda: "sk-deepseek-test")})()
with patch.object(cli_utils.questionary, "password", return_value=fake_prompt): with patch.object(prompts.questionary, "password", return_value=fake_prompt):
result = cli_utils.ensure_api_key("deepseek") result = prompts.ensure_api_key("deepseek")
assert result == "sk-deepseek-test" assert result == "sk-deepseek-test"
assert os.environ["DEEPSEEK_API_KEY"] == "sk-deepseek-test" assert os.environ["DEEPSEEK_API_KEY"] == "sk-deepseek-test"
@@ -114,14 +114,14 @@ def test_ensure_api_key_prompts_and_writes_to_env(monkeypatch, tmp_path, cli_uti
assert "sk-deepseek-test" in env_file.read_text() assert "sk-deepseek-test" in env_file.read_text()
def test_ensure_api_key_user_cancels_returns_none(monkeypatch, tmp_path, cli_utils): def test_ensure_api_key_user_cancels_returns_none(monkeypatch, tmp_path, prompts):
"""Empty prompt response (user cancelled) must not write to .env.""" """Empty prompt response (user cancelled) must not write to .env."""
monkeypatch.delenv("XAI_API_KEY", raising=False) monkeypatch.delenv("XAI_API_KEY", raising=False)
monkeypatch.chdir(tmp_path) monkeypatch.chdir(tmp_path)
fake_prompt = type("P", (), {"ask": staticmethod(lambda: None)})() fake_prompt = type("P", (), {"ask": staticmethod(lambda: None)})()
with patch.object(cli_utils.questionary, "password", return_value=fake_prompt): with patch.object(prompts.questionary, "password", return_value=fake_prompt):
result = cli_utils.ensure_api_key("xai") result = prompts.ensure_api_key("xai")
assert result is None assert result is None
assert "XAI_API_KEY" not in os.environ assert "XAI_API_KEY" not in os.environ
@@ -132,7 +132,7 @@ def test_ensure_api_key_user_cancels_returns_none(monkeypatch, tmp_path, cli_uti
assert "XAI_API_KEY" not in env_file.read_text() assert "XAI_API_KEY" not in env_file.read_text()
def test_ensure_api_key_updates_existing_env_file(monkeypatch, tmp_path, cli_utils): def test_ensure_api_key_updates_existing_env_file(monkeypatch, tmp_path, prompts):
"""An existing .env with other keys must be preserved on writeback.""" """An existing .env with other keys must be preserved on writeback."""
monkeypatch.delenv("OPENROUTER_API_KEY", raising=False) monkeypatch.delenv("OPENROUTER_API_KEY", raising=False)
monkeypatch.chdir(tmp_path) monkeypatch.chdir(tmp_path)
@@ -140,8 +140,8 @@ def test_ensure_api_key_updates_existing_env_file(monkeypatch, tmp_path, cli_uti
env_file.write_text("OPENAI_API_KEY=sk-existing\nOTHER=value\n") env_file.write_text("OPENAI_API_KEY=sk-existing\nOTHER=value\n")
fake_prompt = type("P", (), {"ask": staticmethod(lambda: "sk-openrouter-new")})() fake_prompt = type("P", (), {"ask": staticmethod(lambda: "sk-openrouter-new")})()
with patch.object(cli_utils.questionary, "password", return_value=fake_prompt): with patch.object(prompts.questionary, "password", return_value=fake_prompt):
cli_utils.ensure_api_key("openrouter") prompts.ensure_api_key("openrouter")
content = env_file.read_text() content = env_file.read_text()
assert "OPENAI_API_KEY" in content and "sk-existing" in content assert "OPENAI_API_KEY" in content and "sk-existing" in content
@@ -149,22 +149,22 @@ def test_ensure_api_key_updates_existing_env_file(monkeypatch, tmp_path, cli_uti
assert "OPENROUTER_API_KEY" in content and "sk-openrouter-new" in content assert "OPENROUTER_API_KEY" in content and "sk-openrouter-new" in content
def _prompt_key(cli_utils, monkeypatch, tmp_path, key="sk-typed-in"): def _prompt_key(prompts, monkeypatch, tmp_path, key="sk-typed-in"):
monkeypatch.chdir(tmp_path) monkeypatch.chdir(tmp_path)
monkeypatch.delenv("OPENAI_API_KEY", raising=False) monkeypatch.delenv("OPENAI_API_KEY", raising=False)
monkeypatch.setattr(cli_utils, "find_dotenv", lambda **k: "") monkeypatch.setattr(prompts, "find_dotenv", lambda **k: "")
with patch.object(cli_utils, "questionary") as mock_q: with patch.object(prompts, "questionary") as mock_q:
mock_q.password.return_value.ask.return_value = key mock_q.password.return_value.ask.return_value = key
cli_utils.ensure_api_key("openai") prompts.ensure_api_key("openai")
@pytest.mark.skipif(os.name == "nt", reason="POSIX file modes") @pytest.mark.skipif(os.name == "nt", reason="POSIX file modes")
def test_saved_key_file_is_owner_only(monkeypatch, cli_utils, tmp_path): def test_saved_key_file_is_owner_only(monkeypatch, prompts, tmp_path):
# The prompt writes a real credential; the file must not be readable by # The prompt writes a real credential; the file must not be readable by
# other local users whatever the umask is. # other local users whatever the umask is.
old = os.umask(0o002) old = os.umask(0o002)
try: try:
_prompt_key(cli_utils, monkeypatch, tmp_path) _prompt_key(prompts, monkeypatch, tmp_path)
finally: finally:
os.umask(old) os.umask(old)
env = tmp_path / ".env" env = tmp_path / ".env"
@@ -173,21 +173,21 @@ def test_saved_key_file_is_owner_only(monkeypatch, cli_utils, tmp_path):
@pytest.mark.skipif(os.name == "nt", reason="POSIX file modes") @pytest.mark.skipif(os.name == "nt", reason="POSIX file modes")
def test_existing_key_file_is_tightened_before_writing(monkeypatch, cli_utils, tmp_path): def test_existing_key_file_is_tightened_before_writing(monkeypatch, prompts, tmp_path):
env = tmp_path / ".env" env = tmp_path / ".env"
env.write_text("OTHER=1\n") env.write_text("OTHER=1\n")
os.chmod(env, 0o664) os.chmod(env, 0o664)
_prompt_key(cli_utils, monkeypatch, tmp_path) _prompt_key(prompts, monkeypatch, tmp_path)
assert stat.S_IMODE(env.stat().st_mode) == 0o600 assert stat.S_IMODE(env.stat().st_mode) == 0o600
assert "OTHER=1" in env.read_text() assert "OTHER=1" in env.read_text()
@pytest.mark.skipif(os.name == "nt", reason="POSIX file modes") @pytest.mark.skipif(os.name == "nt", reason="POSIX file modes")
def test_read_only_key_file_is_still_updated(monkeypatch, cli_utils, tmp_path): def test_read_only_key_file_is_still_updated(monkeypatch, prompts, tmp_path):
env = tmp_path / ".env" env = tmp_path / ".env"
env.write_text("OTHER=1\n") env.write_text("OTHER=1\n")
os.chmod(env, 0o400) os.chmod(env, 0o400)
_prompt_key(cli_utils, monkeypatch, tmp_path) _prompt_key(prompts, monkeypatch, tmp_path)
assert "sk-typed-in" in env.read_text() assert "sk-typed-in" in env.read_text()
assert stat.S_IMODE(env.stat().st_mode) == 0o600 assert stat.S_IMODE(env.stat().st_mode) == 0o600
+1 -1
View File
@@ -75,7 +75,7 @@ def test_glm_resolves_to_the_endpoint_its_key_belongs_to():
same platform: glm is Z.AI international (ZHIPU_API_KEY) and glm-cn is same platform: glm is Z.AI international (ZHIPU_API_KEY) and glm-cn is
BigModel China. A mismatch sends the key to the other platform and every BigModel China. A mismatch sends the key to the other platform and every
call fails auth.""" call fails auth."""
from cli.utils import resolve_backend_url from cli.prompts import resolve_backend_url
from tradingagents.llm_clients.api_key_env import get_api_key_env from tradingagents.llm_clients.api_key_env import get_api_key_env
from tradingagents.llm_clients.openai_client import OPENAI_COMPATIBLE_PROVIDERS from tradingagents.llm_clients.openai_client import OPENAI_COMPATIBLE_PROVIDERS
+53 -1
View File
@@ -7,10 +7,16 @@ person to read afterwards. Both got that wrong in ways that hide real content.
from __future__ import annotations from __future__ import annotations
import json import json
import unittest
import pytest import pytest
from cli.main import extract_content_string from cli.display import (
AnalystWallTimeTracker,
extract_content_string,
sync_analyst_tracker_from_chunk,
)
from tradingagents.graph.analyst_execution import build_analyst_execution_plan
@pytest.mark.unit @pytest.mark.unit
@@ -80,3 +86,49 @@ def test_the_live_display_does_not_scroll_the_terminal():
import cli.main as m import cli.main as m
assert "screen=True" in inspect.getsource(m.run_analysis) assert "screen=True" in inspect.getsource(m.run_analysis)
class AnalystWallTimeTrackerTests(unittest.TestCase):
def test_records_wall_time_when_analyst_completes(self):
plan = build_analyst_execution_plan(["market", "news"])
tracker = AnalystWallTimeTracker(plan)
tracker.mark_started("market", started_at=10.0)
tracker.mark_completed("market", completed_at=13.5)
self.assertEqual(tracker.format_summary(), "Analyst wall time: Market 3.50s")
def test_formats_summary_in_plan_order(self):
plan = build_analyst_execution_plan(["news", "market"])
tracker = AnalystWallTimeTracker(plan)
tracker.mark_started("market", started_at=20.0)
tracker.mark_completed("market", completed_at=22.25)
tracker.mark_started("news", started_at=10.0)
tracker.mark_completed("news", completed_at=14.0)
self.assertEqual(
tracker.format_summary(),
"Analyst wall time: News 4.00s | Market 2.25s",
)
def test_syncs_wall_time_from_sequential_chunks(self):
plan = build_analyst_execution_plan(["market", "news"])
tracker = AnalystWallTimeTracker(plan)
sync_analyst_tracker_from_chunk(tracker, {}, now=10.0)
self.assertEqual(tracker.format_summary(), "Analyst wall time: pending")
sync_analyst_tracker_from_chunk(
tracker,
{"market_report": "done"},
now=13.0,
)
self.assertEqual(tracker.format_summary(), "Analyst wall time: Market 3.00s")
sync_analyst_tracker_from_chunk(
tracker,
{"market_report": "done", "news_report": "done"},
now=18.0,
)
self.assertEqual(tracker.format_summary(), "Analyst wall time: Market 3.00s | News 5.00s")
+3 -3
View File
@@ -15,17 +15,17 @@ import pytest
@pytest.mark.unit @pytest.mark.unit
class TestProviderDefaultUrl(unittest.TestCase): class TestProviderDefaultUrl(unittest.TestCase):
def test_known_providers_resolve(self): def test_known_providers_resolve(self):
from cli.utils import provider_default_url from cli.prompts import provider_default_url
self.assertEqual(provider_default_url("openai"), "https://api.openai.com/v1") self.assertEqual(provider_default_url("openai"), "https://api.openai.com/v1")
self.assertEqual(provider_default_url("DeepSeek"), "https://api.deepseek.com") self.assertEqual(provider_default_url("DeepSeek"), "https://api.deepseek.com")
self.assertIsNone(provider_default_url("google")) # uses SDK default self.assertIsNone(provider_default_url("google")) # uses SDK default
def test_unknown_provider_returns_none(self): def test_unknown_provider_returns_none(self):
from cli.utils import provider_default_url from cli.prompts import provider_default_url
self.assertIsNone(provider_default_url("not-a-provider")) self.assertIsNone(provider_default_url("not-a-provider"))
def test_ollama_honors_base_url_env(self): def test_ollama_honors_base_url_env(self):
from cli.utils import provider_default_url from cli.prompts import provider_default_url
with mock.patch.dict(os.environ, {"OLLAMA_BASE_URL": "http://host:1234/v1"}): with mock.patch.dict(os.environ, {"OLLAMA_BASE_URL": "http://host:1234/v1"}):
self.assertEqual(provider_default_url("ollama"), "http://host:1234/v1") self.assertEqual(provider_default_url("ollama"), "http://host:1234/v1")
+2 -2
View File
@@ -134,10 +134,10 @@ def test_selections_are_remembered_after_a_run(monkeypatch):
def test_a_custom_language_is_remembered_without_breaking_the_next_run(): def test_a_custom_language_is_remembered_without_breaking_the_next_run():
"""A free-text answer is not one of the menu's choices, and questionary """A free-text answer is not one of the menu's choices, and questionary
rejects a default it cannot find, so offering it back would crash startup.""" rejects a default it cannot find, so offering it back would crash startup."""
from cli.utils import ask_output_language from cli.prompts import ask_output_language
save_last_run({"output_language": "Turkish"}) save_last_run({"output_language": "Turkish"})
with mock.patch("cli.utils.questionary.select") as select: with mock.patch("cli.prompts.questionary.select") as select:
select.return_value.ask.return_value = "English" select.return_value.ask.return_value = "English"
ask_output_language(load_last_run()["output_language"]) ask_output_language(load_last_run()["output_language"])
assert select.call_args.kwargs["default"] is None assert select.call_args.kwargs["default"] is None
+1 -1
View File
@@ -6,7 +6,7 @@ stock), #982 (BTC-USDT accepted but unpriceable on Yahoo).
import pytest import pytest
from cli.models import AssetType from cli.models import AssetType
from cli.utils import detect_asset_type, is_valid_ticker_input, normalize_ticker_symbol from cli.prompts import detect_asset_type, is_valid_ticker_input, normalize_ticker_symbol
from tradingagents.dataflows.symbols import normalize_symbol from tradingagents.dataflows.symbols import normalize_symbol
+1 -1
View File
@@ -1,7 +1,7 @@
import unittest import unittest
from cli.models import AnalystType, AssetType from cli.models import AnalystType, AssetType
from cli.utils import detect_asset_type, filter_analysts_for_asset_type from cli.prompts import detect_asset_type, filter_analysts_for_asset_type
from tradingagents.graph.propagation import Propagator from tradingagents.graph.propagation import Propagator
+23 -23
View File
@@ -22,7 +22,7 @@ def _console_out(capsys) -> str:
def _resync_reloaded_modules(): def _resync_reloaded_modules():
"""Restore module state after this file's importlib.reload() calls. """Restore module state after this file's importlib.reload() calls.
Several tests below reload ``cli.utils`` to re-evaluate OLLAMA_BASE_URL. 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 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 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 tests that happen to run afterward. Re-sync once on teardown so the reload
@@ -30,8 +30,8 @@ def _resync_reloaded_modules():
""" """
yield yield
import cli.main import cli.main
import cli.utils import cli.prompts
importlib.reload(cli.utils) importlib.reload(cli.prompts)
importlib.reload(cli.main) importlib.reload(cli.main)
@@ -98,14 +98,14 @@ def test_explicit_base_url_overrides_env(monkeypatch):
assert "env-set" not in str(llm.openai_api_base) assert "env-set" not in str(llm.openai_api_base)
# ---- cli.utils side: select_llm_provider dropdown ------------------------- # ---- cli.prompts side: select_llm_provider dropdown -------------------------
def test_cli_dropdown_uses_env(monkeypatch): def test_cli_dropdown_uses_env(monkeypatch):
"""The Ollama entry in the CLI dropdown must reflect OLLAMA_BASE_URL.""" """The Ollama entry in the CLI dropdown must reflect OLLAMA_BASE_URL."""
monkeypatch.setenv("OLLAMA_BASE_URL", "http://cli-remote:11434/v1") monkeypatch.setenv("OLLAMA_BASE_URL", "http://cli-remote:11434/v1")
import cli.utils as cli_utils from cli import prompts
importlib.reload(cli_utils) importlib.reload(prompts)
# Reach inside the function via the same env-read it does at call time # Reach inside the function via the same env-read it does at call time
ollama_url = ( ollama_url = (
__import__("os").environ.get("OLLAMA_BASE_URL") __import__("os").environ.get("OLLAMA_BASE_URL")
@@ -116,8 +116,8 @@ def test_cli_dropdown_uses_env(monkeypatch):
def test_cli_dropdown_default_when_unset(monkeypatch): def test_cli_dropdown_default_when_unset(monkeypatch):
monkeypatch.delenv("OLLAMA_BASE_URL", raising=False) monkeypatch.delenv("OLLAMA_BASE_URL", raising=False)
import cli.utils as cli_utils from cli import prompts
importlib.reload(cli_utils) importlib.reload(prompts)
ollama_url = ( ollama_url = (
__import__("os").environ.get("OLLAMA_BASE_URL") __import__("os").environ.get("OLLAMA_BASE_URL")
or "http://localhost:11434/v1" or "http://localhost:11434/v1"
@@ -130,9 +130,9 @@ def test_cli_dropdown_default_when_unset(monkeypatch):
def test_confirm_endpoint_shows_default(monkeypatch, capsys): def test_confirm_endpoint_shows_default(monkeypatch, capsys):
monkeypatch.delenv("OLLAMA_BASE_URL", raising=False) monkeypatch.delenv("OLLAMA_BASE_URL", raising=False)
import cli.utils as cli_utils from cli import prompts
importlib.reload(cli_utils) importlib.reload(prompts)
cli_utils.confirm_ollama_endpoint("http://localhost:11434/v1") prompts.confirm_ollama_endpoint("http://localhost:11434/v1")
out = _console_out(capsys) out = _console_out(capsys)
assert "http://localhost:11434/v1" in out assert "http://localhost:11434/v1" in out
assert "OLLAMA_BASE_URL" not in out # not from env assert "OLLAMA_BASE_URL" not in out # not from env
@@ -141,9 +141,9 @@ def test_confirm_endpoint_shows_default(monkeypatch, capsys):
def test_confirm_endpoint_marks_env_origin(monkeypatch, capsys): def test_confirm_endpoint_marks_env_origin(monkeypatch, capsys):
monkeypatch.setenv("OLLAMA_BASE_URL", "http://remote-host:11434/v1") monkeypatch.setenv("OLLAMA_BASE_URL", "http://remote-host:11434/v1")
import cli.utils as cli_utils from cli import prompts
importlib.reload(cli_utils) importlib.reload(prompts)
cli_utils.confirm_ollama_endpoint("http://remote-host:11434/v1") prompts.confirm_ollama_endpoint("http://remote-host:11434/v1")
out = _console_out(capsys) out = _console_out(capsys)
assert "http://remote-host:11434/v1" in out assert "http://remote-host:11434/v1" in out
assert "OLLAMA_BASE_URL" in out assert "OLLAMA_BASE_URL" in out
@@ -152,9 +152,9 @@ def test_confirm_endpoint_marks_env_origin(monkeypatch, capsys):
def test_confirm_endpoint_warns_on_missing_scheme(monkeypatch, capsys): def test_confirm_endpoint_warns_on_missing_scheme(monkeypatch, capsys):
"""If user sets OLLAMA_BASE_URL=0.0.0.128, advise on the expected shape.""" """If user sets OLLAMA_BASE_URL=0.0.0.128, advise on the expected shape."""
monkeypatch.setenv("OLLAMA_BASE_URL", "0.0.0.128") monkeypatch.setenv("OLLAMA_BASE_URL", "0.0.0.128")
import cli.utils as cli_utils from cli import prompts
importlib.reload(cli_utils) importlib.reload(prompts)
cli_utils.confirm_ollama_endpoint("0.0.0.128") prompts.confirm_ollama_endpoint("0.0.0.128")
out = _console_out(capsys) out = _console_out(capsys)
assert "missing a scheme" in out assert "missing a scheme" in out
assert "http://<host>:11434/v1" in out assert "http://<host>:11434/v1" in out
@@ -163,9 +163,9 @@ def test_confirm_endpoint_warns_on_missing_scheme(monkeypatch, capsys):
def test_confirm_endpoint_warns_on_non_default_port_remote(monkeypatch, capsys): def test_confirm_endpoint_warns_on_non_default_port_remote(monkeypatch, capsys):
"""A remote host with no :11434 gets a soft hint about port mismatch.""" """A remote host with no :11434 gets a soft hint about port mismatch."""
monkeypatch.setenv("OLLAMA_BASE_URL", "http://remote-host/v1") monkeypatch.setenv("OLLAMA_BASE_URL", "http://remote-host/v1")
import cli.utils as cli_utils from cli import prompts
importlib.reload(cli_utils) importlib.reload(prompts)
cli_utils.confirm_ollama_endpoint("http://remote-host/v1") prompts.confirm_ollama_endpoint("http://remote-host/v1")
out = _console_out(capsys) out = _console_out(capsys)
assert "port 11434" in out assert "port 11434" in out
@@ -173,9 +173,9 @@ def test_confirm_endpoint_warns_on_non_default_port_remote(monkeypatch, capsys):
def test_confirm_endpoint_quiet_on_local_no_port(monkeypatch, capsys): def test_confirm_endpoint_quiet_on_local_no_port(monkeypatch, capsys):
"""Local host without port shouldn't trigger the remote-port hint.""" """Local host without port shouldn't trigger the remote-port hint."""
monkeypatch.setenv("OLLAMA_BASE_URL", "http://localhost/v1") monkeypatch.setenv("OLLAMA_BASE_URL", "http://localhost/v1")
import cli.utils as cli_utils from cli import prompts
importlib.reload(cli_utils) importlib.reload(prompts)
cli_utils.confirm_ollama_endpoint("http://localhost/v1") prompts.confirm_ollama_endpoint("http://localhost/v1")
out = _console_out(capsys) out = _console_out(capsys)
assert "Note" not in out # localhost is fine without explicit port assert "Note" not in out # localhost is fine without explicit port
+1 -1
View File
@@ -68,7 +68,7 @@ def test_any_model_accepted_no_forced_key():
@pytest.mark.unit @pytest.mark.unit
def test_env_backend_url_precedence(): def test_env_backend_url_precedence():
# #978: explicit env URL wins over the menu/default regardless of provider source. # #978: explicit env URL wins over the menu/default regardless of provider source.
from cli.utils import resolve_backend_url from cli.prompts import resolve_backend_url
assert resolve_backend_url("openai", "https://api.openai.com/v1", env_url="http://proxy/v1") == "http://proxy/v1" assert resolve_backend_url("openai", "https://api.openai.com/v1", env_url="http://proxy/v1") == "http://proxy/v1"
assert resolve_backend_url("openai", "https://api.openai.com/v1", env_url=None) == "https://api.openai.com/v1" assert resolve_backend_url("openai", "https://api.openai.com/v1", env_url=None) == "https://api.openai.com/v1"
assert resolve_backend_url("deepseek", None, None) == "https://api.deepseek.com" assert resolve_backend_url("deepseek", None, None) == "https://api.deepseek.com"
+25 -25
View File
@@ -6,7 +6,7 @@ from unittest import mock
import pytest import pytest
from cli import utils from cli import prompts
def _asks(value): def _asks(value):
@@ -23,10 +23,10 @@ class TestOpenRouterPromptLabel:
captured["message"] = message captured["message"] = message
return _asks("openrouter/some-model") return _asks("openrouter/some-model")
with mock.patch.object(utils, "_fetch_openrouter_models", with mock.patch.object(prompts, "_fetch_openrouter_models",
return_value=[("Some Model", "openrouter/some-model")]), \ return_value=[("Some Model", "openrouter/some-model")]), \
mock.patch.object(utils.questionary, "select", side_effect=fake_select): mock.patch.object(prompts.questionary, "select", side_effect=fake_select):
out = utils.select_openrouter_model(mode) out = prompts.select_openrouter_model(mode)
assert label in captured["message"] assert label in captured["message"]
assert out == "openrouter/some-model" assert out == "openrouter/some-model"
@@ -44,7 +44,7 @@ class TestOpenRouterLatestFirst:
resp.json.return_value = payload resp.json.return_value = payload
resp.raise_for_status = mock.Mock() resp.raise_for_status = mock.Mock()
with mock.patch("requests.get", return_value=resp): with mock.patch("requests.get", return_value=resp):
out = utils._fetch_openrouter_models() out = prompts._fetch_openrouter_models()
assert [mid for _, mid in out] == ["new/model", "mid/model", "old/model"] assert [mid for _, mid in out] == ["new/model", "mid/model", "old/model"]
@@ -64,9 +64,9 @@ class TestMainstreamFilter:
captured["values"] = [c.value for c in kwargs["choices"]] captured["values"] = [c.value for c in kwargs["choices"]]
return _asks("anthropic/claude-x") return _asks("anthropic/claude-x")
with mock.patch.object(utils, "_fetch_openrouter_models", return_value=models), \ with mock.patch.object(prompts, "_fetch_openrouter_models", return_value=models), \
mock.patch.object(utils.questionary, "select", side_effect=fake_select): mock.patch.object(prompts.questionary, "select", side_effect=fake_select):
utils.select_openrouter_model("quick") prompts.select_openrouter_model("quick")
assert "anthropic/claude-x" in captured["values"] assert "anthropic/claude-x" in captured["values"]
assert "openai/gpt-x" in captured["values"] assert "openai/gpt-x" in captured["values"]
@@ -82,9 +82,9 @@ class TestMainstreamFilter:
captured["values"] = [c.value for c in kwargs["choices"]] captured["values"] = [c.value for c in kwargs["choices"]]
return _asks("nex-agi/x") return _asks("nex-agi/x")
with mock.patch.object(utils, "_fetch_openrouter_models", return_value=models), \ with mock.patch.object(prompts, "_fetch_openrouter_models", return_value=models), \
mock.patch.object(utils.questionary, "select", side_effect=fake_select): mock.patch.object(prompts.questionary, "select", side_effect=fake_select):
utils.select_openrouter_model("deep") prompts.select_openrouter_model("deep")
assert "nex-agi/x" in captured["values"] # fallback keeps the list usable assert "nex-agi/x" in captured["values"] # fallback keeps the list usable
@@ -92,31 +92,31 @@ class TestMainstreamFilter:
@pytest.mark.unit @pytest.mark.unit
class TestCancelExitsCleanly: class TestCancelExitsCleanly:
def test_dropdown_cancel_exits(self): def test_dropdown_cancel_exits(self):
with mock.patch.object(utils, "_fetch_openrouter_models", return_value=[]), \ with mock.patch.object(prompts, "_fetch_openrouter_models", return_value=[]), \
mock.patch.object(utils.questionary, "select", return_value=_asks(None)), \ mock.patch.object(prompts.questionary, "select", return_value=_asks(None)), \
pytest.raises(SystemExit): pytest.raises(SystemExit):
utils.select_openrouter_model("quick") prompts.select_openrouter_model("quick")
def test_custom_id_cancel_exits(self): def test_custom_id_cancel_exits(self):
with mock.patch.object(utils, "_fetch_openrouter_models", return_value=[]), \ with mock.patch.object(prompts, "_fetch_openrouter_models", return_value=[]), \
mock.patch.object(utils.questionary, "select", return_value=_asks("custom")), \ mock.patch.object(prompts.questionary, "select", return_value=_asks("custom")), \
mock.patch.object(utils.questionary, "text", return_value=_asks(None)), \ mock.patch.object(prompts.questionary, "text", return_value=_asks(None)), \
pytest.raises(SystemExit): pytest.raises(SystemExit):
utils.select_openrouter_model("deep") prompts.select_openrouter_model("deep")
def test_prompt_custom_model_id_cancel_exits(self): def test_prompt_custom_model_id_cancel_exits(self):
with mock.patch.object(utils.questionary, "text", return_value=_asks(None)), \ with mock.patch.object(prompts.questionary, "text", return_value=_asks(None)), \
pytest.raises(SystemExit): pytest.raises(SystemExit):
utils._prompt_custom_model_id() prompts._prompt_custom_model_id()
@pytest.mark.unit @pytest.mark.unit
class TestLanguageDefaultsToEnglish: class TestLanguageDefaultsToEnglish:
def test_select_cancel_defaults_english(self): def test_select_cancel_defaults_english(self):
with mock.patch.object(utils.questionary, "select", return_value=_asks(None)): with mock.patch.object(prompts.questionary, "select", return_value=_asks(None)):
assert utils.ask_output_language() == "English" assert prompts.ask_output_language() == "English"
def test_custom_language_cancel_defaults_english(self): def test_custom_language_cancel_defaults_english(self):
with mock.patch.object(utils.questionary, "select", return_value=_asks("custom")), \ with mock.patch.object(prompts.questionary, "select", return_value=_asks("custom")), \
mock.patch.object(utils.questionary, "text", return_value=_asks(None)): mock.patch.object(prompts.questionary, "text", return_value=_asks(None)):
assert utils.ask_output_language() == "English" assert prompts.ask_output_language() == "English"
+4 -4
View File
@@ -2,7 +2,7 @@ import unittest
import pytest import pytest
from cli.utils import normalize_ticker_symbol from cli.prompts import normalize_ticker_symbol
from tradingagents.agents.context import build_instrument_context from tradingagents.agents.context import build_instrument_context
@@ -19,10 +19,10 @@ class TickerSymbolHandlingTests(unittest.TestCase):
def test_single_get_ticker_no_shadow(self): def test_single_get_ticker_no_shadow(self):
# Regression: cli/main.py had a duplicate get_ticker with an empty # Regression: cli/main.py had a duplicate get_ticker with an empty
# questionary prompt (rendered as a bare "?") that shadowed the # questionary prompt (rendered as a bare "?") that shadowed the
# descriptive one in cli/utils. Keep a single canonical definition. # descriptive one in cli/prompts. Keep a single canonical definition.
import cli.main import cli.main
import cli.utils import cli.prompts
self.assertIs(cli.main.get_ticker, cli.utils.get_ticker) self.assertIs(cli.main.get_ticker, cli.prompts.get_ticker)
if __name__ == "__main__": if __name__ == "__main__":
-62
View File
@@ -1,6 +1,5 @@
from collections.abc import Iterable from collections.abc import Iterable
from dataclasses import dataclass from dataclasses import dataclass
from time import monotonic
from tradingagents.agents.analysts import fundamentals_analyst, market_analyst, news_analyst from tradingagents.agents.analysts import fundamentals_analyst, market_analyst, news_analyst
@@ -73,64 +72,3 @@ def build_analyst_execution_plan(
return AnalystExecutionPlan(specs=specs) return AnalystExecutionPlan(specs=specs)
def get_initial_analyst_node(plan: AnalystExecutionPlan) -> str:
return plan.specs[0].agent_node
class AnalystWallTimeTracker:
def __init__(self, plan: AnalystExecutionPlan):
self.plan = plan
self._started_at: dict[str, float] = {}
self._wall_times: dict[str, float] = {}
def mark_started(self, analyst_key: str, started_at: float | None = None) -> None:
if analyst_key not in ANALYST_NODE_SPECS:
raise ValueError(f"unknown analyst key: {analyst_key}")
self._started_at.setdefault(analyst_key, monotonic() if started_at is None else started_at)
def mark_completed(
self,
analyst_key: str,
completed_at: float | None = None,
) -> None:
if analyst_key not in ANALYST_NODE_SPECS:
raise ValueError(f"unknown analyst key: {analyst_key}")
if analyst_key in self._wall_times:
return
started_at = self._started_at.get(analyst_key)
if started_at is None:
return
finished_at = monotonic() if completed_at is None else completed_at
self._wall_times[analyst_key] = max(0.0, finished_at - started_at)
def format_summary(self) -> str:
parts = []
for spec in self.plan.specs:
duration = self._wall_times.get(spec.key)
if duration is not None:
label = spec.agent_node.removesuffix(" Analyst")
parts.append(f"{label} {duration:.2f}s")
if not parts:
return "Analyst wall time: pending"
return "Analyst wall time: " + " | ".join(parts)
def sync_analyst_tracker_from_chunk(
tracker: AnalystWallTimeTracker,
chunk: dict[str, str],
now: float | None = None,
) -> None:
current_time = monotonic() if now is None else now
active_found = False
for spec in tracker.plan.specs:
has_report = bool(chunk.get(spec.report_key))
if has_report:
tracker.mark_started(spec.key, started_at=current_time)
tracker.mark_completed(spec.key, completed_at=current_time)
continue
if not active_found:
tracker.mark_started(spec.key, started_at=current_time)
active_found = True
+1 -1
View File
@@ -2,7 +2,7 @@
A single source of truth for which environment variable holds the API A single source of truth for which environment variable holds the API
key for each supported LLM provider. Used by the CLI's interactive key key for each supported LLM provider. Used by the CLI's interactive key
prompt (cli/utils.ensure_api_key) and by anything else that needs to prompt (cli/prompts.ensure_api_key) and by anything else that needs to
ask "does this provider require a key, and which env var is it?". ask "does this provider require a key, and which env var is it?".
When adding a new provider, register its env var here so the CLI flow When adding a new provider, register its env var here so the CLI flow
+1 -1
View File
@@ -182,7 +182,7 @@ MODEL_OPTIONS: ProviderModeOptions = {
# endpoint is now configurable via OLLAMA_BASE_URL, so the same labels # endpoint is now configurable via OLLAMA_BASE_URL, so the same labels
# apply whether the user runs ollama-serve on localhost or against a # apply whether the user runs ollama-serve on localhost or against a
# remote host. The actual resolved endpoint is surfaced separately by # remote host. The actual resolved endpoint is surfaced separately by
# cli.utils.confirm_ollama_endpoint() right after provider selection. # cli.prompts.confirm_ollama_endpoint() right after provider selection.
# "Custom model ID" lets users pick any model they have pulled via # "Custom model ID" lets users pick any model they have pulled via
# `ollama pull` beyond the three suggested defaults. # `ollama pull` beyond the three suggested defaults.
"ollama": { "ollama": {