mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
done
This commit is contained in:
parent
6600c86dbd
commit
3253a5147c
2 changed files with 22 additions and 0 deletions
|
|
@ -58,6 +58,7 @@ class AimGuardrail(CustomGuardrail):
|
|||
self.api_base = (
|
||||
api_base or os.environ.get("AIM_API_BASE") or "https://api.aim.security"
|
||||
)
|
||||
self.api_base = self.api_base.rstrip("/")
|
||||
self.ws_api_base = self.api_base.replace("http://", "ws://").replace(
|
||||
"https://", "wss://"
|
||||
)
|
||||
|
|
|
|||
21
tests/guardrails_tests/test_aim_guardrail_base_url.py
Normal file
21
tests/guardrails_tests/test_aim_guardrail_base_url.py
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
import pytest
|
||||
import os
|
||||
from litellm.proxy.guardrails.guardrail_hooks.aim.aim import CustomGuardrail
|
||||
|
||||
def test_aim_base_url_trailing_slash():
|
||||
# Test with trailing slash
|
||||
guardrail = CustomGuardrail(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")
|
||||
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)
|
||||
assert guardrail3.api_base == "https://api.aim.security"
|
||||
assert guardrail3.ws_api_base.startswith("wss://api.aim.security")
|
||||
del os.environ["AIM_API_BASE"]
|
||||
Loading…
Add table
Reference in a new issue