Files
tradingagents/tradingagents/graph/trading_graph.py
T
Yijia-Xiao 9968bd8dd1 feat(graph): run the analysts at the same time (#1255)
- each analyst is a graph of its own (model and tools on a private message history) that returns only its report; all start together and the research debate waits for every report
- the message-clearing nodes are gone; a checkpoint saved by the sequential layout starts fresh
- TradingAgentsGraph.stream_run streams the analysts' messages for debug mode and the CLI, whose status and timing now track the analysts side by side
2026-09-25 06:35:54 +00:00

423 lines
20 KiB
Python

import json
import logging
import os
from contextlib import contextmanager
from datetime import datetime
from pathlib import Path
from typing import Any
import tradingagents
from tradingagents.agents.context import build_instrument_context, resolve_instrument_identity
from tradingagents.agents.rating import run_rating
from tradingagents.dataflows.config import run_config, set_config
from tradingagents.dataflows.date_window import get_current_date
from tradingagents.dataflows.symbols import safe_ticker_component
from tradingagents.default_config import DEFAULT_CONFIG
from tradingagents.llm_clients import build_llm_kwargs, create_llm_client
from tradingagents.memory import TradingMemoryLog, settlement
from tradingagents.memory.reflection import Reflector
from tradingagents.reporting import write_report_tree
from .checkpointer import checkpoint_step, clear_checkpoint, get_checkpointer, thread_id
from .conditional_logic import ConditionalLogic
from .propagation import Propagator
from .setup import GraphSetup
logger = logging.getLogger(__name__)
def _validate_trade_date(trade_date) -> str:
"""The run date as a canonical ``YYYY-MM-DD`` string no later than today."""
value = str(trade_date)
try:
canonical = datetime.strptime(value, "%Y-%m-%d").strftime("%Y-%m-%d") == value
except ValueError:
canonical = False
if not canonical:
raise ValueError(f"trade_date must be a date in YYYY-MM-DD format, got {trade_date!r}")
if value > get_current_date():
raise ValueError(f"trade_date cannot be in the future: {value}")
return value
class TradingAgentsGraph:
"""Main class that orchestrates the trading agents framework."""
def __init__(
self,
selected_analysts=("market", "social", "news", "fundamentals"),
debug=False,
config: dict[str, Any] = None,
callbacks: list | None = None,
):
"""Initialize the trading agents graph and components.
Args:
selected_analysts: List of analyst types to include
debug: Whether to run in debug mode
config: Configuration dictionary. If None, uses default config
callbacks: Optional list of callback handlers (e.g., for tracking LLM/tool stats)
"""
self.debug = debug
self.config = config or DEFAULT_CONFIG
self.callbacks = callbacks or []
set_config(self.config)
os.makedirs(self.config["data_cache_dir"], exist_ok=True)
os.makedirs(self.config["results_dir"], exist_ok=True)
llm_kwargs = build_llm_kwargs(self.config)
if self.callbacks:
llm_kwargs["callbacks"] = self.callbacks
deep_client = create_llm_client(
provider=self.config["llm_provider"],
model=self.config["deep_think_llm"],
base_url=self.config.get("backend_url"),
**llm_kwargs,
)
quick_client = create_llm_client(
provider=self.config["llm_provider"],
model=self.config["quick_think_llm"],
base_url=self.config.get("backend_url"),
**llm_kwargs,
)
self.deep_thinking_llm = deep_client.get_llm()
self.quick_thinking_llm = quick_client.get_llm()
self.memory_log = TradingMemoryLog(self.config)
self.conditional_logic = ConditionalLogic(
max_debate_rounds=self.config["max_debate_rounds"],
max_risk_discuss_rounds=self.config["max_risk_discuss_rounds"],
)
self.graph_setup = GraphSetup(
self.quick_thinking_llm,
self.deep_thinking_llm,
self.conditional_logic,
)
self.propagator = Propagator(
max_recur_limit=self.config.get("max_recur_limit", 100),
)
self.reflector = Reflector(self.quick_thinking_llm)
# Graph-shape-affecting run choices, kept for the checkpoint signature.
self.selected_analysts = tuple(selected_analysts)
# Set up the graph: keep the workflow for recompilation with a checkpointer.
self.workflow = self.graph_setup.setup_graph(selected_analysts)
self.graph = self.workflow.compile()
self._checkpointer_ctx = None
self._resuming = False
def resolve_instrument_context(self, ticker: str, asset_type: str = "stock",
trade_date: str | None = None) -> str:
"""Resolve ticker identity once and return the full instrument context.
Deterministic yfinance lookup (cached, fail-open) injected into a
context string so every agent anchors to the real company instead of
hallucinating one from the price chart (#814). Both the propagate()
path and the CLI call this so the resolved identity reaches the whole
graph regardless of entry point.
"""
identity = resolve_instrument_identity(ticker)
return build_instrument_context(ticker, asset_type, identity, trade_date)
def _memory_as_of(self, trade_date) -> str | None:
"""Point-in-time cutoff for past-context lessons (#1251).
A historical/backtest run (trade date before today) filters lessons to
those already resolved by the trade date. A current-date run returns
None, disabling the filter so live behavior and pre-migration entries
(which have no stored resolution date) are unaffected.
"""
td = str(trade_date)
return td if td < datetime.now().strftime("%Y-%m-%d") else None
def _run_signature(self, asset_type: str, portfolio=None) -> str:
"""Graph-shape inputs that must invalidate a checkpoint if changed.
Keyed into the checkpoint thread ID so a resume under a different analyst
selection, debate/risk depth, or asset mode starts fresh instead of
silently continuing the previous graph (#1089).
"""
return "|".join([
"analysts=" + ",".join(self.selected_analysts),
f"debate={self.config['max_debate_rounds']}",
f"risk={self.config['max_risk_discuss_rounds']}",
f"asset={asset_type}",
# None, an empty book and a changed book are three different runs.
f"portfolio={portfolio.fingerprint() if portfolio is not None else 'none'}",
# The layout itself: a checkpoint saved when analysts ran one after
# another has pending nodes this graph no longer has.
"analysts=parallel",
])
def propagate(self, company_name, trade_date, asset_type: str = "stock", portfolio=None):
"""Run the trading agents graph for a company on a specific date.
``asset_type`` selects between the stock pipeline (default) and the
crypto pipeline (``"crypto"``) shipped in #567 — the CLI auto-detects
from the ticker; programmatic callers pass it explicitly. When
``checkpoint_enabled`` is set in config, the graph is recompiled with
a per-ticker SqliteSaver so a crashed run can resume from the last
successful node on a subsequent invocation with the same ticker+date.
Returns ``(final_state, signal)`` where ``signal`` is one of the 5-tier
ratings (Buy / Overweight / Hold / Underweight / Sell) or ``"REVIEW"``
when the decision had no parseable rating (#1170); guard with
``tradingagents.agents.rating.is_review`` before mapping it to the
PortfolioRating enum.
"""
trade_date = _validate_trade_date(trade_date)
with run_config(self.config), \
self.checkpoint_scope(company_name, trade_date, asset_type, portfolio) as thread_id_value:
return self._run_graph(
company_name, trade_date, asset_type=asset_type,
checkpoint_thread_id=thread_id_value, portfolio=portfolio,
)
def begin_checkpoint(self, company_name, trade_date, asset_type: str = "stock", portfolio=None) -> str | None:
"""Recompile the graph with a per-ticker checkpointer and return the
``thread_id`` to inject into the stream/invoke ``config`` (or ``None``
when checkpointing is disabled).
Pair every call with :meth:`end_checkpoint` in a ``finally``. Both
``propagate`` (via :meth:`checkpoint_scope`) and the CLI stream path use
this so ``--checkpoint`` actually resumes (#1249); previously the setup
lived only inside ``propagate`` and the CLI streamed the checkpointer-less
graph, making the flag a no-op.
"""
self._resuming = False
if not self.config.get("checkpoint_enabled"):
return None
signature = self._run_signature(asset_type, portfolio)
self._checkpointer_ctx = get_checkpointer(self.config["data_cache_dir"], company_name)
saver = self._checkpointer_ctx.__enter__()
self.graph = self.workflow.compile(checkpointer=saver)
step = checkpoint_step(
self.config["data_cache_dir"], company_name, str(trade_date), signature
)
self._resuming = step is not None
if step is not None:
logger.info("Resuming from step %d for %s on %s", step, company_name, trade_date)
else:
logger.info("Starting fresh for %s on %s", company_name, trade_date)
return thread_id(company_name, str(trade_date), signature)
def checkpoint_input(self, init_state):
"""The value to stream/invoke: ``None`` to resume an existing checkpoint,
else the initial state for a fresh run.
LangGraph resumes an interrupted thread when invoked with ``None``;
re-passing the initial state instead appends it through the message
reducer, duplicating messages in the resumed state (#1249).
"""
return None if self._resuming else init_state
def end_checkpoint(self):
"""Restore the plain uncheckpointed graph after a checkpointed run."""
if self._checkpointer_ctx is not None:
self._checkpointer_ctx.__exit__(None, None, None)
self._checkpointer_ctx = None
self.graph = self.workflow.compile()
self._resuming = False
@contextmanager
def checkpoint_scope(self, company_name, trade_date, asset_type: str = "stock", portfolio=None):
"""Context-manager form of begin/end_checkpoint for the propagate path."""
try:
yield self.begin_checkpoint(company_name, trade_date, asset_type, portfolio)
finally:
self.end_checkpoint()
def clear_checkpoint_on_success(self, company_name, trade_date, asset_type: str = "stock", portfolio=None):
"""Drop a completed run's checkpoint so a later run starts fresh (#1249)."""
if self.config.get("checkpoint_enabled"):
clear_checkpoint(
self.config["data_cache_dir"], company_name, str(trade_date),
self._run_signature(asset_type, portfolio),
)
def run_settings(self) -> dict:
"""What produces this graph's runs, for the saved report and state log.
An allowlist: endpoints (a backend_url can carry credentials), keys and
local paths are never recorded.
"""
cfg = self.config
return {
"version": tradingagents.__version__,
"llm_provider": cfg.get("llm_provider"),
"deep_think_llm": cfg.get("deep_think_llm"),
"quick_think_llm": cfg.get("quick_think_llm"),
"analysts": list(self.selected_analysts),
"max_debate_rounds": cfg.get("max_debate_rounds"),
"max_risk_discuss_rounds": cfg.get("max_risk_discuss_rounds"),
"output_language": cfg.get("output_language"),
"data_vendors": dict(cfg.get("data_vendors") or {}),
"tool_vendors": dict(cfg.get("tool_vendors") or {}),
}
def save_reports(self, final_state, ticker, save_path=None) -> Path:
"""Write the markdown report tree for a completed run, like the CLI does.
Programmatic callers get the same on-disk reports the CLI produces. Pass
an explicit ``save_path`` or let it default under ``results_dir``.
"""
if save_path is None:
stamp = datetime.now().strftime("%Y%m%d_%H%M%S")
save_path = (
Path(self.config["results_dir"])
/ "reports"
/ f"{safe_ticker_component(ticker)}_{stamp}"
)
return write_report_tree(final_state, ticker, save_path, settings=self.run_settings())
def create_run_state(self, company_name, trade_date, asset_type: str = "stock", portfolio=None):
"""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 memory log.
"""
self.settle_pending(company_name)
return self.propagator.create_initial_state(
company_name,
trade_date,
asset_type=asset_type,
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, trade_date),
portfolio_context=portfolio.render(company_name) if portfolio is not None else "",
)
def settle_pending(self, company_name):
"""Settle this ticker's decisions whose holding window has now traded.
A run settles the ticker's earlier decisions on its way in, so the most
recent one stays pending until the next run for that ticker. A caller
that is done analyzing a ticker (a backtest sweep, a scheduled job) calls
this to settle it now.
"""
with run_config(self.config):
settlement.settle_pending(company_name, self.memory_log, self.reflector, self.config)
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,
rating=run_rating(final_state),
)
def _run_graph(self, company_name, trade_date, asset_type: str = "stock",
checkpoint_thread_id: str | None = None, portfolio=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, portfolio)
args = self.propagator.get_graph_args()
# Inject the checkpoint thread_id (from checkpoint_scope) so the same
# ticker+date+graph-shape resumes; a different one starts fresh (#1089).
if checkpoint_thread_id is not None:
args.setdefault("config", {}).setdefault("configurable", {})["thread_id"] = checkpoint_thread_id
# None resumes an existing checkpoint; init_agent_state starts fresh (#1249).
graph_input = self.checkpoint_input(init_agent_state)
if self.debug:
# A state repeats the messages before it, so each prints once (#1027).
final_state, printed = {}, set()
for messages, state in self.stream_run(graph_input, **args):
for msg in messages:
key = getattr(msg, "id", None) or (type(msg).__name__, getattr(msg, "content", None))
if key not in printed:
printed.add(key)
msg.pretty_print()
if state is not None:
final_state.update(state)
else:
final_state = self.graph.invoke(graph_input, **args)
# Log state to disk.
self._log_state(trade_date, final_state)
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, portfolio)
return final_state, run_rating(final_state)
def stream_run(self, graph_input, **args):
"""Stream a run as ``(messages, state)`` pairs.
``messages`` are the agents' messages, the analysts' included. ``state``
is the run's state after a top-level step; for a step inside an analyst's
graph it is that analyst's report once filed, else None.
Each analyst works in a graph of its own, and the run's state takes the
analysts' reports only when the slowest has finished, so their messages
and reports come from their own finished steps ("tasks") as they happen.
"""
args = {**args, "stream_mode": ["values", "tasks"]}
for namespace, mode, chunk in self.graph.stream(graph_input, subgraphs=True, **args):
if namespace:
result = chunk.get("result") if mode == "tasks" else None
if isinstance(result, dict):
report = {k: v for k, v in result.items() if k != "messages" and v}
if result.get("messages") or report:
yield result.get("messages", []), report or None
elif mode == "values":
yield chunk.get("messages", []), chunk
def _log_state(self, trade_date, final_state):
"""Write a run's final state to JSON under the run's own ticker."""
entry = {
"company_of_interest": final_state["company_of_interest"],
"trade_date": final_state["trade_date"],
"market_report": final_state["market_report"],
"sentiment_report": final_state["sentiment_report"],
"news_report": final_state["news_report"],
"fundamentals_report": final_state["fundamentals_report"],
"investment_debate_state": {
"bull_history": final_state["investment_debate_state"]["bull_history"],
"bear_history": final_state["investment_debate_state"]["bear_history"],
"history": final_state["investment_debate_state"]["history"],
"current_response": final_state["investment_debate_state"][
"current_response"
],
},
"trader_investment_plan": final_state["trader_investment_plan"],
"risk_debate_state": {
"aggressive_history": final_state["risk_debate_state"]["aggressive_history"],
"conservative_history": final_state["risk_debate_state"]["conservative_history"],
"neutral_history": final_state["risk_debate_state"]["neutral_history"],
"history": final_state["risk_debate_state"]["history"],
},
"investment_plan": final_state["investment_plan"],
"final_trade_decision": final_state["final_trade_decision"],
"final_rating": run_rating(final_state),
"run_settings": self.run_settings(),
}
# A ticker that would escape the results directory is rejected.
safe_ticker = safe_ticker_component(final_state["company_of_interest"])
directory = Path(self.config["results_dir"]) / safe_ticker / "TradingAgentsStrategy_logs"
directory.mkdir(parents=True, exist_ok=True)
log_path = directory / f"full_states_log_{trade_date}.json"
with open(log_path, "w", encoding="utf-8") as f:
# Reports can be in any language and this file is read by a person.
json.dump(entry, f, indent=4, ensure_ascii=False)