File size: 7,979 Bytes
a67fb69
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ba8c515
 
a67fb69
ba8c515
 
 
a67fb69
 
 
e23d022
 
 
a67fb69
 
 
 
 
 
 
625e140
e23d022
625e140
 
 
 
 
 
 
 
 
a67fb69
 
 
ea8ea0a
 
 
 
e23d022
 
625e140
a67fb69
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
"""Unit tests for `lexsi_ds.agent.tools.train_tabular_model`.

Covers:
  - happy path (XGBoost default, foundation model, sample_percentage)
  - missing context df β†’ missing_df
  - target column not in df β†’ bad_target
  - all-null target β†’ empty_target
  - no entity column resolvable β†’ no_entity
  - foundation model without compute_type β†’ missing_compute_type
  - no tab_project β†’ no_tab_project
  - cache contract: last_model_id, last_train_tag, last_model_task,
    last_train_df_label are populated correctly.
"""

from __future__ import annotations

import pandas as pd
import pytest

from lexsi_ds.agent.tools.train_tabular_model import (
    TOOL,
    TrainTabularModelArgs,
    _build_project_config,
    _mint_train_tag,
    _resolve_active_model_name,
    _resolve_entity_col,
)


# ----- happy path -----


def test_happy_path_xgboost(ctx_with_fake, seed_sql_result, loan_context_df):
    seed_sql_result(ctx_with_fake, "ctx_df", loan_context_df)
    args = TrainTabularModelArgs(
        df_label="ctx_df",
        target_column="status",
        entity_column="loan_id",
        task_type="classification",
    )
    out = TOOL.run(args, ctx_with_fake)
    assert out.ok, out.summary
    assert out.payload["model_type"] == "XGBoost"
    assert out.payload["target_column"] == "status"
    assert out.payload["entity_column"] == "loan_id"
    assert out.payload["model_id"].startswith("XGBoost_")
    # Cache contract
    assert ctx_with_fake.cache["last_model_id"] == out.payload["model_id"]
    assert ctx_with_fake.cache["last_train_tag"] == out.payload["tag"]
    assert ctx_with_fake.cache["last_train_df_label"] == "ctx_df"
    assert ctx_with_fake.cache["last_model_task"]["target_column"] == "status"


def test_happy_path_foundation_model_requires_compute(
    ctx_with_fake, seed_sql_result, loan_context_df
):
    seed_sql_result(ctx_with_fake, "ctx_df", loan_context_df)
    args = TrainTabularModelArgs(
        df_label="ctx_df",
        target_column="status",
        entity_column="loan_id",
        model_type="TabPFN",
        compute_type="small",
    )
    out = TOOL.run(args, ctx_with_fake)
    assert out.ok, out.summary
    assert out.payload["model_type"] == "TabPFN"
    assert "TabPFN" in out.payload["model_id"]


def test_happy_path_sample_percentage_passthrough(
    ctx_with_fake, seed_sql_result, loan_context_df
):
    seed_sql_result(ctx_with_fake, "ctx_df", loan_context_df)
    args = TrainTabularModelArgs(
        df_label="ctx_df",
        target_column="status",
        entity_column="loan_id",
        sample_percentage=0.5,
    )
    out = TOOL.run(args, ctx_with_fake)
    assert out.ok


# ----- error branches -----


def test_missing_context_df(ctx_with_fake):
    args = TrainTabularModelArgs(
        df_label="not_present",
        target_column="status",
        entity_column="loan_id",
    )
    out = TOOL.run(args, ctx_with_fake)
    assert not out.ok
    assert out.error == "missing_df"
    assert "no DataFrame cached" in out.summary


def test_bad_target_column(ctx_with_fake, seed_sql_result, loan_context_df):
    seed_sql_result(ctx_with_fake, "ctx_df", loan_context_df)
    args = TrainTabularModelArgs(
        df_label="ctx_df",
        target_column="nonexistent",
        entity_column="loan_id",
    )
    out = TOOL.run(args, ctx_with_fake)
    assert not out.ok
    assert out.error == "bad_target"


def test_empty_target_column(ctx_with_fake, seed_sql_result, loan_context_df):
    df = loan_context_df.copy()
    df["status"] = None
    seed_sql_result(ctx_with_fake, "ctx_df", df)
    args = TrainTabularModelArgs(
        df_label="ctx_df",
        target_column="status",
        entity_column="loan_id",
    )
    out = TOOL.run(args, ctx_with_fake)
    assert not out.ok
    assert out.error == "empty_target"


