GCMD_Keyword_Classifier_MVP / tests /test_llm_config.py
igerasimov's picture
MVP Milestone 5
d840c10
Raw
History Blame Contribute Delete
2.34 kB
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"