refactor(dataflows): name the shared modules by what they hold

- interface -> router; symbol_utils -> symbols, which also takes safe_ticker_component
- utils is split: get_current_date to date_window, the HTTP helpers to net
- dataflows imports are absolute; the NoMarketDataError re-export from symbols is gone
This commit is contained in:
Yijia-Xiao
2026-09-24 04:31:05 +00:00
parent a58aa613fc
commit c42a2f2c61
47 changed files with 217 additions and 219 deletions
+1 -2
View File
@@ -22,6 +22,7 @@ from tradingagents.agents.utils.news_data_tools import (
)
from tradingagents.agents.utils.prediction_markets_tools import get_prediction_markets
from tradingagents.agents.utils.technical_indicators_tools import get_indicators
from tradingagents.dataflows.date_window import get_current_date
from tradingagents.dataflows.y_finance import get_company_profile
# Public surface: the data tools are imported here so agents and the graph
@@ -48,8 +49,6 @@ __all__ = [
logger = logging.getLogger(__name__)
from tradingagents.dataflows.utils import get_current_date # noqa: E402
def get_language_instruction() -> str:
"""Return a prompt instruction for the configured output language.
@@ -4,7 +4,7 @@ from langchain_core.tools import tool
from langgraph.prebuilt import InjectedState
from tradingagents.dataflows.date_window import as_of_window
from tradingagents.dataflows.interface import route_to_vendor
from tradingagents.dataflows.router import route_to_vendor
@tool
@@ -4,7 +4,7 @@ from langchain_core.tools import tool
from langgraph.prebuilt import InjectedState
from tradingagents.dataflows.date_window import as_of
from tradingagents.dataflows.interface import route_to_vendor
from tradingagents.dataflows.router import route_to_vendor
@tool
@@ -4,7 +4,7 @@ from langchain_core.tools import tool
from langgraph.prebuilt import InjectedState
from tradingagents.dataflows.date_window import as_of
from tradingagents.dataflows.interface import route_to_vendor
from tradingagents.dataflows.router import route_to_vendor
@tool
@@ -4,7 +4,7 @@ from langchain_core.tools import tool
from langgraph.prebuilt import InjectedState
from tradingagents.dataflows.date_window import as_of, as_of_window
from tradingagents.dataflows.interface import route_to_vendor
from tradingagents.dataflows.router import route_to_vendor
@tool
@@ -3,7 +3,7 @@ from typing import Annotated
from langchain_core.tools import tool
from langgraph.prebuilt import InjectedState
from tradingagents.dataflows.interface import route_to_vendor
from tradingagents.dataflows.router import route_to_vendor
@tool
@@ -4,7 +4,7 @@ from langchain_core.tools import tool
from langgraph.prebuilt import InjectedState
from tradingagents.dataflows.date_window import as_of
from tradingagents.dataflows.interface import route_to_vendor
from tradingagents.dataflows.router import route_to_vendor
@tool
+2 -1
View File
@@ -23,7 +23,8 @@ from pathlib import Path
from tradingagents.agents.utils.memory import TradingMemoryLog
from tradingagents.agents.utils.rating import RATING_REVIEW
from tradingagents.dataflows.utils import get_current_date, safe_ticker_component
from tradingagents.dataflows.date_window import get_current_date
from tradingagents.dataflows.symbols import safe_ticker_component
from tradingagents.graph.trading_graph import TradingAgentsGraph
logger = logging.getLogger(__name__)
+8 -4
View File
@@ -1,14 +1,18 @@
# Aggregates the per-category Alpha Vantage implementations into one module the
# vendor router imports from; the imports below are the public surface.
from .alpha_vantage_fundamentals import (
from tradingagents.dataflows.alpha_vantage_fundamentals import (
get_balance_sheet,
get_cashflow,
get_fundamentals,
get_income_statement,
)
from .alpha_vantage_indicator import get_indicator
from .alpha_vantage_news import get_global_news, get_insider_transactions, get_news
from .alpha_vantage_stock import get_stock
from tradingagents.dataflows.alpha_vantage_indicator import get_indicator
from tradingagents.dataflows.alpha_vantage_news import (
get_global_news,
get_insider_transactions,
get_news,
)
from tradingagents.dataflows.alpha_vantage_stock import get_stock
__all__ = [
"get_balance_sheet",
@@ -5,8 +5,8 @@ from io import StringIO
import pandas as pd
from .errors import VendorNotConfiguredError, VendorRateLimitError
from .utils import get_scrubbed
from tradingagents.dataflows.errors import VendorNotConfiguredError, VendorRateLimitError
from tradingagents.dataflows.net import get_scrubbed
API_BASE_URL = "https://www.alphavantage.co/query"
@@ -1,7 +1,7 @@
import json
from .alpha_vantage_common import _make_api_request
from .date_window import withhold_live_profile
from tradingagents.dataflows.alpha_vantage_common import _make_api_request
from tradingagents.dataflows.date_window import withhold_live_profile
def _filter_reports_by_date(result, curr_date: str):
@@ -1,7 +1,7 @@
import logging
from .alpha_vantage_common import _make_api_request
from .errors import NoMarketDataError, VendorError
from tradingagents.dataflows.alpha_vantage_common import _make_api_request
from tradingagents.dataflows.errors import NoMarketDataError, VendorError
logger = logging.getLogger(__name__)
@@ -1,7 +1,7 @@
import json
from .alpha_vantage_common import _make_api_request, format_datetime_for_api
from .config import get_config
from tradingagents.dataflows.alpha_vantage_common import _make_api_request, format_datetime_for_api
from tradingagents.dataflows.config import get_config
def get_news(ticker, start_date, end_date) -> dict[str, str] | str:
@@ -1,6 +1,9 @@
from datetime import datetime
from .alpha_vantage_common import _filter_csv_by_date_range, _make_api_request
from tradingagents.dataflows.alpha_vantage_common import (
_filter_csv_by_date_range,
_make_api_request,
)
def get_stock(
+6 -3
View File
@@ -11,9 +11,7 @@ in a backtest we can't prove it isn't future.
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from .utils import get_current_date
from datetime import date, datetime, timedelta, timezone
def to_utc(dt: datetime) -> datetime:
@@ -32,6 +30,11 @@ def in_window(pub_dt: datetime | None, start_dt: datetime, end_dt: datetime) ->
return end >= datetime.now(timezone.utc) - timedelta(days=1)
def get_current_date() -> str:
"""Today's date, YYYY-MM-DD."""
return date.today().strftime("%Y-%m-%d")
def coverage_gap(
dates, start_date: str, end_date: str, source: str, subject: str
) -> str | None:
+2 -2
View File
@@ -14,8 +14,8 @@ from datetime import datetime, timedelta
import pytz
from .errors import VendorNotConfiguredError
from .utils import get_scrubbed
from tradingagents.dataflows.errors import VendorNotConfiguredError
from tradingagents.dataflows.net import get_scrubbed
logger = logging.getLogger(__name__)
+37
View File
@@ -0,0 +1,37 @@
"""HTTP helpers shared by the vendors."""
import requests
def get_scrubbed(url: str, *, params: dict, timeout: float, secret: str, passthrough=()):
"""``requests.get`` plus ``raise_for_status``, with ``secret`` kept out of errors.
Vendors that authenticate with a query parameter put the key in the URL, and
requests quotes the full URL in HTTP, connection and timeout errors, so any
log or traceback that records one would carry the key (#1324). A requests
error is re-raised as the same class with the key replaced and nothing
attached: no request or response (both hold the URL) and no exception chain,
which is why this raises after the ``except`` block rather than inside it.
Statuses in ``passthrough`` are returned for the caller to handle.
"""
try:
response = requests.get(url, params=params, timeout=timeout)
if response.status_code not in passthrough:
response.raise_for_status()
return response
except requests.RequestException as exc:
error = type(exc)(str(exc).replace(secret, "***")) if secret else exc
raise error
def vendor_reachable(url: str, timeout: float = 5.0) -> bool:
"""Whether the vendor answers at all, for telling silence from an outage.
A client that returns an empty result instead of raising leaves those two
cases indistinguishable. Called only when a result is empty.
"""
try:
requests.head(url, timeout=timeout, allow_redirects=True)
return True
except requests.RequestException:
return False
+1 -1
View File
@@ -15,7 +15,7 @@ from datetime import datetime, timezone
import requests
from .utils import get_current_date
from tradingagents.dataflows.date_window import get_current_date
logger = logging.getLogger(__name__)
+2 -2
View File
@@ -29,8 +29,8 @@ from urllib.error import HTTPError
from urllib.parse import urlencode
from urllib.request import Request, urlopen
from .date_window import coverage_gap, in_window
from .symbol_utils import crypto_base
from tradingagents.dataflows.date_window import coverage_gap, in_window
from tradingagents.dataflows.symbols import crypto_base
logger = logging.getLogger(__name__)
@@ -1,6 +1,6 @@
import logging
from .alpha_vantage import (
from tradingagents.dataflows.alpha_vantage import (
get_balance_sheet as get_alpha_vantage_balance_sheet,
get_cashflow as get_alpha_vantage_cashflow,
get_fundamentals as get_alpha_vantage_fundamentals,
@@ -11,20 +11,22 @@ from .alpha_vantage import (
get_news as get_alpha_vantage_news,
get_stock as get_alpha_vantage_stock,
)
from .config import get_config
from .errors import (
from tradingagents.dataflows.config import get_config
from tradingagents.dataflows.errors import (
NoMarketDataError,
VendorNotConfiguredError,
VendorRateLimitError,
)
from .fred import get_macro_data as get_fred_macro_data
from .polymarket import get_prediction_markets as get_polymarket_prediction_markets
from .sec_edgar import (
from tradingagents.dataflows.fred import get_macro_data as get_fred_macro_data
from tradingagents.dataflows.polymarket import (
get_prediction_markets as get_polymarket_prediction_markets,
)
from tradingagents.dataflows.sec_edgar import (
get_balance_sheet as get_sec_edgar_balance_sheet,
get_cashflow as get_sec_edgar_cashflow,
get_income_statement as get_sec_edgar_income_statement,
)
from .y_finance import (
from tradingagents.dataflows.y_finance import (
get_balance_sheet as get_yfinance_balance_sheet,
get_cashflow as get_yfinance_cashflow,
get_fundamentals as get_yfinance_fundamentals,
@@ -33,7 +35,7 @@ from .y_finance import (
get_stock_stats_indicators_window,
get_YFin_data_online,
)
from .yfinance_news import get_global_news_yfinance, get_news_yfinance
from tradingagents.dataflows.yfinance_news import get_global_news_yfinance, get_news_yfinance
logger = logging.getLogger(__name__)
+2 -2
View File
@@ -27,8 +27,8 @@ from pathlib import Path
import requests
from .config import get_config
from .errors import NoMarketDataError, VendorRateLimitError
from tradingagents.dataflows.config import get_config
from tradingagents.dataflows.errors import NoMarketDataError, VendorRateLimitError
logger = logging.getLogger(__name__)
+4 -4
View File
@@ -8,10 +8,10 @@ import yfinance as yf
from stockstats import wrap
from yfinance.exceptions import YFRateLimitError
from .config import get_config
from .errors import VendorRateLimitError
from .symbol_utils import NoMarketDataError, normalize_symbol
from .utils import safe_ticker_component, vendor_reachable
from tradingagents.dataflows.config import get_config
from tradingagents.dataflows.errors import NoMarketDataError, VendorRateLimitError
from tradingagents.dataflows.net import vendor_reachable
from tradingagents.dataflows.symbols import normalize_symbol, safe_ticker_component
logger = logging.getLogger(__name__)
+2 -2
View File
@@ -21,8 +21,8 @@ import logging
from datetime import datetime
from urllib.request import Request, urlopen
from .date_window import coverage_gap, in_window
from .symbol_utils import crypto_base
from tradingagents.dataflows.date_window import coverage_gap, in_window
from tradingagents.dataflows.symbols import crypto_base
logger = logging.getLogger(__name__)
@@ -1,4 +1,4 @@
"""Symbol normalization and market-data error types for vendor calls.
"""Symbol normalization for vendor calls, and ticker values safe to use in a path.
Yahoo Finance (the default vendor) uses specific ticker conventions that
differ from the broker / TradingView / MT5 style symbols users often type:
@@ -25,10 +25,6 @@ from __future__ import annotations
import logging
import re
# NoMarketDataError lives in the vendor-error taxonomy (errors.py); re-exported
# here for the many call sites that import it alongside normalize_symbol.
from .errors import NoMarketDataError as NoMarketDataError
logger = logging.getLogger(__name__)
@@ -147,3 +143,38 @@ def normalize_symbol(raw: str) -> str:
logger.info("Resolved symbol %r to Yahoo symbol %r", raw, canonical)
return canonical
# Tickers can contain letters, digits, dot, dash, underscore, caret
# (index symbols like ^GSPC), equals (futures like GC=F), and plus
# (forex/CFD symbols like XAUUSD+). None of these enable directory
# traversal, so the value never escapes a containing directory when
# interpolated into a path. Anything else is rejected.
_TICKER_PATH_RE = re.compile(r"^[A-Za-z0-9._\-\^=+]+$")
def safe_ticker_component(value: str, *, max_len: int = 32) -> str:
"""Validate ``value`` is safe to interpolate into a filesystem path.
Tickers come from user CLI input or from LLM tool calls, both of which
can be influenced by attacker-controlled content (e.g. prompt injection
embedded in fetched news). Without validation, a value like
``"../../../etc/foo"`` flows into ``os.path.join`` / ``Path /`` and
escapes the configured cache, checkpoint, or results directory.
Returns ``value`` unchanged when it matches the allowed pattern; raises
``ValueError`` otherwise.
"""
if not isinstance(value, str) or not value:
raise ValueError(f"ticker must be a non-empty string, got {value!r}")
if len(value) > max_len:
raise ValueError(f"ticker exceeds {max_len} chars: {value!r}")
if not _TICKER_PATH_RE.fullmatch(value):
raise ValueError(
f"ticker contains characters not allowed in a filesystem path: {value!r}"
)
# The regex above allows '.', so values like '.', '..', '...' would pass,
# and as a path component they traverse the parent directory. Reject any
# value that's only dots.
if set(value) == {"."}:
raise ValueError(f"ticker cannot consist solely of dots: {value!r}")
return value
-77
View File
@@ -1,77 +0,0 @@
import re
from datetime import date
import requests
# Tickers can contain letters, digits, dot, dash, underscore, caret
# (index symbols like ^GSPC), equals (futures like GC=F), and plus
# (forex/CFD symbols like XAUUSD+). None of these enable directory
# traversal, so the value never escapes a containing directory when
# interpolated into a path. Anything else is rejected.
_TICKER_PATH_RE = re.compile(r"^[A-Za-z0-9._\-\^=+]+$")
def safe_ticker_component(value: str, *, max_len: int = 32) -> str:
"""Validate ``value`` is safe to interpolate into a filesystem path.
Tickers come from user CLI input or from LLM tool calls, both of which
can be influenced by attacker-controlled content (e.g. prompt injection
embedded in fetched news). Without validation, a value like
``"../../../etc/foo"`` flows into ``os.path.join`` / ``Path /`` and
escapes the configured cache, checkpoint, or results directory.
Returns ``value`` unchanged when it matches the allowed pattern; raises
``ValueError`` otherwise.
"""
if not isinstance(value, str) or not value:
raise ValueError(f"ticker must be a non-empty string, got {value!r}")
if len(value) > max_len:
raise ValueError(f"ticker exceeds {max_len} chars: {value!r}")
if not _TICKER_PATH_RE.fullmatch(value):
raise ValueError(
f"ticker contains characters not allowed in a filesystem path: {value!r}"
)
# The regex above allows '.', so values like '.', '..', '...' would pass,
# and as a path component they traverse the parent directory. Reject any
# value that's only dots.
if set(value) == {"."}:
raise ValueError(f"ticker cannot consist solely of dots: {value!r}")
return value
def get_current_date():
return date.today().strftime("%Y-%m-%d")
def get_scrubbed(url: str, *, params: dict, timeout: float, secret: str, passthrough=()):
"""``requests.get`` plus ``raise_for_status``, with ``secret`` kept out of errors.
Vendors that authenticate with a query parameter put the key in the URL, and
requests quotes the full URL in HTTP, connection and timeout errors, so any
log or traceback that records one would carry the key (#1324). A requests
error is re-raised as the same class with the key replaced and nothing
attached: no request or response (both hold the URL) and no exception chain,
which is why this raises after the ``except`` block rather than inside it.
Statuses in ``passthrough`` are returned for the caller to handle.
"""
try:
response = requests.get(url, params=params, timeout=timeout)
if response.status_code not in passthrough:
response.raise_for_status()
return response
except requests.RequestException as exc:
error = type(exc)(str(exc).replace(secret, "***")) if secret else exc
raise error
def vendor_reachable(url: str, timeout: float = 5.0) -> bool:
"""Whether the vendor answers at all, for telling silence from an outage.
A client that returns an empty result instead of raising leaves those two
cases indistinguishable. Called only when a result is empty.
"""
try:
requests.head(url, timeout=timeout, allow_redirects=True)
return True
except requests.RequestException:
return False
+5 -5
View File
@@ -6,9 +6,10 @@ import pandas as pd
import yfinance as yf
from dateutil.relativedelta import relativedelta
from .date_window import withhold_live_profile
from .errors import VendorError, VendorRateLimitError
from .stockstats_utils import (
from tradingagents.dataflows.date_window import withhold_live_profile
from tradingagents.dataflows.errors import NoMarketDataError, VendorError, VendorRateLimitError
from tradingagents.dataflows.net import vendor_reachable
from tradingagents.dataflows.stockstats_utils import (
StockstatsUtils,
_assert_ohlcv_not_stale,
filter_financials_by_date,
@@ -16,8 +17,7 @@ from .stockstats_utils import (
raise_for_empty,
yf_retry,
)
from .symbol_utils import NoMarketDataError, normalize_symbol
from .utils import vendor_reachable
from tradingagents.dataflows.symbols import normalize_symbol
_YAHOO_HOST = "https://query2.finance.yahoo.com"
+5 -5
View File
@@ -6,11 +6,11 @@ from datetime import datetime, timezone
import yfinance as yf
from dateutil.relativedelta import relativedelta
from .config import get_config
from .date_window import coverage_gap, in_window
from .errors import NoMarketDataError
from .stockstats_utils import yf_retry
from .symbol_utils import normalize_symbol
from tradingagents.dataflows.config import get_config
from tradingagents.dataflows.date_window import coverage_gap, in_window
from tradingagents.dataflows.errors import NoMarketDataError
from tradingagents.dataflows.stockstats_utils import yf_retry
from tradingagents.dataflows.symbols import normalize_symbol
def _extract_article_data(article: dict) -> dict:
+1 -1
View File
@@ -13,7 +13,7 @@ from pathlib import Path
from langgraph.checkpoint.sqlite import SqliteSaver
from tradingagents.dataflows.utils import safe_ticker_component
from tradingagents.dataflows.symbols import safe_ticker_component
def _db_path(data_dir: str | Path, ticker: str) -> Path:
+3 -2
View File
@@ -15,7 +15,8 @@ from tradingagents.agents.utils.agent_utils import (
from tradingagents.agents.utils.memory import TradingMemoryLog
from tradingagents.agents.utils.rating import parse_rating
from tradingagents.dataflows.config import run_config, set_config
from tradingagents.dataflows.utils import get_current_date, safe_ticker_component
from tradingagents.dataflows.date_window import get_current_date
from tradingagents.dataflows.symbols import safe_ticker_component
from tradingagents.dataflows.y_finance import get_closes
from tradingagents.default_config import DEFAULT_CONFIG
from tradingagents.llm_clients import create_llm_client
@@ -207,7 +208,7 @@ class TradingAgentsGraph:
entry, which is the right default because the alpha calculation works
in USD.
"""
from tradingagents.dataflows.symbol_utils import normalize_symbol
from tradingagents.dataflows.symbols import normalize_symbol
explicit = self.config.get("benchmark_ticker")
if explicit: