"""Tests for synthetic company generation.""" import pytest from hr_env.server.data_gen import generate_company from hr_env.server.company import DEPARTMENT_NAMES class TestGenerateCompany: def test_generates_correct_size(self): company = generate_company(seed=42, size=300) # Allow ±10% due to rounding total = company.total_headcount assert 250 <= total <= 350, f"Expected ~300 employees, got {total}" def test_has_all_departments(self): company = generate_company(seed=42, size=200) for dept_name in DEPARTMENT_NAMES: assert dept_name in company.departments assert company.departments[dept_name].headcount > 0 def test_deterministic_with_seed(self): c1 = generate_company(seed=42, size=200) c2 = generate_company(seed=42, size=200) assert c1.total_headcount == c2.total_headcount for dept in DEPARTMENT_NAMES: assert c1.departments[dept].headcount == c2.departments[dept].headcount def test_different_seeds_differ(self): c1 = generate_company(seed=42, size=200) c2 = generate_company(seed=99, size=200) # Should differ in at least some employee attributes e1 = c1.all_active_employees[0] e2 = c2.all_active_employees[0] assert e1.name != e2.name or e1.salary != e2.salary def test_employee_attributes_in_range(self): company = generate_company(seed=42, size=300) for emp in company.all_active_employees: assert 1 <= emp.level <= 5 assert emp.salary >= 35000 assert 1.0 <= emp.performance_score <= 5.0 assert 0 <= emp.engagement <= 100 assert 0 <= emp.flight_risk <= 1.0 assert 0 <= emp.promotability <= 1.0 assert 0 <= emp.transferability <= 1.0 assert 0 <= emp.retainability <= 1.0 assert emp.is_active def test_salary_distribution(self): company = generate_company(seed=42, size=500) salaries = [e.salary for e in company.all_active_employees] avg = sum(salaries) / len(salaries) # Average salary should be reasonable (60K-120K) assert 60000 <= avg <= 120000, f"Avg salary {avg} out of range" def test_level_distribution_pyramid(self): company = generate_company(seed=42, size=500) levels = [e.level for e in company.all_active_employees] level_counts = {i: levels.count(i) for i in range(1, 6)} # More level 1s than level 5s (pyramid) assert level_counts[1] > level_counts[5] assert level_counts[1] > level_counts[4]