mirror of
https://github.com/TauricResearch/TradingAgents.git
synced 2026-09-27 06:56:39 +03:00
- 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
423 lines
20 KiB
Python
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)
|