diff --git a/src/llm/prompt/builder.py b/src/llm/prompt/builder.py index ffa0ed1c6..57ffd84ef 100644 --- a/src/llm/prompt/builder.py +++ b/src/llm/prompt/builder.py @@ -118,7 +118,7 @@ _PROVIDER_TEMPLATE_MAP: dict[str, dict[str, Any]] = { def get_provider_builder(provider_name: str) -> PromptBuilder: - config_dict = _PROVIDER_TEMPLATE_MAP.get(provider_name.lower(), {}) + config_dict = _PROVIDER_TEMPLATE_MAP.get(provider_name.strip().lower(), {}) config = PromptConfig(**config_dict) return PromptBuilder(config) diff --git a/tests/test_builder.py b/tests/test_builder.py index 439967e91..8fcd2e742 100644 --- a/tests/test_builder.py +++ b/tests/test_builder.py @@ -83,3 +83,13 @@ class TestAdaptMessagesForProvider: messages = [Message(role=Role.USER, content="Hello")] result = adapt_messages_for_provider(messages, "ollama") assert len(result) == 1 + + def test_provider_names_allow_outer_whitespace(self): + messages = [Message(role=Role.USER, content="Hello")] + tools = [ToolDefinition(name="search", description="Search the web", parameters={})] + + result = adapt_messages_for_provider(messages, " ollama ", tools) + + assert len(result) == 2 + assert result[0].role == Role.SYSTEM + assert "Available Tools" in result[0].content