mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: encode additional provider path identifiers
This commit is contained in:
parent
1d7778673a
commit
124379e42e
6 changed files with 161 additions and 17 deletions
|
|
@ -36,6 +36,7 @@ from typing import (
|
|||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.azure_ai.agents.transformation import (
|
||||
AzureAIAgentsConfig,
|
||||
AzureAIAgentsError,
|
||||
|
|
@ -75,20 +76,29 @@ class AzureAIAgentsHandler:
|
|||
def _build_messages_url(
|
||||
self, api_base: str, thread_id: str, api_version: str
|
||||
) -> str:
|
||||
return f"{api_base}/threads/{thread_id}/messages?api-version={api_version}"
|
||||
encoded_thread_id = encode_url_path_segment(thread_id, field_name="thread_id")
|
||||
return (
|
||||
f"{api_base}/threads/{encoded_thread_id}/messages?api-version={api_version}"
|
||||
)
|
||||
|
||||
def _build_runs_url(self, api_base: str, thread_id: str, api_version: str) -> str:
|
||||
return f"{api_base}/threads/{thread_id}/runs?api-version={api_version}"
|
||||
encoded_thread_id = encode_url_path_segment(thread_id, field_name="thread_id")
|
||||
return f"{api_base}/threads/{encoded_thread_id}/runs?api-version={api_version}"
|
||||
|
||||
def _build_run_status_url(
|
||||
self, api_base: str, thread_id: str, run_id: str, api_version: str
|
||||
) -> str:
|
||||
return f"{api_base}/threads/{thread_id}/runs/{run_id}?api-version={api_version}"
|
||||
encoded_thread_id = encode_url_path_segment(thread_id, field_name="thread_id")
|
||||
encoded_run_id = encode_url_path_segment(run_id, field_name="run_id")
|
||||
return f"{api_base}/threads/{encoded_thread_id}/runs/{encoded_run_id}?api-version={api_version}"
|
||||
|
||||
def _build_list_messages_url(
|
||||
self, api_base: str, thread_id: str, api_version: str
|
||||
) -> str:
|
||||
return f"{api_base}/threads/{thread_id}/messages?api-version={api_version}"
|
||||
encoded_thread_id = encode_url_path_segment(thread_id, field_name="thread_id")
|
||||
return (
|
||||
f"{api_base}/threads/{encoded_thread_id}/messages?api-version={api_version}"
|
||||
)
|
||||
|
||||
def _build_create_thread_and_run_url(self, api_base: str, api_version: str) -> str:
|
||||
"""URL for the create-thread-and-run endpoint (supports streaming)."""
|
||||
|
|
|
|||
|
|
@ -11,13 +11,14 @@ import httpx
|
|||
from httpx import Headers
|
||||
|
||||
import litellm
|
||||
from litellm.types.utils import all_litellm_params
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.text_to_speech.transformation import (
|
||||
BaseTextToSpeechConfig,
|
||||
TextToSpeechRequestData,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.utils import all_litellm_params
|
||||
|
||||
from ..common_utils import ElevenLabsException
|
||||
|
||||
|
|
@ -321,7 +322,8 @@ class ElevenLabsTextToSpeechConfig(BaseTextToSpeechConfig):
|
|||
"ElevenLabs voice_id is required. Pass `voice` when calling `litellm.speech()`."
|
||||
)
|
||||
|
||||
url = f"{base_url}{self.TTS_ENDPOINT_PATH}/{voice_id}"
|
||||
encoded_voice_id = encode_url_path_segment(voice_id, field_name="voice_id")
|
||||
url = f"{base_url}{self.TTS_ENDPOINT_PATH}/{encoded_voice_id}"
|
||||
|
||||
query_params = litellm_params.get(self.ELEVENLABS_QUERY_PARAMS_KEY, {})
|
||||
if query_params:
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import httpx
|
|||
from openai.types.file_deleted import FileDeleted
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
|
||||
from litellm.llms.base_llm.files.transformation import (
|
||||
BaseFilesConfig,
|
||||
|
|
@ -258,10 +259,14 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
|
|||
normalized_file_id = file_id
|
||||
|
||||
normalized_file_id = normalized_file_id.strip("/")
|
||||
if not normalized_file_id.startswith("files/"):
|
||||
normalized_file_id = f"files/{normalized_file_id}"
|
||||
if normalized_file_id.startswith("files/"):
|
||||
normalized_file_id = normalized_file_id.removeprefix("files/")
|
||||
|
||||
return normalized_file_id
|
||||
encoded_file_id = encode_url_path_segment(
|
||||
normalized_file_id, field_name="file_id"
|
||||
)
|
||||
|
||||
return f"files/{encoded_file_id}"
|
||||
|
||||
def transform_retrieve_file_response(
|
||||
self,
|
||||
|
|
@ -337,13 +342,8 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
|
|||
if not api_key:
|
||||
raise ValueError("api_key is required")
|
||||
|
||||
# Extract file name from URI if full URI is provided
|
||||
# file_id could be "files/abc123" or "https://generativelanguage.googleapis.com/v1beta/files/abc123"
|
||||
if file_id.startswith("http"):
|
||||
# Extract the file path from full URI
|
||||
file_name = file_id.split("/v1beta/")[-1]
|
||||
else:
|
||||
file_name = file_id if file_id.startswith("files/") else f"files/{file_id}"
|
||||
# Normalize and encode the file name before interpolating it into the URL.
|
||||
file_name = self._normalize_gemini_file_id(file_id)
|
||||
|
||||
# Construct the delete URL
|
||||
url = f"{api_base}/v1beta/{file_name}"
|
||||
|
|
|
|||
|
|
@ -0,0 +1,57 @@
|
|||
import pytest
|
||||
|
||||
from litellm.llms.azure_ai.agents.handler import AzureAIAgentsHandler
|
||||
|
||||
|
||||
def test_should_encode_thread_id_in_azure_ai_agent_urls():
|
||||
handler = AzureAIAgentsHandler()
|
||||
|
||||
assert (
|
||||
handler._build_messages_url(
|
||||
"https://example.services.ai.azure.com/api/projects/proj",
|
||||
"../../threads/other?x=1#frag",
|
||||
"2024-05-01-preview",
|
||||
)
|
||||
== "https://example.services.ai.azure.com/api/projects/proj/threads/..%2F..%2Fthreads%2Fother%3Fx%3D1%23frag/messages?api-version=2024-05-01-preview"
|
||||
)
|
||||
assert (
|
||||
handler._build_runs_url(
|
||||
"https://example.services.ai.azure.com/api/projects/proj",
|
||||
"thread/abc",
|
||||
"2024-05-01-preview",
|
||||
)
|
||||
== "https://example.services.ai.azure.com/api/projects/proj/threads/thread%2Fabc/runs?api-version=2024-05-01-preview"
|
||||
)
|
||||
|
||||
|
||||
def test_should_encode_thread_and_run_ids_in_azure_ai_agent_status_url():
|
||||
handler = AzureAIAgentsHandler()
|
||||
|
||||
assert (
|
||||
handler._build_run_status_url(
|
||||
"https://example.services.ai.azure.com/api/projects/proj",
|
||||
"thread/abc",
|
||||
"../runs/other#frag",
|
||||
"2024-05-01-preview",
|
||||
)
|
||||
== "https://example.services.ai.azure.com/api/projects/proj/threads/thread%2Fabc/runs/..%2Fruns%2Fother%23frag?api-version=2024-05-01-preview"
|
||||
)
|
||||
|
||||
|
||||
def test_should_reject_dot_segments_in_azure_ai_agent_urls():
|
||||
handler = AzureAIAgentsHandler()
|
||||
|
||||
with pytest.raises(ValueError, match="thread_id cannot be a dot path segment"):
|
||||
handler._build_messages_url(
|
||||
"https://example.services.ai.azure.com/api/projects/proj",
|
||||
"..",
|
||||
"2024-05-01-preview",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="run_id cannot be a dot path segment"):
|
||||
handler._build_run_status_url(
|
||||
"https://example.services.ai.azure.com/api/projects/proj",
|
||||
"thread_123",
|
||||
"..",
|
||||
"2024-05-01-preview",
|
||||
)
|
||||
|
|
@ -0,0 +1,33 @@
|
|||
import pytest
|
||||
|
||||
from litellm.llms.elevenlabs.text_to_speech.transformation import (
|
||||
ElevenLabsTextToSpeechConfig,
|
||||
)
|
||||
|
||||
|
||||
def test_should_encode_elevenlabs_voice_id_path_segment():
|
||||
config = ElevenLabsTextToSpeechConfig()
|
||||
|
||||
url = config.get_complete_url(
|
||||
model="elevenlabs/tts",
|
||||
api_base="https://api.elevenlabs.io",
|
||||
litellm_params={
|
||||
config.ELEVENLABS_VOICE_ID_KEY: "voice/../../models?x=1#frag",
|
||||
},
|
||||
)
|
||||
|
||||
assert (
|
||||
url
|
||||
== "https://api.elevenlabs.io/v1/text-to-speech/voice%2F..%2F..%2Fmodels%3Fx%3D1%23frag"
|
||||
)
|
||||
|
||||
|
||||
def test_should_reject_dot_segment_elevenlabs_voice_id():
|
||||
config = ElevenLabsTextToSpeechConfig()
|
||||
|
||||
with pytest.raises(ValueError, match="voice_id cannot be a dot path segment"):
|
||||
config.get_complete_url(
|
||||
model="elevenlabs/tts",
|
||||
api_base="https://api.elevenlabs.io",
|
||||
litellm_params={config.ELEVENLABS_VOICE_ID_KEY: ".."},
|
||||
)
|
||||
|
|
@ -2,7 +2,6 @@
|
|||
Test Google AI Studio (Gemini) files transformation functionality
|
||||
"""
|
||||
|
||||
import os
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -93,6 +92,30 @@ class TestGoogleAIStudioFilesTransformation:
|
|||
assert "key=" not in url
|
||||
assert params == {}
|
||||
|
||||
def test_transform_retrieve_file_request_encodes_file_id_path_segment(self):
|
||||
file_id = "files/../../models/gemini-pro?x=1#frag"
|
||||
litellm_params = {"api_key": "test-api-key"}
|
||||
|
||||
url, params = self.handler.transform_retrieve_file_request(
|
||||
file_id=file_id,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert (
|
||||
url
|
||||
== "https://generativelanguage.googleapis.com/v1beta/files/..%2F..%2Fmodels%2Fgemini-pro%3Fx%3D1%23frag"
|
||||
)
|
||||
assert params == {}
|
||||
|
||||
def test_transform_retrieve_file_request_rejects_dot_path_segment(self):
|
||||
with pytest.raises(ValueError, match="file_id cannot be a dot path segment"):
|
||||
self.handler.transform_retrieve_file_request(
|
||||
file_id="files/..",
|
||||
optional_params={},
|
||||
litellm_params={"api_key": "test-api-key"},
|
||||
)
|
||||
|
||||
@patch.dict("os.environ", {}, clear=True)
|
||||
@patch("litellm.llms.gemini.common_utils.get_secret_str", return_value=None)
|
||||
def test_transform_retrieve_file_request_missing_api_key(self, mock_get_secret):
|
||||
|
|
@ -322,3 +345,22 @@ class TestGoogleAIStudioFilesTransformation:
|
|||
assert file_id in url
|
||||
assert "generativelanguage.googleapis.com" in url
|
||||
assert params == {}
|
||||
|
||||
def test_transform_delete_file_request_encodes_file_id_path_segment(self):
|
||||
file_id = "files/../../models/gemini-pro?x=1#frag"
|
||||
litellm_params = {
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://generativelanguage.googleapis.com",
|
||||
}
|
||||
|
||||
url, params = self.handler.transform_delete_file_request(
|
||||
file_id=file_id,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert (
|
||||
url
|
||||
== "https://generativelanguage.googleapis.com/v1beta/files/..%2F..%2Fmodels%2Fgemini-pro%3Fx%3D1%23frag"
|
||||
)
|
||||
assert params == {}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue