hale-api / tests /test_rl_engine_extra.py
Raunak211006's picture
Deploy HALE API v3.1
cc38116 verified
Raw
History Blame Contribute Delete
5.42 kB
# tests/test_rl_engine_extra.py
import pytest
from datetime import datetime, date, timedelta, timezone
from unittest.mock import MagicMock, patch
from app.services.rl_engine import (
_compute_circadian_bonus,
_apply_memory_decay,
CIRCADIAN_WINDOWS,
MEMORY_DECAY_GAMMA
)
def make_user(user_id="testuser", rl_bandit_params=None):
user = MagicMock()
user.user_id = user_id
user.rl_bandit_params = rl_bandit_params
user.rl_history_count = 10
user.rl_last_decay_date = None
user.days_since_last_login = 0
user.current_streak = 0
return user
def test_compute_circadian_bonus_no_db():
user = make_user()
bonus = _compute_circadian_bonus(user, None, "Deep Work")
assert bonus == 0.0
def test_compute_circadian_bonus_wrong_action():
user = make_user()
db = MagicMock()
bonus = _compute_circadian_bonus(user, db, "Normal")
assert bonus == 0.0
@patch("app.services.rl_engine.datetime")
def test_compute_circadian_bonus_with_db(mock_dt):
# Mock datetime to a specific hour, say 14 (afternoon)
mock_now = datetime(2023, 1, 1, 14, 0, 0)
mock_dt.now.return_value = mock_now
class FakeProgress:
def __init__(self, hour, satisfaction):
self.completed_at = datetime(2023, 1, 1, hour, 0, 0, tzinfo=timezone.utc)
self.satisfaction = satisfaction
rows = [
FakeProgress(14, 1.0),
FakeProgress(15, 0.9),
FakeProgress(16, 1.0),
FakeProgress(8, 0.2),
FakeProgress(9, 0.3),
FakeProgress(22, 0.1),
FakeProgress(23, 0.2),
]
db = MagicMock()
db.query.return_value.filter.return_value.all.return_value = rows
user = make_user()
bonus = _compute_circadian_bonus(user, db, "Deep Work")
assert bonus > 0
assert bonus <= 3.0
def test_apply_memory_decay_no_decay_needed():
user = make_user()
user.last_login_at = datetime.now(timezone.utc) - timedelta(days=1)
user.rl_last_decay_date = (datetime.now(timezone.utc) - timedelta(days=1)).isoformat()
params = {
"Deep Work": {"alpha": 10.0, "beta": 10.0},
"Learning": {"alpha": 5.0, "beta": 5.0},
"Maintenance": {"alpha": 3.0, "beta": 3.0},
"Recovery": {"alpha": 1.0, "beta": 1.0},
}
db = MagicMock()
new_params = _apply_memory_decay(user, params, db)
assert new_params["Deep Work"]["alpha"] == 10.0
def test_apply_memory_decay_performs_decay():
user = make_user()
user.last_login_at = datetime.now(timezone.utc) - timedelta(days=10)
params = {
"Deep Work": {"alpha": 10.0, "beta": 10.0},
"Learning": {"alpha": 5.0, "beta": 5.0},
"Maintenance": {"alpha": 3.0, "beta": 3.0},
"Recovery": {"alpha": 1.0, "beta": 1.0},
}
db = MagicMock()
new_params = _apply_memory_decay(user, params, db)
if new_params["Deep Work"]["alpha"] == 10.0:
print("Decay failed to apply. Params returned unchanged.")
expected_alpha = 10.0 * MEMORY_DECAY_GAMMA + 2.0 * (1 - MEMORY_DECAY_GAMMA)
expected_beta = 10.0 * MEMORY_DECAY_GAMMA + 2.0 * (1 - MEMORY_DECAY_GAMMA)
assert abs(new_params["Deep Work"]["alpha"] - expected_alpha) < 0.001
assert abs(new_params["Deep Work"]["beta"] - expected_beta) < 0.001
def test_get_next_day_signal_does_not_crash():
from app.services.rl_engine import get_next_day_signal
user = make_user()
user.competence = 0.5
user.mood = 0.5
user.fatigue = 0.5
user.sleep_quality = 0.5
user.shock = 0.0
user.sick = 0.0
user.stress = 0.0
# Should not raise TypeError
signal = get_next_day_signal(user)
assert signal["preview"] is True
assert "action_label" in signal
@patch("app.services.rl_engine._compute_historical_features")
def test_reward_hacking_discount(mock_hist):
from app.services.rl_engine import update_reward
user = make_user("hacker", "{}")
db = MagicMock()
# Simulate high fatigue to trigger reward hacking squash
mock_hist.return_value = {"ema_fatigue": 0.9}
update_reward(db, user, "Normal", 1.0)
# 1.0 should be squashed toward 0.5 -> 0.75
# Default alpha is 3.0 for Normal. alpha += 0.75 -> 3.75
import json
params = json.loads(user.rl_bandit_params)
assert params["Normal"]["alpha"] < 4.0
assert params["Normal"]["alpha"] >= 3.70
@patch("app.services.rl_engine._compute_historical_features")
def test_circuit_breaker_debt(mock_hist):
from app.services.rl_engine import get_behavior_signal
user = make_user("burnout", '{"_recovery_debt": 16.0}')
db = MagicMock()
mock_hist.return_value = {
"consecutive_hard_days": 5.0,
"recent_completion_rate": 0.5,
"avg_satisfaction_7d": 0.5,
"fatigue_trend": 0.0,
"overwork_signal": 0.0,
"deep_work_streak": 0.0,
"ema_fatigue": 0.0,
"ema_mood": 0.5,
"last_action": "Deep Work"
}
signal = get_behavior_signal(user, state_override=[0.5, 0.5, 0.9, 0.5, 0.0, 0.0, 0.0], db=db)
# Debt is > 15, so Deep Work should be circuit broken into Recovery.
# Also fatigue is 0.9, triggering ceiling.
assert signal["action_label"] in ("Recovery", "Light Review")