def test_foundation_model_without_compute_type(
    ctx_with_fake, seed_sql_result, loan_context_df
):
    seed_sql_result(ctx_with_fake, "ctx_df", loan_context_df)
    args = TrainTabularModelArgs(
        df_label="ctx_df",
        target_column="status",
        entity_column="loan_id",
        model_type="OrionMSP",
    )
    out = TOOL.run(args, ctx_with_fake)
    assert not out.ok
    assert out.error == "missing_compute_type"
    assert "OrionMSP" in out.summary


def test_no_tab_project(ctx_offline, seed_sql_result, loan_context_df):
    seed_sql_result(ctx_offline, "ctx_df", loan_context_df)
    args = TrainTabularModelArgs(
        df_label="ctx_df",
        target_column="status",
        entity_column="loan_id",
    )
    out = TOOL.run(args, ctx_offline)
    assert not out.ok
    assert out.error == "no_tab_project"


# ----- helpers -----


def test_resolve_entity_col_explicit(loan_context_df):
    args = TrainTabularModelArgs(
        df_label="x", target_column="status", entity_column="loan_id"
    )
    col, warn = _resolve_entity_col(args, loan_context_df)
    assert col == "loan_id"
    assert warn is None


def test_resolve_entity_col_auto_pick(loan_context_df):
    args = TrainTabularModelArgs(df_label="x", target_column="status")
    col, warn = _resolve_entity_col(args, loan_context_df)
    assert col == "loan_id"
    assert warn is not None and "auto-picked" in warn


def test_resolve_entity_col_no_unique_column():
    df = pd.DataFrame({"a": [1, 1, 1], "b": ["x", "x", "x"], "y": [0, 1, 0]})
    args = TrainTabularModelArgs(df_label="x", target_column="y")
    col, warn = _resolve_entity_col(args, df)
    assert col is None


def test_mint_train_tag_is_safe():
    """Lexsi rejects ANY non-alphanumeric in tag names (including `_`).
    Strip everything except [a-zA-Z0-9]; uniqueness comes from run_id."""
    tag = _mint_train_tag("run123", "ctx-df with spaces & chars!")
    assert tag.isalnum(), f"non-alphanumeric chars in tag: {tag!r}"
    assert tag.startswith("agenttrain")
    assert "run123" in tag  # run_id preserved


def test_build_project_config_has_required_keys():
    """ProjectConfig dict carries the identifier / hygiene / sampling / XAI
    keys. `pred_label` is intentionally omitted (no longer required by the
    current lexsi-sdk upload+train path)."""
    cfg = _build_project_config(
        unique_identifier="loan_id",
        true_label="status",
        tag="agent_train_x",
        model_name="XGBoost",
        sample_percentage=None,
    )
    expected_keys = {
        "unique_identifier", "true_label", "tag", "model_name",
        "feature_exclude", "feature_encodings", "drop_duplicate_uid",
        "handle_errors", "handle_data_imbalance", "sample_percentage", "xai_method",
    }
    assert set(cfg.keys()) == expected_keys, (
        f"ProjectConfig key set drift detected. "
        f"Missing: {expected_keys - set(cfg.keys())}, "
        f"Extra: {set(cfg.keys()) - expected_keys}"
    )
    # Values
    assert cfg["unique_identifier"] == "loan_id"
    assert cfg["true_label"] == "status"
    assert cfg["tag"] == "agent_train_x"
    # XAI defaults to [] β€” training-time SHAP fails on fresh projects
    # with "could not convert string to float: 'class_X'" (activity log
    # #05-013). Per-case SHAP runs later via case_predict.
    assert cfg["xai_method"] == []
    # pred_label is no longer emitted β€” the SDK doesn't require it.
    assert "pred_label" not in cfg
    assert cfg["sample_percentage"] is None


def test_build_project_config_includes_sample_percentage():
    cfg = _build_project_config(
        unique_identifier="loan_id",
        true_label="status",
        tag="x",
        model_name="XGBoost",
        sample_percentage=0.5,
    )
    assert cfg["sample_percentage"] == 0.5


def test_resolve_active_model_name_uses_models_fallback(fake_tab_project, loan_context_df):
    # Train one model
    fake_tab_project.upload_data(
        loan_context_df, tag="t1",
        config={"unique_identifier": "loan_id", "true_label": "status"},
    )
    name = _resolve_active_model_name(fake_tab_project)
    assert name.startswith("XGBoost_")