hcm21 / tests /test_data_gen.py
ParetoOptimal's picture
Initial release: HCM:21 HR Productivity Measurement Environment
8697ae6
Raw
History Blame Contribute Delete
2.64 kB
"""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]