diff --git a/tests/guardrails_tests/test_aim_guardrail_base_url.py b/tests/guardrails_tests/test_aim_guardrail_base_url.py index 3e230bb32cc..105ddc89b17 100644 --- a/tests/guardrails_tests/test_aim_guardrail_base_url.py +++ b/tests/guardrails_tests/test_aim_guardrail_base_url.py @@ -1,21 +1,24 @@ import pytest +from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail import os -from litellm.proxy.guardrails.guardrail_hooks.aim.aim import CustomGuardrail + +# 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 = CustomGuardrail(api_base="https://api.aim.security/") + guardrail = AimGuardrail(api_base="https://api.aim.security/") assert guardrail.api_base == "https://api.aim.security" assert guardrail.ws_api_base.startswith("wss://api.aim.security") # Test without trailing slash - guardrail2 = CustomGuardrail(api_base="https://api.aim.security") + 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 = CustomGuardrail(api_base=None) + 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"]