CoolFace
Apppublic

MaximeSzymanski/TradingAssistant

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
test_tradercomp.py447 linesDownload Raw Back to root
1import pytest2from unittest.mock import MagicMock, patch, ANY3from datetime import datetime, timedelta4import pandas as pd5import numpy as np6from langchain_core.messages import HumanMessage, AIMessage7 8# --- IMPORT YOUR MODULES ---9# Adjust paths if your folder structure is different10from my_agent.utils.nodes import (11    node_extract_entities,12    node_validate_ticker,13    node_validate_dates,14    node_fetch_data,15    node_technical_analysis,16    node_sentiment_analysis,17    node_forecast,18    node_rag_search,19    node_company_profile,20    node_clarify_preference,21    node_generate_viz,22    node_ask_user,23    node_risk_disclaimer24)25# Import tools explicitly to test them26from my_agent.utils.tools import search_financial_news, get_company_fundamentals27 28# =============================================================================29# 0. FIXTURES30# =============================================================================31 32@pytest.fixture33def base_state():34    """Returns a fresh, empty state for every test."""35    return {36        "messages": [HumanMessage(content="Test message")],37        "ticker": None,38        "start_date": None,39        "end_date": None,40        "stock_data": None,41        "forecast_data": None,42        "output_preference": None,43        "error_message": None,44        "news_summary": None,45        "rag_fallback": False,46        "is_rag_active": False47    }48 49@pytest.fixture50def mock_stock_data():51    """Creates a dummy pandas DataFrame resembling yfinance data."""52    dates = pd.date_range(start="2023-01-01", periods=50)53    data = {54        "Open": np.random.rand(50) * 100,55        "High": np.random.rand(50) * 100,56        "Low": np.random.rand(50) * 100,57        "Close": np.linspace(100, 150, 50), # Linear upward trend for forecast testing58        "Volume": np.random.randint(1000, 5000, 50)59    }60    df = pd.DataFrame(data, index=dates)61    return df.to_dict(orient="index")62 63# =============================================================================64# 1. TEST ENTITY EXTRACTION65# =============================================================================66 67@patch("my_agent.utils.nodes.ChatOllama")68def test_extract_entities_success(mock_chat, base_state):69    """Test successful JSON extraction from LLM."""70    base_state["messages"] = [HumanMessage(content="Analyze Apple from 2023-01-01 to 2023-02-01")]71    72    # Mock chain execution: chain.invoke(...) -> returns dict73    # Since the node constructs the chain dynamically, we mock the final parser output74    # But since it's hard to mock the pipe '|', we mock the invoke call on the result of the prompt.75    # ALTERNATIVE: Mock ChatOllama to return a JSON string, and the real parser handles it?76    # Simpler: The node code calls `chain.invoke`. Let's assume we can patch the chain construction.77    # Given the complexity of patching inside the function, let's mock the `chain.invoke` return value 78    # by patching the class method if possible, or just the LLM response if the chain is robust.79    80    # Let's try mocking the object returned by the chain construction:81    with patch("my_agent.utils.nodes.ChatPromptTemplate") as mock_prompt:82        mock_chain = MagicMock()83        mock_chain.invoke.return_value = {84            "company_name": "Apple",85            "start_date": "2023-01-01",86            "end_date": "2023-02-01",87            "visualization_type": "plot"88        }89        # We make the pipe operator return our mock chain90        mock_prompt.from_messages.return_value.partial.return_value.__or__.return_value.__or__.return_value = mock_chain91        92        result = node_extract_entities(base_state)93        94        assert result["ticker"] == "Apple"95        assert result["start_date"] == "2023-01-01"96        assert result["output_preference"] == "plot"97 98@patch("my_agent.utils.nodes.ChatOllama")99def test_extract_entities_failure_handling(mock_chat, base_state):100    """Test that it doesn't crash if LLM fails."""101    with patch("my_agent.utils.nodes.ChatPromptTemplate") as mock_prompt:102        mock_chain = MagicMock()103        mock_chain.invoke.side_effect = Exception("LLM Timeout")104        mock_prompt.from_messages.return_value.partial.return_value.__or__.return_value.__or__.return_value = mock_chain105        106        result = node_extract_entities(base_state)107        assert result == {} # Graceful exit108 109# =============================================================================110# 2. TEST TICKER VALIDATION (The Complex Node)111# =============================================================================112 113@patch("my_agent.utils.nodes.SP500_DF", pd.DataFrame({'Symbol': ['AAPL'], 'Security': ['Apple Inc.']}))114def test_validate_ticker_sp500_fast_path(base_state):115    """Test fast lookup in S&P 500 CSV."""116    base_state["ticker"] = "Apple"117    result = node_validate_ticker(base_state)118    assert result["ticker"] == "AAPL"119    assert result["error_message"] is None120 121@patch("my_agent.utils.nodes.SP500_DF", pd.DataFrame(columns=['Symbol', 'Security']))122@patch("my_agent.utils.nodes.ChatOllama")123@patch("my_agent.utils.nodes.DuckDuckGoSearchRun")124@patch("my_agent.utils.nodes.yf.Ticker")125def test_validate_ticker_collision_boralex(mock_ticker, mock_search, mock_chat, base_state):126    """127    CRITICAL TEST: The Boralex vs Banco Latinoamericano collision.128    1. Web Search finds "BLX".129    2. Loop 1: Checks BLX. Info says "Banco". Mismatch.130    3. Loop 2: Checks BLX.TO. Info says "Boralex". Match.131    """132    base_state["ticker"] = "Boralex"133    134    # 1. Web Search finds the generic ticker135    mock_search.return_value.invoke.return_value = "Boralex Inc (BLX) Stock..."136    mock_chat.return_value.invoke.return_value = MagicMock(content="BLX")137 138    # 2. Mock YFinance for different calls139    # Mock Object for BLX (Banco)140    mock_banco = MagicMock()141    mock_banco.history.return_value = pd.DataFrame({'Close': [10]})142    mock_banco.info = {"longName": "Banco Latinoamericano de Comercio Exterior, S. A."}143 144    # Mock Object for BLX.TO (Boralex)145    mock_boralex = MagicMock()146    mock_boralex.history.return_value = pd.DataFrame({'Close': [20]})147    mock_boralex.info = {"longName": "Boralex Inc."}148 149    # Side Effect: Return Banco first, then Boralex when suffix added150    def side_effect(ticker):151        if ticker == "BLX": return mock_banco152        if ticker == "BLX.TO": return mock_boralex153        return MagicMock(history=lambda period: pd.DataFrame()) # Return empty for others154 155    mock_ticker.side_effect = side_effect156 157    # Run158    result = node_validate_ticker(base_state)159 160    # Assert161    assert result["ticker"] == "BLX.TO"162    assert result["error_message"] is None163 164@patch("my_agent.utils.nodes.SP500_DF", pd.DataFrame(columns=['Symbol', 'Security']))165@patch("my_agent.utils.nodes.ChatOllama")166@patch("my_agent.utils.nodes.DuckDuckGoSearchRun")167@patch("my_agent.utils.nodes.yf.Ticker")168def test_validate_ticker_retry_root(mock_ticker, mock_search, mock_chat, base_state):169    """170    Test logic: If STLA.PA fails, try STLA (Root).171    """172    base_state["ticker"] = "Stellantis"173    174    # Web search says STLA.PA175    mock_search.return_value.invoke.return_value = "..."176    mock_chat.return_value.invoke.return_value = MagicMock(content="STLA.PA")177 178    # STLA.PA fails (empty history), STLA succeeds179    mock_fail = MagicMock()180    mock_fail.history.return_value = pd.DataFrame()181    182    mock_success = MagicMock()183    mock_success.history.return_value = pd.DataFrame({'Close': [10]})184    mock_success.info = {"longName": "Stellantis N.V."}185 186    mock_ticker.side_effect = lambda t: mock_success if t == "STLA" else mock_fail187 188    result = node_validate_ticker(base_state)189    # The robust node should strip .PA and try STLA190    assert result["ticker"] == "STLA"191 192# =============================================================================193# 3. TEST DATE VALIDATION194# =============================================================================195 196def test_validate_dates_missing_inputs(base_state):197    """If dates are missing and no preference is set, pass."""198    base_state["ticker"] = "AAPL"199    result = node_validate_dates(base_state)200    assert result["error_message"] is None201 202def test_validate_dates_missing_inputs_with_plot_pref(base_state):203    """If preference is PLOT, missing dates is an error."""204    base_state["output_preference"] = "plot"205    result = node_validate_dates(base_state)206    assert "provide start and end dates" in result["error_message"]207 208def test_validate_dates_future_error(base_state):209    future = (datetime.now() + timedelta(days=365)).strftime("%Y-%m-%d")210    base_state["start_date"] = "2023-01-01"211    base_state["end_date"] = future212    result = node_validate_dates(base_state)213    assert "cannot be in the future" in result["error_message"]214 215def test_validate_dates_start_after_end(base_state):216    base_state["start_date"] = "2023-02-01"217    base_state["end_date"] = "2023-01-01"218    result = node_validate_dates(base_state)219    assert "Start date cannot be after end date" in result["error_message"]220 221# =============================================================================222# 4. TEST DATA FETCHING & ROUTING (Keyword Guards & RSS)223# =============================================================================224 225def test_fetch_data_keyword_guard(base_state):226    """Test that 'price' keyword skips LLM router."""227    base_state["ticker"] = "AAPL"228    base_state["start_date"] = "2023-01-01"229    base_state["end_date"] = "2023-01-02"230    base_state["messages"] = [HumanMessage(content="What is the price history?")] # Contains 'price'231 232    with patch("my_agent.utils.nodes.yf.Ticker") as mock_ticker:233        with patch("my_agent.utils.nodes.ChatOllama") as mock_chat:234            mock_stock = MagicMock()235            mock_stock.history.return_value = pd.DataFrame({'Close': [100]}, index=pd.to_datetime(["2023-01-01"]))236            mock_ticker.return_value = mock_stock237 238            result = node_fetch_data(base_state)239            240            # ChatOllama should NOT be called because of keyword guard241            mock_chat.assert_not_called()242            assert result["stock_data"] is not None243 244@patch("my_agent.utils.nodes.search_financial_news")245@patch("my_agent.utils.nodes.ChatOllama")246def test_fetch_data_routes_to_news(mock_chat, mock_tool, base_state):247    """Test routing to news tool."""248    base_state["messages"] = [HumanMessage(content="Why is it down?")]249    base_state["ticker"] = "AAPL"250 251    # Mock Router Response252    mock_msg = MagicMock()253    mock_msg.tool_calls = [{"name": "search_financial_news", "args": {"ticker": "AAPL"}}]254    mock_chat.return_value.bind_tools.return_value.invoke.return_value = mock_msg255 256    # Mock Tool execution257    mock_tool.invoke.return_value = "Fake RSS Data"258 259    result = node_fetch_data(base_state)260    261    assert result["news_summary"] == "Fake RSS Data"262    mock_tool.invoke.assert_called_with("AAPL") # Ensure ticker was passed, not search string263 264@patch("my_agent.utils.nodes.get_company_fundamentals")265@patch("my_agent.utils.nodes.ChatOllama")266def test_fetch_data_routes_to_fundamentals(mock_chat, mock_tool, base_state):267    """Test routing to fundamentals."""268    base_state["messages"] = [HumanMessage(content="What does this company do?")]269    base_state["ticker"] = "AAPL"270 271    mock_msg = MagicMock()272    mock_msg.tool_calls = [{"name": "get_company_fundamentals", "args": {"ticker": "AAPL"}}]273    mock_chat.return_value.bind_tools.return_value.invoke.return_value = mock_msg274 275    mock_tool.invoke.return_value = {"name": "Apple", "summary": "Tech stuff", "pe_ratio": 30}276 277    result = node_fetch_data(base_state)278    assert isinstance(result["messages"][0], AIMessage)279    assert "Tech stuff" in result["messages"][0].content280 281@patch("my_agent.utils.nodes.yf.Ticker")282@patch("my_agent.utils.nodes.ChatOllama")283def test_fetch_data_missing_params(mock_chat, mock_ticker, base_state):284    """Test error when routing to data but dates are missing."""285    base_state["messages"] = [HumanMessage(content="Show me price")] # triggers keyword guard286    base_state["ticker"] = "AAPL"287    base_state["start_date"] = None # Missing288 289    result = node_fetch_data(base_state)290    assert "Missing ticker or date range" in result["error_message"]291 292# =============================================================================293# 5. TEST RSS & FUNDAMENTALS TOOLS (Unit Tests)294# =============================================================================295 296@patch("requests.get")297def test_tool_rss_success(mock_get):298    """Test the RSS parser tool."""299    xml = """<rss><channel><item>300             <title>Test News</title><link>http://link</link><pubDate>Mon</pubDate>301             </item></channel></rss>"""302    mock_get.return_value.status_code = 200303    mock_get.return_value.content = xml.encode()304 305    result = search_financial_news.invoke("AAPL")306    assert "Test News" in result307 308@patch("requests.get")309def test_tool_rss_failure(mock_get):310    """Test RSS network failure."""311    mock_get.return_value.status_code = 404312    result = search_financial_news.invoke("AAPL")313    assert "Failed to retrieve" in result314 315@patch("my_agent.utils.tools.yf.Ticker")316def test_tool_fundamentals(mock_ticker):317    """Test fundamentals extraction."""318    mock_info = {"longName": "Test Co", "marketCap": 1000000000, "sector": "Tech"}319    mock_ticker.return_value.info = mock_info320    321    result = get_company_fundamentals.invoke("TEST")322    assert result["name"] == "Test Co"323    assert "Billion" in result["market_cap"]324 325# =============================================================================326# 6. TEST MATH NODES (Forecast & Tech Analysis)327# =============================================================================328 329def test_forecast_linear_regression(base_state, mock_stock_data):330    """Test that forecasting produces valid numbers and R2 score."""331    base_state["stock_data"] = mock_stock_data332    333    result = node_forecast(base_state)334    335    assert "forecast_data" in result336    assert "forecast_meta" in result337    meta = result["forecast_meta"]338    # Since mock data is perfectly linear (linspace), R2 should be close to 1.0339    assert meta["r2_score"] > 0.95 340    assert meta["trend"] == "upward"341    assert len(result["forecast_data"]) == 29342 343def test_technical_analysis_calcs(base_state, mock_stock_data):344    """Test RSI and SMA addition."""345    base_state["stock_data"] = mock_stock_data346    result = node_technical_analysis(base_state)347    first_row = list(result["stock_data"].values())[-1]348    assert "RSI" in first_row349    assert "SMA_20" in first_row350 351@patch("my_agent.utils.nodes.ChatOllama")352def test_sentiment_analysis(mock_chat, base_state):353    """Test sentiment JSON parsing."""354    base_state["news_summary"] = "Good news"355    356    # Mock LLM returning JSON string357    mock_chat.return_value.invoke.return_value = AIMessage(content='{"score": 0.8, "verdict": "Bullish", "explanation": "Good stuff"}')358    359    result = node_sentiment_analysis(base_state)360    assert result["sentiment"]["score"] == 0.8361    assert result["news_summary"] is None # Should be cleared362 363# =============================================================================364# 7. TEST RAG & OTHER NODES365# =============================================================================366 367@patch("my_agent.utils.nodes.query_rag")368@patch("my_agent.utils.nodes.ChatOllama")369def test_rag_success(mock_chat, mock_query, base_state):370    """Test successful retrieval."""371    base_state["messages"] = [HumanMessage(content="What is in the doc?")]372    mock_query.return_value = "Context Chunk 1..."373    mock_chat.return_value.invoke.return_value = AIMessage(content="The doc says X.")374    375    result = node_rag_search(base_state)376    assert result["rag_fallback"] is False377 378@patch("my_agent.utils.nodes.query_rag")379def test_rag_no_context(mock_query, base_state):380    """Test fallback when no context found."""381    base_state["messages"] = [HumanMessage(content="Unknown")]382    mock_query.return_value = None383    result = node_rag_search(base_state)384    assert result["rag_fallback"] is True385 386def test_company_profile(base_state):387    """Test company profile formatting."""388    base_state["ticker"] = "AAPL"389    base_state["stock_data"] = {"info": {"longName": "Apple", "marketCap": 2000000000}}390    result = node_company_profile(base_state)391    assert "Apple" in result["messages"][0].content392    assert "$2.00B" in result["messages"][0].content393 394def test_risk_disclaimer(base_state):395    """Test disclaimer triggers only if signals exist."""396    # 1. Test Negative Case (No data -> No disclaimer)397    base_state["signals"] = None398    base_state["forecast_data"] = None399    assert node_risk_disclaimer(base_state) == {}400 401    # 2. Test Positive Case (Forecast exists -> Add disclaimer)402    base_state["forecast_data"] = {"2024-01-01": 100}403    404    # Run Node405    result = node_risk_disclaimer(base_state)406    407    # Assert408    # FIX: Check the LAST message [-1], not the first one [0]409    last_message = result["messages"][-1]410    411    assert isinstance(last_message, AIMessage)412    assert "educational purposes" in last_message.content413 414def test_clarify_preference(base_state):415    """Test preference standardization."""416    base_state["output_preference"] = "graph"417    assert node_clarify_preference(base_state)["output_preference"] == "plot"418    419    base_state["output_preference"] = "unknown"420    result = node_clarify_preference(base_state)421    assert result["output_preference"] is None422    assert result["error_message"] is not None423 424# =============================================================================425# 8. TEST VISUALIZATION (Plotly & Table)426# =============================================================================427 428@patch("my_agent.utils.nodes.os.makedirs")429@patch("my_agent.utils.nodes.go.Figure.write_html")430def test_generate_viz_plot(mock_write, mock_dirs, base_state, mock_stock_data):431    """Test interactive plot generation."""432    base_state["stock_data"] = mock_stock_data433    base_state["output_preference"] = "plot"434    base_state["ticker"] = "TEST"435    436    result = node_generate_viz(base_state)437    assert "charts/" in result["messages"][0].content438    mock_write.assert_called_once()439 440def test_generate_viz_table(base_state, mock_stock_data):441    """Test markdown table generation."""442    base_state["stock_data"] = mock_stock_data443    base_state["output_preference"] = "table"444    445    result = node_generate_viz(base_state)446    assert "Close" in result["messages"][0].content447    assert "|" in result["messages"][0].content