File size: 2,638 Bytes
8697ae6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
"""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]