diff --git a/tests/guardrails_tests/test_aim_guardrail_base_url.py b/tests/guardrails_tests/test_aim_guardrail_base_url.py index 3adb40db2c3..c235e83b934 100644 --- a/tests/guardrails_tests/test_aim_guardrail_base_url.py +++ b/tests/guardrails_tests/test_aim_guardrail_base_url.py @@ -1,24 +1,21 @@ +import pytest from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail -from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail -import os -# Set a dummy API key for testing -os.environ["AIM_API_KEY"] = "test-key" -def test_aim_base_url_trailing_slash(): - # Test with trailing slash - guardrail = AimGuardrail(api_base="https://api.aim.security/") +@pytest.mark.parametrize("api_base", [ + "https://api.aim.security/", + "https://api.aim.security", +]) +def test_aim_base_url_trailing_slash(monkeypatch, api_base): + monkeypatch.setenv("AIM_API_KEY", "test-key") + guardrail = AimGuardrail(api_base=api_base) assert guardrail.api_base == "https://api.aim.security" - assert guardrail.ws_api_base.startswith("wss://api.aim.security") + assert guardrail.ws_api_base == "wss://api.aim.security" - # Test without trailing slash - guardrail2 = AimGuardrail(api_base="https://api.aim.security") - assert guardrail2.api_base == "https://api.aim.security" - assert guardrail2.ws_api_base.startswith("wss://api.aim.security") - # Test with environment variable - os.environ["AIM_API_BASE"] = "https://api.aim.security/" - guardrail3 = AimGuardrail(api_base=None) - assert guardrail3.api_base == "https://api.aim.security" - assert guardrail3.ws_api_base.startswith("wss://api.aim.security") - del os.environ["AIM_API_BASE"] +def test_aim_base_url_from_env(monkeypatch): + monkeypatch.setenv("AIM_API_KEY", "test-key") + monkeypatch.setenv("AIM_API_BASE", "https://api.aim.security/") + guardrail = AimGuardrail(api_base=None) + assert guardrail.api_base == "https://api.aim.security" + assert guardrail.ws_api_base == "wss://api.aim.security"