Files
EvoScientist/tests/test_llm.py
T
X-iZhang 418f2d58a7 Enhance documentation and refactor code for clarity and functionality
- Updated README.md to clarify tool allowlist supports glob wildcards.
- Enhanced docstrings in prompts.py, utils.py, and formatter.py for better understanding of function parameters and return values.
- Refactored test cases in test_llm.py and test_onboard.py to use consistent mocking style with unittest.mock.patch.
- Improved test coverage and clarity in test_skills_manager.py and test_stream_state.py by adding descriptive comments and organizing sections.
- Adjusted tool result formatting logic in formatter.py to streamline success checks.
- Ensured backward compatibility in various modules while enhancing functionality.
2026-02-13 18:54:33 +00:00

234 lines
8.5 KiB
Python

"""Tests for EvoScientist LLM module."""
from unittest.mock import patch
from EvoScientist.llm import (
MODELS,
DEFAULT_MODEL,
get_chat_model,
get_models_for_provider,
list_models,
get_model_info,
)
from EvoScientist.llm.models import _MODEL_ENTRIES
# =============================================================================
# Test MODELS registry
# =============================================================================
class TestModelsRegistry:
def test_models_is_dict(self):
"""Test that MODELS is a dictionary."""
assert isinstance(MODELS, dict)
def test_entries_has_all_providers(self):
"""Test that _MODEL_ENTRIES covers native providers."""
providers = {p for _, _, p in _MODEL_ENTRIES}
assert "anthropic" in providers
assert "openai" in providers
assert "google-genai" in providers
assert "nvidia" in providers
def test_entries_are_valid_tuples(self):
"""Test that _MODEL_ENTRIES contains valid (name, model_id, provider) tuples."""
valid_providers = {"anthropic", "openai", "google-genai", "nvidia"}
for entry in _MODEL_ENTRIES:
assert len(entry) == 3, f"Entry {entry} doesn't have 3 elements"
name, model_id, provider = entry
assert isinstance(name, str)
assert isinstance(model_id, str)
assert provider in valid_providers, f"Unknown provider '{provider}' for '{name}'"
def test_get_models_for_provider(self):
"""Test that get_models_for_provider returns correct models."""
anthropic_models = get_models_for_provider("anthropic")
assert len(anthropic_models) > 0
for name, model_id in anthropic_models:
assert isinstance(name, str)
assert isinstance(model_id, str)
# Third-party providers have no registered models (user types model name)
openrouter_models = get_models_for_provider("openrouter")
assert len(openrouter_models) == 0
# =============================================================================
# Test DEFAULT_MODEL
# =============================================================================
class TestDefaultModel:
def test_default_model_exists_in_registry(self):
"""Test that DEFAULT_MODEL is a valid model in MODELS."""
assert DEFAULT_MODEL in MODELS
def test_default_model_is_anthropic(self):
"""Test that default model uses Anthropic."""
_, provider = MODELS[DEFAULT_MODEL]
assert provider == "anthropic"
# =============================================================================
# Test list_models
# =============================================================================
class TestListModels:
def test_returns_list(self):
"""Test that list_models returns a list."""
result = list_models()
assert isinstance(result, list)
def test_returns_all_model_names(self):
"""Test that list_models returns all model names."""
result = list_models()
assert set(result) == set(MODELS.keys())
def test_list_is_not_empty(self):
"""Test that the list is not empty."""
assert len(list_models()) > 0
# =============================================================================
# Test get_model_info
# =============================================================================
class TestGetModelInfo:
def test_returns_tuple_for_valid_model(self):
"""Test that get_model_info returns tuple for valid model."""
result = get_model_info("claude-sonnet-4-5")
assert result is not None
assert isinstance(result, tuple)
assert len(result) == 2
def test_returns_none_for_invalid_model(self):
"""Test that get_model_info returns None for invalid model."""
result = get_model_info("nonexistent-model")
assert result is None
def test_returns_correct_info(self):
"""Test that get_model_info returns correct info."""
model_id, provider = get_model_info("gpt-5-nano")
assert model_id == "gpt-5-nano-2025-08-07"
assert provider == "openai"
# =============================================================================
# Test get_chat_model
# =============================================================================
class TestGetChatModel:
@patch("EvoScientist.llm.models.init_chat_model")
def test_uses_default_model_when_none(self, mock_init):
"""Test that get_chat_model uses default model when model=None."""
mock_init.return_value = "mock_model"
get_chat_model()
mock_init.assert_called_once()
call_kwargs = mock_init.call_args[1]
# Default model should be resolved from MODELS
expected_model_id, expected_provider = MODELS[DEFAULT_MODEL]
assert call_kwargs["model"] == expected_model_id
assert call_kwargs["model_provider"] == expected_provider
@patch("EvoScientist.llm.models.init_chat_model")
def test_resolves_short_name(self, mock_init):
"""Test that get_chat_model resolves short names correctly."""
mock_init.return_value = "mock_model"
get_chat_model("claude-opus-4-5")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["model"] == "claude-opus-4-5-20251101"
assert call_kwargs["model_provider"] == "anthropic"
@patch("EvoScientist.llm.models.init_chat_model")
def test_resolves_openai_short_name(self, mock_init):
"""Test that get_chat_model resolves OpenAI short names."""
mock_init.return_value = "mock_model"
get_chat_model("gpt-5-mini")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["model"] == "gpt-5-mini-2025-08-07"
assert call_kwargs["model_provider"] == "openai"
@patch("EvoScientist.llm.models.init_chat_model")
def test_uses_full_model_id(self, mock_init):
"""Test that get_chat_model accepts full model IDs."""
mock_init.return_value = "mock_model"
get_chat_model("claude-3-opus-20240229")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["model"] == "claude-3-opus-20240229"
# Should infer anthropic from the model prefix
assert call_kwargs["model_provider"] == "anthropic"
@patch("EvoScientist.llm.models.init_chat_model")
def test_provider_override(self, mock_init):
"""Test that provider can be overridden."""
mock_init.return_value = "mock_model"
get_chat_model("claude-sonnet-4-5", provider="custom_provider")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["model_provider"] == "custom_provider"
@patch("EvoScientist.llm.models.init_chat_model")
def test_passes_kwargs(self, mock_init):
"""Test that additional kwargs are passed through."""
mock_init.return_value = "mock_model"
get_chat_model("gpt-5-nano", temperature=0.7, max_tokens=1000)
call_kwargs = mock_init.call_args[1]
assert call_kwargs["temperature"] == 0.7
assert call_kwargs["max_tokens"] == 1000
@patch("EvoScientist.llm.models.init_chat_model")
def test_infers_openai_from_gpt_prefix(self, mock_init):
"""Test that OpenAI is inferred from gpt- prefix."""
mock_init.return_value = "mock_model"
get_chat_model("gpt-4-turbo-preview")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["model_provider"] == "openai"
@patch("EvoScientist.llm.models.init_chat_model")
def test_infers_openai_from_o1_prefix(self, mock_init):
"""Test that OpenAI is inferred from o1 prefix."""
mock_init.return_value = "mock_model"
get_chat_model("o1-preview")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["model_provider"] == "openai"
@patch("EvoScientist.llm.models.init_chat_model")
def test_infers_google_from_gemini_prefix(self, mock_init):
"""Test that google-genai is inferred from gemini prefix."""
mock_init.return_value = "mock_model"
get_chat_model("gemini-2.0-flash")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["model_provider"] == "google-genai"
@patch("EvoScientist.llm.models.init_chat_model")
def test_defaults_to_anthropic_for_unknown(self, mock_init):
"""Test that anthropic is default for unknown model prefixes."""
mock_init.return_value = "mock_model"
get_chat_model("some-unknown-model")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["model_provider"] == "anthropic"