test: run the whole graph end to end with scripted models

- every analyst calls its tools through the router, offline; the decision is parsed and logged
- free-text and structured-output paths, and resume from a checkpoint
- remove the unused mock_llm_client fixture
This commit is contained in:
Yijia-Xiao
2026-09-24 00:41:56 +00:00
parent 632fb51d18
commit 1e439352c0
2 changed files with 155 additions and 12 deletions
-12
View File
@@ -1,7 +1,6 @@
"""Shared pytest fixtures that prevent CI hangs when API keys are absent.""" """Shared pytest fixtures that prevent CI hangs when API keys are absent."""
import os import os
from unittest.mock import MagicMock, patch
import pytest import pytest
@@ -75,14 +74,3 @@ def _isolate_config():
config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG) config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG)
yield yield
config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG) config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG)
@pytest.fixture()
def mock_llm_client():
client = MagicMock()
client.get_llm.return_value = MagicMock()
with patch(
"tradingagents.llm_clients.factory.create_llm_client",
return_value=client,
):
yield client
+155
View File
@@ -0,0 +1,155 @@
"""The whole graph, end to end, with scripted models and no network.
Every tool-using analyst calls each of its tools once through the vendor router,
the debates and managers run, and the decision is parsed and logged. This pins
the wiring: a restructure that drops a node, a tool or an edge fails here.
"""
from __future__ import annotations
import copy
import pandas as pd
import pytest
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.messages import AIMessage, ToolMessage
from langchain_core.outputs import ChatGeneration, ChatResult
from langchain_core.runnables import RunnableLambda
from pydantic import Field
from tradingagents.agents import schemas
from tradingagents.agents.analysts import sentiment_analyst
from tradingagents.agents.utils import agent_utils
from tradingagents.dataflows import interface, market_data_validator
from tradingagents.default_config import DEFAULT_CONFIG
from tradingagents.graph import trading_graph
TRADE_DATE = "2026-01-09"
# Enough for every free-text reader: the PM's labelled rating and the trader's
# closing proposal line.
TEXT = "Report.\n\n**Rating**: Overweight\n\nFINAL TRANSACTION PROPOSAL: **BUY**"
STRUCTURED = {
schemas.ResearchPlan: schemas.ResearchPlan(
recommendation=schemas.PortfolioRating.OVERWEIGHT, rationale="r", strategic_actions="a"),
schemas.TraderProposal: schemas.TraderProposal(action=schemas.TraderAction.BUY, reasoning="r"),
schemas.PortfolioDecision: schemas.PortfolioDecision(
rating=schemas.PortfolioRating.OVERWEIGHT, executive_summary="s", investment_thesis="t"),
schemas.SentimentReport: schemas.SentimentReport(
overall_band=schemas.SentimentBand.NEUTRAL, overall_score=5.0, confidence="low", narrative="n"),
}
ARGS = {"symbol": "NVDA", "ticker": "NVDA", "curr_date": TRADE_DATE, "start_date": "2026-01-02",
"end_date": TRADE_DATE, "indicator": "rsi", "topic": "Fed rate cut", "freq": "quarterly"}
class ScriptedModel(BaseChatModel):
"""Calls every bound tool once, then answers with TEXT."""
structured: bool = False
tools: tuple = ()
calls: list = Field(default_factory=list) # shared across bound copies
fail_at: int | None = None # raise on this call, once
@property
def _llm_type(self) -> str:
return "scripted"
def bind_tools(self, tools, **kwargs):
return self.model_copy(update={"tools": tuple(tools)})
def with_structured_output(self, schema, **kwargs):
if not self.structured:
raise NotImplementedError
return RunnableLambda(lambda _: self._count() or STRUCTURED[schema])
def _count(self) -> None:
self.calls.append(1)
if len(self.calls) == self.fail_at:
raise RuntimeError("provider unavailable")
def _generate(self, messages, stop=None, run_manager=None, **kwargs) -> ChatResult:
self._count()
if self.tools and not isinstance(messages[-1], ToolMessage):
calls = [{"name": t.name, "id": f"call_{i}",
"args": {k: v for k, v in ARGS.items()
if k in t.tool_call_schema.model_json_schema()["properties"]}}
for i, t in enumerate(self.tools)]
message = AIMessage(content="", tool_calls=calls)
else:
message = AIMessage(content=TEXT)
return ChatResult(generations=[ChatGeneration(message=message)])
class _Client:
def __init__(self, model):
self.model = model
def get_llm(self):
return self.model
@pytest.fixture
def offline(monkeypatch, tmp_path):
"""Every vendor answers offline; returns the set of router methods called."""
called: set[str] = set()
for method, vendors in interface.VENDOR_METHODS.items():
for vendor in vendors:
monkeypatch.setitem(vendors, vendor,
lambda *a, _m=method, **k: called.add(_m) or f"{_m} data")
prices = pd.DataFrame({
"Date": pd.bdate_range(end=TRADE_DATE, periods=60),
"Open": 100.0, "High": 101.0, "Low": 99.0, "Close": 100.5, "Volume": 1_000_000,
})
monkeypatch.setattr(market_data_validator, "load_ohlcv",
lambda *a, **k: called.add("ohlcv") or prices.copy())
monkeypatch.setattr(sentiment_analyst, "fetch_stocktwits_messages", lambda *a, **k: "no posts")
monkeypatch.setattr(sentiment_analyst, "fetch_reddit_posts", lambda *a, **k: "no posts")
monkeypatch.setattr(agent_utils.yf, "Ticker", lambda s: type("T", (), {"info": {"longName": "NVIDIA"}})())
agent_utils.resolve_instrument_identity.cache_clear()
return called
def _graph(tmp_path, monkeypatch, model, **config):
cfg = copy.deepcopy(DEFAULT_CONFIG)
cfg.update(results_dir=str(tmp_path / "results"), data_cache_dir=str(tmp_path / "cache"),
memory_log_path=str(tmp_path / "log.md"), **config)
monkeypatch.setattr(trading_graph, "create_llm_client", lambda **k: _Client(model))
return trading_graph.TradingAgentsGraph(config=cfg)
@pytest.mark.unit
@pytest.mark.parametrize("structured", [False, True], ids=["free-text", "structured"])
def test_a_full_run_reaches_a_logged_decision(tmp_path, monkeypatch, offline, structured):
graph = _graph(tmp_path, monkeypatch, ScriptedModel(structured=structured))
state, signal = graph.propagate("NVDA", TRADE_DATE)
assert signal == "Overweight"
for key in ("market_report", "sentiment_report", "news_report", "fundamentals_report",
"investment_plan", "trader_investment_plan", "final_trade_decision"):
assert state[key].strip(), key
tool_methods = {"get_stock_data", "get_indicators", "get_news", "get_global_news",
"get_macro_indicators", "get_prediction_markets", "get_fundamentals",
"get_balance_sheet", "get_cashflow", "get_income_statement", "ohlcv"}
assert offline == tool_methods
assert [e["rating"] for e in graph.memory_log.load_entries()] == ["Overweight"]
@pytest.mark.unit
def test_an_interrupted_run_resumes_from_its_checkpoint(tmp_path, monkeypatch, offline):
model = ScriptedModel(fail_at=12) # past the analysts, before the decision
graph = _graph(tmp_path, monkeypatch, model, checkpoint_enabled=True)
with pytest.raises(RuntimeError, match="provider unavailable"):
graph.propagate("NVDA", TRADE_DATE)
calls_before = len(model.calls)
_, signal = graph.propagate("NVDA", TRADE_DATE)
assert signal == "Overweight"
resumed_calls = len(model.calls) - calls_before
full_run = ScriptedModel()
_graph(tmp_path / "fresh", monkeypatch, full_run).propagate("NVDA", TRADE_DATE)
# The resumed run makes only the calls the interrupted one had not completed.
assert resumed_calls == len(full_run.calls) - (model.fail_at - 1)