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