File size: 6,330 Bytes
d95323c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Tests for the fine-tuning dataset prep script.

No network / no OpenAI calls — pure JSON transformation tests.
"""
from __future__ import annotations

import json
import subprocess
import sys
from pathlib import Path

ROOT = Path(__file__).resolve().parents[2]


def _write_smoke(tmp_path: Path, rows: list[dict]) -> Path:
    """Write a JSONL fixture in the same shape as evaluation/smoke_*_sample.jsonl."""
    p = tmp_path / "smoke.jsonl"
    with p.open("w") as f:
        for r in rows:
            f.write(json.dumps(r) + "\n")
    return p


def _run_prep(input_path: Path, out_prefix: Path, doc_type: str = "receipt") -> subprocess.CompletedProcess:
    """Always pass absolute paths — the script handles both absolute and relative."""
    return subprocess.run(
        [sys.executable, "scripts/prep_ft_dataset.py",
         "--input", str(input_path),
         "--doc-type", doc_type,
         "--out", str(out_prefix),
         "--val-frac", "0.2",
         "--seed", "42"],
        cwd=ROOT,
        capture_output=True,
        text=True,
    )


def _valid_smoke_rows(n: int) -> list[dict]:
    rows = []
    for i in range(n):
        rows.append({
            "id": f"synthetic_{i}",
            "source": "test",
            "text": f"MERCHANT {i}\nTOTAL USD {10 + i}.99\n",
            "ground_truth": {
                "merchant": f"MERCHANT {i}",
                "total": 10.0 + i + 0.99,
                "currency": "USD",
            },
        })
    return rows


# --- shape ------------------------------------------------------------

def test_prep_produces_openai_chat_format(tmp_path):
    input_path = _write_smoke(tmp_path, _valid_smoke_rows(10))
    out = tmp_path / "ft" / "receipt"
    r = _run_prep(input_path, out)
    assert r.returncode == 0, r.stderr

    train_path = out.with_name(out.name + "_train.jsonl")
    val_path   = out.with_name(out.name + "_val.jsonl")
    assert train_path.exists()
    assert val_path.exists()

    with train_path.open() as f:
        first = json.loads(f.readline())
    assert "messages" in first
    roles = [m["role"] for m in first["messages"]]
    assert roles == ["system", "user", "assistant"]


def test_assistant_content_is_valid_envelope_json(tmp_path):
    input_path = _write_smoke(tmp_path, _valid_smoke_rows(10))
    out = tmp_path / "ft" / "receipt"
    _run_prep(input_path, out)

    train_path = out.with_name(out.name + "_train.jsonl")
    with train_path.open() as f:
        for line in f:
            row = json.loads(line)
            asst = row["messages"][2]["content"]
            envelope = json.loads(asst)
            assert set(envelope.keys()) == {"data", "field_confidences", "warnings"}
            assert isinstance(envelope["data"], dict)
            assert isinstance(envelope["field_confidences"], list)
            assert isinstance(envelope["warnings"], list)


def test_system_prompt_is_the_production_prompt(tmp_path):
    input_path = _write_smoke(tmp_path, _valid_smoke_rows(10))
    out = tmp_path / "ft" / "receipt"
    _run_prep(input_path, out)

    train_path = out.with_name(out.name + "_train.jsonl")
    with train_path.open() as f:
        first = json.loads(f.readline())
    sys_msg = first["messages"][0]["content"]
    # The receipt prompt starts with a specific opening line — verify we\'re
    # baking THAT into the fine-tune, not some ad-hoc other prompt.
    assert "receipts" in sys_msg.lower()
    assert "CRITICAL EXTRACTION RULES" in sys_msg


def test_user_message_matches_production_shape(tmp_path):
    """The user-message format has to match what DocumentExtractor sends at
    inference. If they diverge, the fine-tune won\'t transfer."""
    input_path = _write_smoke(tmp_path, _valid_smoke_rows(10))
    out = tmp_path / "ft" / "receipt"
    _run_prep(input_path, out)

    train_path = out.with_name(out.name + "_train.jsonl")
    with train_path.open() as f:
        first = json.loads(f.readline())
    user_msg = first["messages"][1]["content"]
    assert "---BEGIN DOCUMENT TEXT---" in user_msg
    assert "---END DOCUMENT TEXT---" in user_msg


# --- split ------------------------------------------------------------

def test_split_80_20_and_deterministic(tmp_path):
    input_path = _write_smoke(tmp_path, _valid_smoke_rows(10))
    out = tmp_path / "ft" / "receipt"
    _run_prep(input_path, out)

    n_train = sum(1 for _ in (out.with_name(out.name + "_train.jsonl")).open())
    n_val   = sum(1 for _ in (out.with_name(out.name + "_val.jsonl")).open())
    assert n_train + n_val == 10
    # 20% of 10 = 2 val rows, 8 train rows.
    assert n_val == 2
    assert n_train == 8


# --- error handling ---------------------------------------------------

def test_skips_rows_missing_text_without_crashing(tmp_path):
    rows = _valid_smoke_rows(9) + [{
        "id": "no_text", "source": "test",
        "ground_truth": {"merchant": "X", "total": 1.0, "currency": "USD"},
    }]
    input_path = _write_smoke(tmp_path, rows)
    out = tmp_path / "ft" / "receipt"
    r = _run_prep(input_path, out)
    assert r.returncode == 0
    assert "no_text" in r.stderr  # skipped-line warning printed to stderr


def test_exit_code_2_when_input_missing(tmp_path):
    r = _run_prep(tmp_path / "does_not_exist.jsonl", tmp_path / "ft" / "r")
    assert r.returncode == 2


def test_warns_below_openai_minimum(tmp_path):
    input_path = _write_smoke(tmp_path, _valid_smoke_rows(5))
    out = tmp_path / "ft" / "receipt"
    r = _run_prep(input_path, out)
    assert r.returncode == 0
    assert ">= 10" in r.stderr  # the explicit-minimum warning


def test_filing_doc_type_uses_filing_prompt(tmp_path):
    rows = [{
        "id": "syn_filing",
        "source": "test",
        "text": "Company FORM 10-K\nRevenue: $100M",
        "ground_truth": {
            "cover":       {"company_name": "Test Co"},
            "financials":  {"currency": "USD"},
            "top_risk_factors": [],
        },
    }] * 10
    input_path = _write_smoke(tmp_path, rows)
    out = tmp_path / "ft" / "filing"
    r = _run_prep(input_path, out, doc_type="filing")
    assert r.returncode == 0, r.stderr

    train_path = out.with_name(out.name + "_train.jsonl")
    with train_path.open() as f:
        first = json.loads(f.readline())
    assert "10-K" in first["messages"][0]["content"]