mirror of
https://github.com/TauricResearch/TradingAgents.git
synced 2026-09-19 11:15:24 +03:00
fix(cli): read and write the decision log on the CLI path
- shared create_run_state and record_decision for propagate() and the CLI #1332 #1347
This commit is contained in:
34
cli/main.py
34
cli/main.py
@@ -1110,18 +1110,10 @@ def run_analysis(checkpoint: bool | None = None):
|
||||
)
|
||||
update_display(layout, spinner_text, stats_handler=stats_handler, start_time=start_time)
|
||||
|
||||
# Initialize state and get graph args with callbacks.
|
||||
# Resolve the instrument identity once here so all agents anchor to
|
||||
# the real company (#814); the CLI builds state directly rather than
|
||||
# going through propagate(), so this must happen on the CLI path too.
|
||||
instrument_context = graph.resolve_instrument_context(
|
||||
selections["ticker"], selections["asset_type"]
|
||||
)
|
||||
init_agent_state = graph.propagator.create_initial_state(
|
||||
selections["ticker"],
|
||||
selections["analysis_date"],
|
||||
asset_type=selections["asset_type"],
|
||||
instrument_context=instrument_context,
|
||||
# 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"]
|
||||
)
|
||||
# Pass callbacks to graph config for tool execution tracking
|
||||
# (LLM tracking is handled separately via LLM constructor)
|
||||
@@ -1243,8 +1235,16 @@ def run_analysis(checkpoint: bool | None = None):
|
||||
|
||||
trace.append(chunk)
|
||||
|
||||
# Clean run: drop this run's checkpoint so a later run starts fresh.
|
||||
# A mid-stream failure skips this, keeping the checkpoint for resume.
|
||||
# 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"]
|
||||
)
|
||||
@@ -1252,12 +1252,6 @@ def run_analysis(checkpoint: bool | None = None):
|
||||
# Always restore the plain uncheckpointed graph, even on failure.
|
||||
graph.end_checkpoint()
|
||||
|
||||
# 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)
|
||||
|
||||
# Update all agent statuses to completed
|
||||
for agent in message_buffer.agent_status:
|
||||
message_buffer.update_agent_status(agent, "completed")
|
||||
|
||||
163
tests/test_cli_decision_log.py
Normal file
163
tests/test_cli_decision_log.py
Normal file
@@ -0,0 +1,163 @@
|
||||
"""The CLI must use the decision log the same way propagate() does.
|
||||
|
||||
The CLI streams the graph itself instead of calling propagate(), so memory steps
|
||||
that lived only in propagate() never ran on the primary entry point: pending
|
||||
decisions were not settled, the Portfolio Manager got no past context, and the
|
||||
finished decision was not recorded. Both paths now build their initial state and
|
||||
record their decision through the same graph methods.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from tradingagents.agents.utils.memory import TradingMemoryLog
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
|
||||
|
||||
def _bare_graph(tmp_path):
|
||||
"""A graph without __init__ (no LLM clients), wired to a temp log."""
|
||||
graph = object.__new__(TradingAgentsGraph)
|
||||
graph.config = {"memory_log_path": str(tmp_path / "trading_memory.md")}
|
||||
graph.memory_log = TradingMemoryLog(graph.config)
|
||||
return graph
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_create_run_state_settles_pending_and_carries_context(tmp_path, monkeypatch):
|
||||
from tradingagents.graph.propagation import Propagator
|
||||
|
||||
graph = _bare_graph(tmp_path)
|
||||
graph.propagator = Propagator()
|
||||
settled = []
|
||||
monkeypatch.setattr(graph, "_resolve_pending_entries", settled.append, raising=False)
|
||||
monkeypatch.setattr(graph, "resolve_instrument_context", lambda t, a="stock": f"id:{t}", raising=False)
|
||||
monkeypatch.setattr(graph, "_memory_as_of", lambda d: d, raising=False)
|
||||
graph.memory_log.store_decision("NVDA", "2026-01-05", "Rating: Buy\nold call")
|
||||
graph.memory_log.update_with_outcome("NVDA", "2026-01-05", 0.01, 0.005, 5, "great trade", "2026-01-12")
|
||||
|
||||
state = graph.create_run_state("NVDA", "2026-02-01")
|
||||
|
||||
assert settled == ["NVDA"]
|
||||
assert "great trade" in state["past_context"]
|
||||
assert state["instrument_context"] == "id:NVDA"
|
||||
assert state["company_of_interest"] == "NVDA"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_record_decision_appends_a_pending_entry(tmp_path):
|
||||
graph = _bare_graph(tmp_path)
|
||||
graph.record_decision("NVDA", "2026-01-10", {"final_trade_decision": "Rating: Buy\n\nBuy NVDA."})
|
||||
entries = graph.memory_log.load_entries()
|
||||
assert [(e["ticker"], e["pending"], e["rating"]) for e in entries] == [("NVDA", True, "Buy")]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_record_decision_skips_a_run_without_a_decision(tmp_path):
|
||||
graph = _bare_graph(tmp_path)
|
||||
graph.record_decision("NVDA", "2026-01-10", {})
|
||||
assert graph.memory_log.load_entries() == []
|
||||
|
||||
|
||||
# --- the CLI path ----------------------------------------------------------------
|
||||
|
||||
class _FakeGraph:
|
||||
"""Records the lifecycle calls run_analysis makes."""
|
||||
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
self.graph = self
|
||||
self.propagator = self
|
||||
|
||||
def create_run_state(self, ticker, trade_date, asset_type="stock"):
|
||||
self.calls.append(("create_run_state", ticker, trade_date))
|
||||
return {"messages": [], "company_of_interest": ticker}
|
||||
|
||||
def record_decision(self, ticker, trade_date, final_state):
|
||||
self.calls.append(("record_decision", ticker, trade_date, final_state.get("final_trade_decision")))
|
||||
|
||||
def get_graph_args(self, callbacks=None):
|
||||
return {}
|
||||
|
||||
def begin_checkpoint(self, *a, **k):
|
||||
return None
|
||||
|
||||
def checkpoint_input(self, state):
|
||||
return state
|
||||
|
||||
def clear_checkpoint_on_success(self, *a, **k):
|
||||
self.calls.append(("clear_checkpoint",))
|
||||
|
||||
def end_checkpoint(self):
|
||||
pass
|
||||
|
||||
def stream(self, graph_input, **kwargs):
|
||||
yield {"messages": [], "market_report": "M"}
|
||||
yield {"messages": [], "final_trade_decision": "Rating: Buy\n\nBuy NVDA."}
|
||||
|
||||
|
||||
class _NullLive:
|
||||
def __init__(self, *a, **k):
|
||||
pass
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
|
||||
class _FakeBuffer:
|
||||
def __init__(self):
|
||||
self.messages = []
|
||||
self.tool_calls = []
|
||||
self.report_sections = {}
|
||||
self.agent_status = {}
|
||||
self.selected_analysts = []
|
||||
self._processed_message_ids = set()
|
||||
|
||||
def init_for_analysis(self, selected_analysts):
|
||||
self.selected_analysts = [a.lower() for a in selected_analysts]
|
||||
|
||||
def add_message(self, kind, content):
|
||||
self.messages.append((0.0, kind, content))
|
||||
|
||||
def add_tool_call(self, name, args):
|
||||
self.tool_calls.append((0.0, name, args))
|
||||
|
||||
def update_report_section(self, *a):
|
||||
pass
|
||||
|
||||
def update_agent_status(self, agent, status):
|
||||
self.agent_status[agent] = status
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_cli_run_uses_the_decision_log_like_propagate(tmp_path, monkeypatch):
|
||||
import cli.main as m
|
||||
from cli.models import AnalystType
|
||||
|
||||
fake = _FakeGraph()
|
||||
monkeypatch.setattr(m, "TradingAgentsGraph", lambda *a, **k: fake)
|
||||
monkeypatch.setattr(m, "message_buffer", _FakeBuffer())
|
||||
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: {
|
||||
"ticker": "NVDA", "analysis_date": "2026-01-10",
|
||||
"analysts": [AnalystType.MARKET], "asset_type": "stock",
|
||||
})
|
||||
monkeypatch.setattr(m, "_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()
|
||||
|
||||
assert fake.calls == [
|
||||
("create_run_state", "NVDA", "2026-01-10"),
|
||||
# The decision is recorded from the merged stream, before the checkpoint
|
||||
# is cleared, matching propagate().
|
||||
("record_decision", "NVDA", "2026-01-10", "Rating: Buy\n\nBuy NVDA."),
|
||||
("clear_checkpoint",),
|
||||
]
|
||||
@@ -903,6 +903,9 @@ class TestLegacyRemoval:
|
||||
mock_graph._run_graph = functools.partial(
|
||||
TradingAgentsGraph._run_graph, mock_graph
|
||||
)
|
||||
mock_graph.record_decision = functools.partial(
|
||||
TradingAgentsGraph.record_decision, mock_graph
|
||||
)
|
||||
TradingAgentsGraph.propagate(mock_graph, "NVDA", "2026-01-10")
|
||||
entries = mock_graph.memory_log.load_entries()
|
||||
assert len(entries) == 1
|
||||
|
||||
@@ -419,9 +419,6 @@ class TradingAgentsGraph:
|
||||
"""
|
||||
self.ticker = company_name
|
||||
|
||||
# Resolve any pending memory-log entries for this ticker before the pipeline runs.
|
||||
self._resolve_pending_entries(company_name)
|
||||
|
||||
with self.checkpoint_scope(company_name, trade_date, asset_type) as thread_id_value:
|
||||
return self._run_graph(
|
||||
company_name, trade_date, asset_type=asset_type,
|
||||
@@ -506,24 +503,39 @@ class TradingAgentsGraph:
|
||||
)
|
||||
return write_report_tree(final_state, ticker, save_path)
|
||||
|
||||
def _run_graph(self, company_name, trade_date, asset_type: str = "stock",
|
||||
checkpoint_thread_id: str | None = None):
|
||||
"""Execute the graph and write the resulting state to disk and memory log."""
|
||||
# Initialize state — inject memory log context for PM and the
|
||||
# deterministically resolved instrument identity for all agents. On a
|
||||
# historical run, gate lessons to those whose outcome was known by the
|
||||
# trade date so a backtest can't learn from the future (#1251).
|
||||
past_context = self.memory_log.get_past_context(
|
||||
company_name, as_of=self._memory_as_of(trade_date)
|
||||
)
|
||||
instrument_context = self.resolve_instrument_context(company_name, asset_type)
|
||||
init_agent_state = self.propagator.create_initial_state(
|
||||
def create_run_state(self, company_name, trade_date, asset_type: str = "stock"):
|
||||
"""Build a run's initial state; propagate() and the CLI both start here.
|
||||
|
||||
Settles this ticker's pending decisions first, then injects the lessons
|
||||
known by the trade date for the Portfolio Manager (#1251) and the
|
||||
resolved instrument identity for every agent (#814). An entry point that
|
||||
assembled the state itself would skip the decision log.
|
||||
"""
|
||||
self._resolve_pending_entries(company_name)
|
||||
return self.propagator.create_initial_state(
|
||||
company_name,
|
||||
trade_date,
|
||||
asset_type=asset_type,
|
||||
past_context=past_context,
|
||||
instrument_context=instrument_context,
|
||||
past_context=self.memory_log.get_past_context(
|
||||
company_name, as_of=self._memory_as_of(trade_date)
|
||||
),
|
||||
instrument_context=self.resolve_instrument_context(company_name, asset_type),
|
||||
)
|
||||
|
||||
def record_decision(self, company_name, trade_date, final_state):
|
||||
"""Log a finished run's decision for reflection on the next same-ticker run."""
|
||||
decision = final_state.get("final_trade_decision")
|
||||
if not decision:
|
||||
logger.warning("No final decision for %s on %s; nothing logged", company_name, trade_date)
|
||||
return
|
||||
self.memory_log.store_decision(
|
||||
ticker=company_name, trade_date=trade_date, final_trade_decision=decision
|
||||
)
|
||||
|
||||
def _run_graph(self, company_name, trade_date, asset_type: str = "stock",
|
||||
checkpoint_thread_id: str | None = None):
|
||||
"""Execute the graph and write the resulting state to disk and memory log."""
|
||||
init_agent_state = self.create_run_state(company_name, trade_date, asset_type)
|
||||
args = self.propagator.get_graph_args()
|
||||
|
||||
# Inject the checkpoint thread_id (from checkpoint_scope) so the same
|
||||
@@ -561,12 +573,7 @@ class TradingAgentsGraph:
|
||||
# Log state to disk.
|
||||
self._log_state(trade_date, final_state)
|
||||
|
||||
# Store decision for deferred reflection on the next same-ticker run.
|
||||
self.memory_log.store_decision(
|
||||
ticker=company_name,
|
||||
trade_date=trade_date,
|
||||
final_trade_decision=final_state["final_trade_decision"],
|
||||
)
|
||||
self.record_decision(company_name, trade_date, final_state)
|
||||
|
||||
# Clear checkpoint on successful completion to avoid stale state.
|
||||
self.clear_checkpoint_on_success(company_name, trade_date, asset_type)
|
||||
|
||||
Reference in New Issue
Block a user