Spaces:
Paused
Paused
| """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] | |