"""
test_core_logic.py

Pure-logic unit tests (no Alpaca API calls). Run with:
    python -m pytest tests/test_core_logic.py -v
or:
    python tests/test_core_logic.py
"""

import sys
import os
import unittest
from datetime import datetime, timedelta, timezone

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

import indicators
import risk_manager
import symbol_filters
from scorer import score_premarket_candidate


def make_bars(prices, volumes, start=None):
    start = start or datetime(2026, 8, 17, 9, 0, tzinfo=timezone.utc)
    bars = []
    for i, (p, v) in enumerate(zip(prices, volumes)):
        bars.append({
            "t": start + timedelta(minutes=i),
            "o": p - 0.01, "h": p + 0.02, "l": p - 0.02, "c": p, "v": v,
        })
    return bars


class TestIndicators(unittest.TestCase):
    def test_vwap_basic(self):
        bars = make_bars([10, 10, 10], [100, 100, 100])
        self.assertAlmostEqual(indicators.vwap(bars), 10.0, places=1)

    def test_vwap_slope_rising(self):
        bars = make_bars([10, 10.1, 10.2, 10.3, 10.4, 10.5], [100] * 6)
        self.assertGreater(indicators.vwap_slope(bars, lookback=3), 0)

    def test_volume_acceleration_developing(self):
        # Matches the project's "developing momentum" example
        volumes = [450_000, 620_000, 890_000, 1_020_000]
        self.assertGreater(indicators.volume_acceleration(volumes), 1.15)

    def test_volume_acceleration_fading(self):
        # Matches the project's "fading momentum" example
        volumes = [1_200_000, 1_250_000, 1_270_000, 1_270_000]
        self.assertLess(indicators.volume_acceleration(volumes), 1.15)

    def test_higher_highs_higher_lows_true(self):
        bars = make_bars([8.0, 8.1, 8.2, 8.35, 8.5], [100] * 5)
        self.assertTrue(indicators.is_higher_highs_higher_lows(bars, lookback=4))

    def test_higher_highs_higher_lows_false_choppy(self):
        bars = make_bars([8.5, 8.2, 8.5, 8.2, 8.2], [100] * 5)
        self.assertFalse(indicators.is_higher_highs_higher_lows(bars, lookback=4))

    def test_breakout_confirmed_holds(self):
        bars = make_bars([9.9, 10.05, 10.08, 10.10], [100] * 4)
        self.assertTrue(indicators.breakout_confirmed(bars, resistance=10.0, buffer_pct=0.1, hold_bars=2))

    def test_breakout_not_confirmed_spike_fails(self):
        bars = make_bars([9.9, 10.05, 9.95, 9.90], [100] * 4)
        self.assertFalse(indicators.breakout_confirmed(bars, resistance=10.0, buffer_pct=0.1, hold_bars=2))

    def test_spread_pct(self):
        self.assertAlmostEqual(indicators.spread_pct(9.98, 10.02), 0.4, places=1)

    def test_extended_downtrend_flags_bounce_still_off_the_high(self):
        # Mirrors ACHR's 2026-09-03 shape: a real decline from a session
        # high, then a partial bounce that hasn't reclaimed it.
        decline = [10.0 - i * (1.0 / 19) for i in range(20)]  # 10.0 -> 9.0
        bounce = [9.0 + i * (0.3 / 9) for i in range(10)]      # 9.0 -> 9.3
        bars = make_bars(decline + bounce, [100] * 30)
        self.assertTrue(indicators.is_extended_downtrend(
            bars, min_bars=20, min_decline_pct=1.5, recovery_threshold_pct=1.5))

    def test_extended_downtrend_false_once_high_is_reclaimed(self):
        # Same decline, but the bounce fully reclaims (and exceeds) the
        # session high -- this is a reversal, not a bounce, and must not
        # be flagged regardless of the whole-window regression's sign.
        decline = [10.0 - i * (1.0 / 19) for i in range(20)]   # 10.0 -> 9.0
        recovery = [9.0 + i * (1.5 / 9) for i in range(10)]    # 9.0 -> 10.5
        bars = make_bars(decline + recovery, [100] * 30)
        self.assertFalse(indicators.is_extended_downtrend(
            bars, min_bars=20, min_decline_pct=1.5, recovery_threshold_pct=1.5))

    def test_extended_downtrend_false_with_insufficient_history(self):
        decline = [10.0 - i * 0.05 for i in range(10)]
        bars = make_bars(decline, [100] * 10)
        self.assertFalse(indicators.is_extended_downtrend(
            bars, min_bars=20, min_decline_pct=1.5, recovery_threshold_pct=1.5))

    def test_regression_fit_pct_change_matches_direction(self):
        rising = make_bars([10, 10.1, 10.2, 10.3, 10.4], [100] * 5)
        falling = make_bars([10.4, 10.3, 10.2, 10.1, 10.0], [100] * 5)
        self.assertGreater(
            indicators.regression_fit_pct_change([b["c"] for b in rising]), 0)
        self.assertLess(
            indicators.regression_fit_pct_change([b["c"] for b in falling]), 0)

    def test_atr_reasonable(self):
        bars = make_bars([10, 10.2, 10.1, 10.3, 10.25], [100] * 5)
        a = indicators.atr(bars, period=4)
        self.assertGreater(a, 0)
        self.assertLess(a, 1.0)


