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"
|