fix(assemblyai): detect EU region from passthrough path

This commit is contained in:
Minh Vu 2026-05-24 10:51:03 +02:00
parent 4148667671
commit b23ca933cc
2 changed files with 82 additions and 64 deletions

View file

@ -326,7 +326,11 @@ class AssemblyAIPassthroughLoggingHandler:
"""
if url is None:
return None
if urlparse(url).hostname == "eu.assemblyai.com":
parsed_url = urlparse(url)
if parsed_url.hostname in {"eu.assemblyai.com", "api.eu.assemblyai.com"}:
return "eu"
path_parts = [part for part in parsed_url.path.split("/") if part]
if "eu.assemblyai" in path_parts:
return "eu"
return None

View file

@ -1,34 +1,13 @@
import json
import os
import sys
from datetime import datetime
from unittest.mock import AsyncMock, Mock, patch
from types import ModuleType, SimpleNamespace
from unittest.mock import patch
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system-path
import httpx
import pytest
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
import json
import os
import sys
from datetime import datetime
from unittest.mock import AsyncMock, Mock, patch
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system-path
import httpx
import pytest
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.assembly_passthrough_logging_handler import (
AssemblyAIPassthroughLoggingHandler,
AssemblyAITranscriptResponse,
@ -56,10 +35,25 @@ def mock_transcript_response():
}
def _mock_llm_passthrough_endpoint_router(api_key: str = "test-key"):
mock_module = ModuleType(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints"
)
mock_module.passthrough_endpoint_router = SimpleNamespace(
get_credentials=lambda **_: api_key
)
return patch.dict(
sys.modules,
{
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints": mock_module
},
)
def test_should_log_request():
handler = AssemblyAIPassthroughLoggingHandler()
assert handler._should_log_request("POST") == True
assert handler._should_log_request("GET") == False
assert handler._should_log_request("POST")
assert not handler._should_log_request("GET")
def test_get_assembly_transcript(assembly_handler, mock_transcript_response):
@ -68,10 +62,7 @@ def test_get_assembly_transcript(assembly_handler, mock_transcript_response):
and uses the test key returned by the mocked get_credentials.
"""
# Patch get_credentials to return "test-key"
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with _mock_llm_passthrough_endpoint_router():
with patch("httpx.get") as mock_get:
mock_get.return_value.json.return_value = mock_transcript_response
mock_get.return_value.raise_for_status.return_value = None
@ -94,10 +85,7 @@ def test_poll_assembly_for_transcript_response(
"""
Test that the _poll_assembly_for_transcript_response method returns the correct transcript response
"""
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with _mock_llm_passthrough_endpoint_router():
with patch("httpx.get") as mock_get:
mock_get.return_value.json.return_value = mock_transcript_response
mock_get.return_value.raise_for_status.return_value = None
@ -121,56 +109,85 @@ def test_is_assemblyai_route():
handler = PassThroughEndpointLogging()
# Test positive cases
assert (
handler.is_assemblyai_route("https://api.assemblyai.com/v2/transcript") == True
)
assert handler.is_assemblyai_route("https://api.assemblyai.com/other/path") == True
assert handler.is_assemblyai_route("https://api.assemblyai.com/transcript") == True
assert handler.is_assemblyai_route("https://api.assemblyai.com/v2/transcript")
assert handler.is_assemblyai_route("https://api.assemblyai.com/other/path")
assert handler.is_assemblyai_route("https://api.assemblyai.com/transcript")
# Test negative cases
assert handler.is_assemblyai_route("https://example.com/other") == False
assert (
handler.is_assemblyai_route("https://api.openai.com/v1/chat/completions")
== False
assert not handler.is_assemblyai_route("https://example.com/other")
assert not handler.is_assemblyai_route(
"https://api.openai.com/v1/chat/completions"
)
assert not handler.is_assemblyai_route("")
def test_get_assembly_region_from_url_returns_eu_for_proxy_path():
handler = AssemblyAIPassthroughLoggingHandler()
assert (
handler._get_assembly_region_from_url(
"https://proxy.company.com/eu.assemblyai/v2/transcript"
)
== "eu"
)
def test_get_assembly_region_from_url_returns_eu_for_root_path_proxy_path():
handler = AssemblyAIPassthroughLoggingHandler()
assert (
handler._get_assembly_region_from_url(
"https://proxy.company.com/litellm/eu.assemblyai/v2/transcript"
)
== "eu"
)
def test_get_assembly_region_from_url_returns_eu_for_api_eu_host():
handler = AssemblyAIPassthroughLoggingHandler()
assert (
handler._get_assembly_region_from_url(
"https://api.eu.assemblyai.com/v2/transcript"
)
== "eu"
)
def test_get_assembly_region_from_url_returns_none_for_default_route():
handler = AssemblyAIPassthroughLoggingHandler()
assert (
handler._get_assembly_region_from_url(
"https://proxy.company.com/assemblyai/v2/transcript"
)
is None
)
assert handler.is_assemblyai_route("") == False
# --- Security: SSRF via transcript_id path traversal ---
def test_get_assembly_transcript_rejects_slash_in_id(assembly_handler):
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with _mock_llm_passthrough_endpoint_router():
with pytest.raises(ValueError, match="disallowed characters"):
assembly_handler._get_assembly_transcript("../../admin/credentials")
def test_get_assembly_transcript_rejects_dotdot_in_id(assembly_handler):
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with _mock_llm_passthrough_endpoint_router():
with pytest.raises(ValueError, match="disallowed characters"):
assembly_handler._get_assembly_transcript("..evil")
def test_get_assembly_transcript_rejects_fragment_in_id(assembly_handler):
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with _mock_llm_passthrough_endpoint_router():
with pytest.raises(ValueError, match="disallowed characters"):
assembly_handler._get_assembly_transcript("abc#suffix")
def test_get_assembly_transcript_rejects_query_in_id(assembly_handler):
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with _mock_llm_passthrough_endpoint_router():
with pytest.raises(ValueError, match="disallowed characters"):
assembly_handler._get_assembly_transcript("abc?x=1")
@ -178,10 +195,7 @@ def test_get_assembly_transcript_rejects_query_in_id(assembly_handler):
def test_get_assembly_transcript_allows_valid_id(
assembly_handler, mock_transcript_response
):
with patch(
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
return_value="test-key",
):
with _mock_llm_passthrough_endpoint_router():
with patch("httpx.get") as mock_get:
mock_get.return_value.json.return_value = mock_transcript_response
mock_get.return_value.raise_for_status.return_value = None