class TestRiskManager(unittest.TestCase):
    def test_trailing_stop_only_moves_up(self):
        # Reproduces the exact worked example from the spec.
        distance = 0.10
        stop = 10.00 - distance
        for price in [10.00, 10.05, 10.10, 10.20]:
            new_stop = max(stop, price - distance)
            self.assertGreaterEqual(new_stop, stop)
            stop = new_stop
        self.assertAlmostEqual(stop, 10.10, places=2)

    def test_stop_never_moves_backward_on_pullback(self):
        highest = 10.50
        current_stop = 10.40
        # price pulls back to 10.39 but highest_price (10.50) hasn't changed
        new_stop = risk_manager.update_trailing_stop(highest, current_stop, bars=None)
        self.assertEqual(new_stop, 10.40)

    def test_initial_stop_fixed_cents(self):
        stop = risk_manager.compute_initial_stop(10.00)
        self.assertAlmostEqual(stop, 9.90, places=2)

    def test_position_sizing_respects_risk_pct(self):
        shares = risk_manager.compute_position_size(
            account_equity=100_000, entry_price=10.00, initial_stop=9.90
        )
        # 1% risk of 100k = $1000 / $0.10 risk-per-share = 10,000 shares by risk,
        # but capped by max_position_notional_pct_of_equity (20% = $20,000 / $10 = 2000)
        self.assertEqual(shares, 2000)

    def test_position_sizing_zero_on_bad_inputs(self):
        shares = risk_manager.compute_position_size(
            account_equity=0, entry_price=10.00, initial_stop=9.90
        )
        self.assertEqual(shares, 0)


class TestScorer(unittest.TestCase):
    def test_developing_candidate_scores_higher_than_fading(self):
        developing_bars = make_bars(
            [8.42, 8.51, 8.58, 8.61], [450_000, 620_000, 890_000, 1_020_000]
        )
        fading_bars = make_bars(
            [8.80, 8.76, 8.65, 8.60], [1_200_000, 1_250_000, 1_270_000, 1_270_000]
        )
        dev = score_premarket_candidate("DEV", developing_bars, pm_high=8.65, pm_low=8.40,
                                         avg_vol_baseline=2_000_000, bid=8.60, ask=8.62)
        fade = score_premarket_candidate("FADE", fading_bars, pm_high=8.85, pm_low=8.55,
                                          avg_vol_baseline=2_000_000, bid=8.58, ask=8.62)
        self.assertGreater(dev["total_score"], fade["total_score"])

    def test_score_breakdown_present_and_explains_total(self):
        bars = make_bars([10.0, 10.1, 10.2], [50_000, 60_000, 70_000])
        result = score_premarket_candidate("XYZ", bars, pm_high=10.25, pm_low=9.95,
                                            avg_vol_baseline=1_000_000)
        self.assertIn("breakdown", result)
        self.assertIn("price_quality", result["breakdown"])
        self.assertIn("volume_quality", result["breakdown"])

    def test_out_of_price_band_scores_zero_price_quality(self):
        bars = make_bars([20.0, 20.1, 20.2], [50_000, 60_000, 70_000])
        result = score_premarket_candidate("HIGH", bars, pm_high=20.25, pm_low=19.9,
                                            avg_vol_baseline=1_000_000)
        self.assertEqual(result["breakdown"]["price_quality"], 0.0)


class TestSymbolFilters(unittest.TestCase):
    def test_leveraged_etf_rejected(self):
        ok, reason = symbol_filters.passes_symbol_filters(
            "SOXL", "Direxion Daily Semiconductor Bull 3X Shares"
        )
        self.assertFalse(ok)
        self.assertIn("ETF", reason)  # caught by fund keyword before leverage check

    def test_plain_etf_rejected(self):
        ok, reason = symbol_filters.passes_symbol_filters("SPY", "SPDR S&P 500 ETF Trust")
        self.assertFalse(ok)
        self.assertEqual(reason, "ETF/fund/trust")

    def test_inverse_product_rejected(self):
        ok, reason = symbol_filters.passes_symbol_filters(
            "SQQQ", "ProShares UltraPro Short QQQ"
        )
        self.assertFalse(ok)

    def test_clean_common_stock_accepted(self):
        ok, reason = symbol_filters.passes_symbol_filters("AAPL", "Apple Inc.")
        self.assertTrue(ok)
        self.assertEqual(reason, "")

    def test_long_name_rejected_over_syllable_cap(self):
        ok, reason = symbol_filters.passes_symbol_filters(
            "XYZ", "Telecommunication Systems International Holdings Corporation"
        )
        self.assertFalse(ok)
        self.assertIn("syllables", reason)

    def test_legal_suffixes_dont_inflate_syllable_count(self):
        # "Ford Motor Company" should count ~3 syllables (Ford, Mo-tor),
        # with "Company" stripped as a legal suffix rather than counted.
        syl = symbol_filters.count_syllables("Ford Motor Company")
        self.assertLessEqual(syl, 4)

    def test_filter_symbol_list_end_to_end(self):
        class FakeAsset:
            def __init__(self, symbol, name):
                self.symbol = symbol
                self.name = name

        assets = [
            FakeAsset("AAPL", "Apple Inc."),
            FakeAsset("SPY", "SPDR S&P 500 ETF Trust"),
            FakeAsset("SOXL", "Direxion Daily Semiconductor Bull 3X Shares"),
            FakeAsset("F", "Ford Motor Company"),
        ]
        survivors = symbol_filters.filter_symbol_list(assets, log_rejections=False)
        self.assertIn("AAPL", survivors)
        self.assertIn("F", survivors)
        self.assertNotIn("SPY", survivors)
        self.assertNotIn("SOXL", survivors)


if __name__ == "__main__":
    unittest.main(verbosity=2)
