File size: 3,071 Bytes
f1b0176
06feb89
f1b0176
 
 
 
06feb89
f1b0176
 
 
06feb89
f1b0176
 
 
 
 
06feb89
 
 
f1b0176
 
06feb89
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f1b0176
06feb89
 
f1b0176
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import unittest
from unittest.mock import patch

from training_coach.parser_runtime import (
    DEFAULT_TRANSFORMERS_MODEL,
    PARSER_MODEL,
    MODEL_CACHE_ENV_VAR,
    ParserRuntimeUnavailableError,
    _load_transformers,
    generate_parser_response,
    load_parser_model,
    parse_check_in_with_model,
)


class ParserRuntimeTest(unittest.TestCase):
    def tearDown(self):
        load_parser_model.cache_clear()

    def test_runtime_exports_model_name(self):
        self.assertEqual(PARSER_MODEL, "Qwen/Qwen2.5-1.5B-Instruct")
        self.assertEqual(DEFAULT_TRANSFORMERS_MODEL, "Qwen/Qwen3-1.7B")

    def test_runtime_uses_default_huggingface_cache_by_default(self):
        self.assertEqual(MODEL_CACHE_ENV_VAR, "PARSER_MODEL_CACHE_DIR")

    def test_load_parser_model_uses_default_hf_cache_without_override(self):
        class FakeTokenizer:
            @classmethod
            def from_pretrained(cls, model_name, cache_dir=None):
                return {"model_name": model_name, "cache_dir": cache_dir}

        class FakeModel:
            @classmethod
            def from_pretrained(cls, model_name, **kwargs):
                return {"model_name": model_name, **kwargs}

        load_parser_model.cache_clear()
        with patch.dict("os.environ", {}, clear=True), patch(
            "training_coach.parser_runtime._load_transformers",
            return_value=(FakeModel, FakeTokenizer),
        ):
            tokenizer, model = load_parser_model("test/model")

        self.assertIsNone(tokenizer["cache_dir"])
        self.assertIsNone(model["cache_dir"])

    def test_load_parser_model_allows_explicit_cache_override(self):
        class FakeTokenizer:
            @classmethod
            def from_pretrained(cls, model_name, cache_dir=None):
                return {"model_name": model_name, "cache_dir": cache_dir}

        class FakeModel:
            @classmethod
            def from_pretrained(cls, model_name, **kwargs):
                return {"model_name": model_name, **kwargs}

        load_parser_model.cache_clear()
        with patch.dict(
            "os.environ",
            {MODEL_CACHE_ENV_VAR: "/tmp/parser-cache"},
            clear=True,
        ), patch(
            "training_coach.parser_runtime._load_transformers",
            return_value=(FakeModel, FakeTokenizer),
        ):
            tokenizer, model = load_parser_model("test/model")

        self.assertEqual(tokenizer["cache_dir"], "/tmp/parser-cache")
        self.assertEqual(model["cache_dir"], "/tmp/parser-cache")

    def test_runtime_can_import_transformers_or_reports_clear_error(self):
        try:
            loaded = _load_transformers()
        except ParserRuntimeUnavailableError as error:
            self.assertIn("Install transformers", str(error))
        else:
            self.assertEqual(len(loaded), 2)

    def test_runtime_functions_are_callable(self):
        self.assertTrue(callable(generate_parser_response))
        self.assertTrue(callable(parse_check_in_with_model))


if __name__ == "__main__":
    unittest.main()