File size: 2,337 Bytes
d840c10
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

from gcmd_classifier.config import ModelSettings
from gcmd_classifier.llm import ModelRequest, ModelStage
from gcmd_classifier.llm.schemas import TopicResponse


def test_default_model_settings_load() -> None:
    settings = ModelSettings.from_environment({})

    assert settings.provider == "fake"
    assert settings.model_name == "fake-model"
    assert settings.temperature == 0.0
    assert settings.timeout_seconds == 30.0
    assert settings.max_retries == 2


def test_environment_overrides_work() -> None:
    settings = ModelSettings.from_environment(
        {
            "MODEL_PROVIDER": "openai",
            "MODEL_NAME": "configured-model",
            "MODEL_TEMPERATURE": "0.2",
            "MODEL_TIMEOUT_SECONDS": "45",
            "MODEL_MAX_RETRIES": "4",
            "PROMPT_VERSION_TOPIC": "topic-x",
            "PROMPT_VERSION_TERM": "term-x",
            "PROMPT_VERSION_VARIABLE": "variable-x",
            "MODEL_API_KEY_ENV_VAR": "CUSTOM_API_KEY",
            "MODEL_INCLUDE_COST_METADATA": "false",
        }
    )

    assert settings.provider == "openai"
    assert settings.model_name == "configured-model"
    assert settings.temperature == 0.2
    assert settings.timeout_seconds == 45.0
    assert settings.max_retries == 4
    assert settings.prompt_version_topic == "topic-x"
    assert settings.prompt_version_term == "term-x"
    assert settings.prompt_version_variable == "variable-x"
    assert settings.api_key_env_var == "CUSTOM_API_KEY"
    assert settings.include_cost_metadata is False


def test_model_provider_and_name_are_configurable_in_requests() -> None:
    settings = ModelSettings(provider="fake-provider", model_name="fake-name")
    request = ModelRequest.from_settings(
        stage=ModelStage.TOPIC,
        prompt="prompt",
        response_schema=TopicResponse,
        settings=settings,
    )

    assert request.provider == "fake-provider"
    assert request.model_name == "fake-name"


def test_prompt_versions_are_configurable_in_requests() -> None:
    settings = ModelSettings(prompt_version_topic="topic-2")
    request = ModelRequest.from_settings(
        stage=ModelStage.TOPIC,
        prompt="prompt",
        response_schema=TopicResponse,
        settings=settings,
    )

    assert request.prompt_version == "topic-2"