"""Unit tests for the AI agent pipeline."""
from app.agents.indicators import atr, bollinger, ema, macd, rsi, sma
from app.agents.market_scanner import MarketScannerAgent
from app.agents.options import OptionsAnalysisAgent
from app.agents.probability import ProbabilityAgent
from app.agents.risk_manager import RiskManagerAgent
from app.agents.sentiment import SentimentAgent
from app.agents.technical import TechnicalAnalysisAgent
from app.agents.volatility import VolatilityAgent
from app.services.market_data import MockMarketDataProvider

provider = MockMarketDataProvider()


def _context(symbol: str = "TSLA") -> dict:
    return {
        "symbol": symbol,
        "history": provider.get_history(symbol),
        "quote": provider.get_quote(symbol),
    }


def test_indicators_basic():
    closes = [float(i) for i in range(1, 61)]
    assert sma(closes, 5) == 58.0
    assert ema(closes, 9) is not None
    assert rsi(closes) == 100.0  # monotonic gains
    assert macd(closes) is not None
    bb = bollinger(closes)
    assert bb and bb[0] > bb[2]


def test_atr_positive():
    bars = provider.get_history("NVDA")
    a = atr(bars)
    assert a is not None and a > 0


def test_scanner_returns_candidates():
    result = MarketScannerAgent(provider, max_candidates=10).run({})
    assert len(result.data["candidates"]) == 10
    scores = [r["activity_score"] for r in result.data["rows"]]
    assert scores == sorted(scores, reverse=True)


def test_volatility_score_range():
    result = VolatilityAgent().run(_context())
    assert 0 <= result.score <= 100
    assert result.data["atr"] > 0


def test_technical_direction_and_indicators():
    result = TechnicalAnalysisAgent().run(_context())
    assert result.direction in ("CALL", "PUT")
    assert 0 <= result.score <= 100
    assert result.data["indicators"]["rsi"] is not None


def test_options_agent_selects_contract():
    ctx = _context()
    ctx["direction"] = "CALL"
    result = OptionsAnalysisAgent(provider).run(ctx)
    assert result.score > 0
    contract = result.data["contract"]
    assert contract["type"] == "call"
    assert contract["ask"] >= contract["bid"]


def test_probability_combines_agents():
    ctx = _context()
    vol = VolatilityAgent().run(ctx)
    tech = TechnicalAnalysisAgent().run(ctx)
    ctx["direction"] = tech.direction
    opts = OptionsAnalysisAgent(provider).run(ctx)
    sent = SentimentAgent().run(ctx)
    result = ProbabilityAgent().run({**ctx, "agent_results": {
        "volatility": vol, "technical": tech, "options": opts, "sentiment": sent}})
    assert 0 <= result.data["confidence"] <= 98
    assert 5 <= result.data["probability"] <= 90


def test_risk_manager_math_and_gating():
    ctx = _context()
    ctx.update({"direction": "CALL", "atr": 5.0, "confidence": 80.0,
                "probability": 60.0, "open_risk_usd": 0.0})
    result = RiskManagerAgent().run(ctx)
    d = result.data
    entry, stop = d["entry"], d["stop_loss"]
    assert stop < entry                      # long stop below entry
    assert d["target1"] > entry
    assert d["risk_reward"] >= 2.0
    assert d["position_size"] > 0
    assert d["approved"] is True


def test_risk_manager_rejects_low_confidence():
    ctx = _context()
    ctx.update({"direction": "PUT", "atr": 5.0, "confidence": 10.0,
                "probability": 20.0, "open_risk_usd": 0.0})
    result = RiskManagerAgent().run(ctx)
    assert result.data["approved"] is False
    assert result.data["rejection_reasons"]
