From 1e439352c0cc9b7b781cb82fdec955be9c6243f6 Mon Sep 17 00:00:00 2001 From: Yijia-Xiao Date: Wed, 23 Sep 2026 21:05:38 +0000 Subject: [PATCH] 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 --- tests/conftest.py | 12 --- tests/test_graph_end_to_end.py | 155 +++++++++++++++++++++++++++++++++ 2 files changed, 155 insertions(+), 12 deletions(-) create mode 100644 tests/test_graph_end_to_end.py diff --git a/tests/conftest.py b/tests/conftest.py index 00e542461..9a8e3cd59 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,7 +1,6 @@ """Shared pytest fixtures that prevent CI hangs when API keys are absent.""" import os -from unittest.mock import MagicMock, patch import pytest @@ -75,14 +74,3 @@ def _isolate_config(): config_module._config = copy.deepcopy(default_config.DEFAULT_CONFIG) yield 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 diff --git a/tests/test_graph_end_to_end.py b/tests/test_graph_end_to_end.py new file mode 100644 index 000000000..6790afc39 --- /dev/null +++ b/tests/test_graph_end_to_end.py @@ -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)