mirror of
https://github.com/TauricResearch/TradingAgents.git
synced 2026-09-27 06:56:39 +03:00
refactor(dataflows): read company identity and settlement prices through the data layer
- the Yahoo vendor owns get_company_profile and get_closes; agent_utils and the graph no longer import yfinance - a test keeps vendor libraries inside dataflows
This commit is contained in:
@@ -20,7 +20,7 @@ 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.dataflows import interface, market_data_validator, y_finance
|
||||
from tradingagents.default_config import DEFAULT_CONFIG
|
||||
from tradingagents.graph import trading_graph
|
||||
|
||||
@@ -106,7 +106,7 @@ def offline(monkeypatch, tmp_path):
|
||||
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"}})())
|
||||
monkeypatch.setattr(y_finance.yf, "Ticker", lambda s: type("T", (), {"info": {"longName": "NVIDIA"}})())
|
||||
agent_utils.resolve_instrument_identity.cache_clear()
|
||||
return called
|
||||
|
||||
|
||||
@@ -21,7 +21,7 @@ class ResolveInstrumentIdentityTests(unittest.TestCase):
|
||||
resolve_instrument_identity.cache_clear()
|
||||
|
||||
def test_resolves_company_metadata_from_yfinance(self):
|
||||
with patch("tradingagents.agents.utils.agent_utils.yf.Ticker") as mock:
|
||||
with patch("tradingagents.dataflows.y_finance.yf.Ticker") as mock:
|
||||
mock.return_value.info = {
|
||||
"longName": "TOTO LTD.",
|
||||
"shortName": "TOTO",
|
||||
@@ -38,26 +38,26 @@ class ResolveInstrumentIdentityTests(unittest.TestCase):
|
||||
self.assertEqual(identity["exchange"], "PNK")
|
||||
|
||||
def test_falls_back_to_short_name(self):
|
||||
with patch("tradingagents.agents.utils.agent_utils.yf.Ticker") as mock:
|
||||
with patch("tradingagents.dataflows.y_finance.yf.Ticker") as mock:
|
||||
mock.return_value.info = {"shortName": "TOTO", "sector": "Industrials"}
|
||||
identity = resolve_instrument_identity("TOTDY")
|
||||
self.assertEqual(identity["company_name"], "TOTO")
|
||||
|
||||
def test_skips_placeholder_values(self):
|
||||
with patch("tradingagents.agents.utils.agent_utils.yf.Ticker") as mock:
|
||||
with patch("tradingagents.dataflows.y_finance.yf.Ticker") as mock:
|
||||
mock.return_value.info = {"longName": " ", "sector": "None", "industry": "n/a"}
|
||||
identity = resolve_instrument_identity("TOTDY")
|
||||
self.assertEqual(identity, {})
|
||||
|
||||
def test_fails_open_on_exception(self):
|
||||
with patch(
|
||||
"tradingagents.agents.utils.agent_utils.yf.Ticker",
|
||||
"tradingagents.dataflows.y_finance.yf.Ticker",
|
||||
side_effect=RuntimeError("rate limited"),
|
||||
):
|
||||
self.assertEqual(resolve_instrument_identity("TOTDY"), {})
|
||||
|
||||
def test_result_is_cached(self):
|
||||
with patch("tradingagents.agents.utils.agent_utils.yf.Ticker") as mock:
|
||||
with patch("tradingagents.dataflows.y_finance.yf.Ticker") as mock:
|
||||
mock.return_value.info = {"longName": "TOTO LTD."}
|
||||
first = resolve_instrument_identity("TOTDY")
|
||||
second = resolve_instrument_identity("TOTDY")
|
||||
@@ -104,7 +104,7 @@ class GetInstrumentContextFromStateTests(unittest.TestCase):
|
||||
|
||||
def test_fallback_is_network_free_ticker_only(self):
|
||||
# No instrument_context and no yfinance call — must not hit the network.
|
||||
with patch("tradingagents.agents.utils.agent_utils.yf.Ticker") as mock:
|
||||
with patch("tradingagents.dataflows.y_finance.yf.Ticker") as mock:
|
||||
context = get_instrument_context_from_state(
|
||||
{"company_of_interest": "NVDA", "asset_type": "stock"}
|
||||
)
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
"""Only the data layer imports vendor libraries.
|
||||
|
||||
Vendor calls belong in dataflows, where failures are raised as VendorError
|
||||
subclasses; a call made elsewhere can report an outage as a fact about the market.
|
||||
"""
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
VENDOR_LIBRARIES = {"yfinance"}
|
||||
|
||||
|
||||
def _imports(path: Path) -> set[str]:
|
||||
names = set()
|
||||
for node in ast.walk(ast.parse(path.read_text(encoding="utf-8"))):
|
||||
if isinstance(node, ast.Import):
|
||||
names |= {a.name.split(".")[0] for a in node.names}
|
||||
elif isinstance(node, ast.ImportFrom) and node.module and not node.level:
|
||||
names.add(node.module.split(".")[0])
|
||||
return names
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_vendor_libraries_are_imported_only_by_the_data_layer():
|
||||
data_layer = ROOT / "tradingagents" / "dataflows"
|
||||
offenders = sorted(
|
||||
str(path.relative_to(ROOT))
|
||||
for package in ("tradingagents", "cli")
|
||||
for path in (ROOT / package).rglob("*.py")
|
||||
if data_layer not in path.parents and _imports(path) & VENDOR_LIBRARIES
|
||||
)
|
||||
assert offenders == []
|
||||
@@ -1043,7 +1043,7 @@ def test_a_longer_window_asks_for_enough_price_history(monkeypatch):
|
||||
days = pd.bdate_range(start, end)
|
||||
return pd.DataFrame({"Close": range(len(days))}, index=days)
|
||||
|
||||
monkeypatch.setattr("tradingagents.graph.trading_graph.yf.Ticker", _Ticker)
|
||||
monkeypatch.setattr("tradingagents.dataflows.y_finance.yf.Ticker", _Ticker)
|
||||
|
||||
raw, alpha, days, resolved = graph._fetch_returns("NVDA", "2026-06-01", 21, benchmark="SPY")
|
||||
|
||||
|
||||
@@ -8,8 +8,8 @@ hit the right instrument instead of failing/mismatching.
|
||||
import pandas as pd
|
||||
|
||||
import tradingagents.agents.utils.agent_utils as au
|
||||
import tradingagents.dataflows.y_finance as y_finance
|
||||
import tradingagents.dataflows.yfinance_news as ynews
|
||||
import tradingagents.graph.trading_graph as tg
|
||||
from tradingagents.graph.trading_graph import TradingAgentsGraph
|
||||
|
||||
|
||||
@@ -24,7 +24,7 @@ def test_identity_lookup_normalizes_symbol(monkeypatch):
|
||||
def info(self):
|
||||
return {"longName": "Gold Futures", "quoteType": "FUTURE"}
|
||||
|
||||
monkeypatch.setattr(au.yf, "Ticker", FakeTicker)
|
||||
monkeypatch.setattr(y_finance.yf, "Ticker", FakeTicker)
|
||||
au.resolve_instrument_identity.cache_clear()
|
||||
|
||||
identity = au.resolve_instrument_identity("XAUUSD")
|
||||
@@ -45,7 +45,7 @@ def test_fetch_returns_normalizes_symbol(monkeypatch):
|
||||
idx = pd.date_range(start="2025-01-02", periods=len(prices), freq="D")
|
||||
return pd.DataFrame({"Close": prices}, index=idx)
|
||||
|
||||
monkeypatch.setattr(tg.yf, "Ticker", FakeTicker)
|
||||
monkeypatch.setattr(y_finance.yf, "Ticker", FakeTicker)
|
||||
|
||||
# _fetch_returns does not use ``self``; call unbound to avoid building the graph.
|
||||
raw, alpha, days, resolved = TradingAgentsGraph._fetch_returns(
|
||||
|
||||
Reference in New Issue
Block a user