MaximeSzymanski/TradingAssistant
0
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