diff --git a/docs/my-website/docs/proxy/pass_through.md b/docs/my-website/docs/proxy/pass_through.md index cf8168764b8..f47d7064140 100644 --- a/docs/my-website/docs/proxy/pass_through.md +++ b/docs/my-website/docs/proxy/pass_through.md @@ -58,6 +58,17 @@ Configure the required authentication and pricing: - The Bria API requires an `api_token` header - Enter your Bria API key as the value for the `api_token` header +**Default Query Parameters (Optional):** +- Add query parameters that will be automatically sent with every request +- Perfect for API versioning, format specifications, or default configurations +- Clients can override these parameters by providing their own values +- Example: `version=v1`, `format=json`, `timeout=30` + + + **Pricing Configuration:** - Set a cost per request (e.g., $12.00 in this example) - This enables cost tracking and billing for your users @@ -112,6 +123,9 @@ general_settings: content-type: application/json accept: application/json forward_headers: true # Forward all incoming headers + default_query_params: # Optional: Default query parameters + version: "v1" # Always send version=v1 + format: "json" # Default format (can be overridden) ``` ### Start and Test @@ -166,6 +180,9 @@ general_settings: auth: boolean # Enable LiteLLM authentication (Enterprise) forward_headers: boolean # Forward all incoming headers include_subpath: boolean # If true, forwards requests to sub-paths (default: false) + methods: list[string] # Optional: HTTP methods (e.g., ["GET", "POST"]). If not specified, all methods are supported. + default_query_params: # Optional: Default query parameters sent with every request + : string # Key-value pairs (e.g., version: "v1", format: "json") headers: # Custom headers to add Authorization: string # Auth header for target API content-type: string # Request content type @@ -177,11 +194,17 @@ general_settings: ### Header Options - **Authorization**: Authentication for the target API -- **content-type**: Request body format specification +- **content-type**: Request body format specification - **accept**: Expected response format - **LANGFUSE_PUBLIC_KEY/SECRET_KEY**: For Langfuse integration - **Custom headers**: Any additional key-value pairs +### Default Query Parameters +- **Parameter precedence**: Client params > URL params > default params +- **Use cases**: API versioning, authentication tokens, format control, feature flags +- **Override capability**: Clients can override any default parameter +- **Examples**: `version: "v1"`, `format: "json"`, `timeout: "30"` + ### Sub-path Routing By default, pass-through endpoints only match the **exact path** specified. To forward requests to sub-paths, set `include_subpath: true`: @@ -201,6 +224,92 @@ general_settings: --- +### Default Query Parameters + +Pass-through endpoints support default query parameters that are automatically added to every request. This is useful for API versioning, format specifications, authentication tokens, or any default configuration. + +#### How It Works + +**Parameter Precedence (highest to lowest priority):** +1. **Client-provided parameters** (in the request URL) +2. **URL parameters** (from the target URL) +3. **Default parameters** (from configuration) + +#### Example Configuration + +```yaml +general_settings: + pass_through_endpoints: + - path: "/api/v1" + target: "https://external-api.com/service?timeout=60" # URL has timeout=60 + default_query_params: + version: "v1" # Always add version=v1 + format: "json" # Default format=json (can be overridden) + auth_level: "basic" # Always add auth_level=basic +``` + +#### Request Examples + +**Client Request:** `GET /api/v1/users` +**Actual Backend Call:** `https://external-api.com/service?version=v1&format=json&auth_level=basic&timeout=60` + +**Client Request:** `GET /api/v1/users?format=xml&custom=value` +**Actual Backend Call:** `https://external-api.com/service?version=v1&auth_level=basic&timeout=60&format=xml&custom=value` +- Client `format=xml` overrides default `format=json` +- Default `version=v1` and `auth_level=basic` are preserved +- URL `timeout=60` is preserved +- Client `custom=value` is added + +#### Use Cases + +- **API Versioning**: Always send `version=v2` to maintain compatibility +- **Authentication**: Add authentication tokens like `api_key=default_key` +- **Format Control**: Default to `format=json` but allow client override +- **Rate Limiting**: Set `rate_limit=standard` as default +- **Feature Flags**: Enable `experimental=false` by default + +--- + +You can configure different target URLs for the same path using different HTTP methods. This is useful when different backends handle different operations: + + + +```yaml +general_settings: + pass_through_endpoints: + # GET requests to /azure/kb go to read API + - path: "/azure/kb" + target: "https://read-api.example.com/knowledge-base" + methods: ["GET"] + headers: + Authorization: "bearer os.environ/READ_API_KEY" + + # POST requests to /azure/kb go to write API + - path: "/azure/kb" + target: "https://write-api.example.com/knowledge-base" + methods: ["POST"] + headers: + Authorization: "bearer os.environ/WRITE_API_KEY" + + # PUT requests to /azure/kb go to update API + - path: "/azure/kb" + target: "https://update-api.example.com/knowledge-base" + methods: ["PUT"] + headers: + Authorization: "bearer os.environ/UPDATE_API_KEY" +``` + +**Key Points:** +- If `methods` is not specified, the endpoint supports all HTTP methods (GET, POST, PUT, DELETE, PATCH) +- Multiple endpoints can share the same path as long as they have different methods +- You can specify multiple methods for a single endpoint: `methods: ["GET", "POST"]` +- This allows you to route to different backends based on the operation type + +--- + ## Advanced: Custom Adapters For complex integrations (like Anthropic/Bedrock clients), you can create custom adapters that translate between different API schemas. diff --git a/docs/my-website/img/passthrough_method_setup.png b/docs/my-website/img/passthrough_method_setup.png new file mode 100644 index 00000000000..584e3b966c6 Binary files /dev/null and b/docs/my-website/img/passthrough_method_setup.png differ diff --git a/docs/my-website/img/passthrough_query_default.png b/docs/my-website/img/passthrough_query_default.png new file mode 100644 index 00000000000..fb97e69001e Binary files /dev/null and b/docs/my-website/img/passthrough_query_default.png differ diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index fbbf9cd2581..74af02e5e62 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -9,7 +9,9 @@ from litellm.constants import PASS_THROUGH_HEADER_PREFIX class BasePassthroughUtils: @staticmethod def get_merged_query_parameters( - existing_url: httpx.URL, request_query_params: Dict[str, Union[str, list]] + existing_url: httpx.URL, + request_query_params: Dict[str, Union[str, list]], + default_query_params: Optional[Dict[str, Union[str, list]]] = None ) -> Dict[str, Union[str, List[str]]]: # Get the existing query params from the target URL existing_query_string = existing_url.query.decode("utf-8") @@ -19,8 +21,19 @@ class BasePassthroughUtils: updated_existing_query_params = { k: v[0] if len(v) == 1 else v for k, v in existing_query_params.items() } - # Merge the query params, giving priority to the existing ones - return {**request_query_params, **updated_existing_query_params} + + # Start with default query params (lowest priority) + merged_params = {} + if default_query_params: + merged_params.update(default_query_params) + + # Override with existing URL query params (medium priority) + merged_params.update(updated_existing_query_params) + + # Override with request query params (highest priority - client can override anything) + merged_params.update(request_query_params) + + return merged_params @staticmethod def forward_headers_from_request( diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index f253f5d88ca..9a208d48392 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1916,6 +1916,10 @@ class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase): default={}, description="Key-value pairs of headers to be forwarded with the request. You can set any key value pair here and it will be forwarded to your target endpoint", ) + default_query_params: dict = Field( + default={}, + description="Key-value pairs of default query parameters to be sent with every request to this endpoint. These can be overridden by client-provided query parameters. For example: {'key': 'default_value', 'api_version': '2023-01'}", + ) include_subpath: bool = Field( default=False, description="If True, requests to subpaths of the path will be forwarded to the target endpoint. For example, if the path is /bria and include_subpath is True, requests to /bria/v1/text-to-image/base/2.3 will be forwarded to the target endpoint.", @@ -1936,6 +1940,10 @@ class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase): default=False, description="True if this endpoint is defined in the config file, False if from DB. Config-defined endpoints cannot be edited via the UI.", ) + methods: Optional[List[str]] = Field( + default=None, + description="List of HTTP methods this endpoint handles (e.g., ['GET', 'POST']). If None or empty, all methods (GET, POST, PUT, DELETE, PATCH) are supported for backward compatibility. This allows the same path to have different targets for different HTTP methods.", + ) class PassThroughEndpointResponse(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 2104a606045..883e90f8fef 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -68,7 +68,9 @@ router = APIRouter() pass_through_endpoint_logging = PassThroughEndpointLogging() # Global registry to track registered pass-through routes and prevent memory leaks -_registered_pass_through_routes: Dict[str, Dict[str, Union[str, Dict[str, Any]]]] = {} +_registered_pass_through_routes: Dict[ + str, Dict[str, Union[str, List[str], Dict[str, Any]]] +] = {} def get_response_body(response: httpx.Response) -> Optional[dict]: @@ -601,6 +603,7 @@ async def pass_through_request( # noqa: PLR0915 forward_headers: Optional[bool] = False, merge_query_params: Optional[bool] = False, query_params: Optional[dict] = None, + default_query_params: Optional[dict] = None, stream: Optional[bool] = None, cost_per_request: Optional[float] = None, custom_llm_provider: Optional[str] = None, @@ -618,6 +621,7 @@ async def pass_through_request( # noqa: PLR0915 forward_headers: Whether to forward headers merge_query_params: Whether to merge query params query_params: The query params + default_query_params: The default query params to be applied if not overridden by client stream: Whether to stream the response cost_per_request: Optional field - cost per request to the target endpoint custom_llm_provider: Optional field - custom LLM provider for the endpoint @@ -650,13 +654,18 @@ async def pass_through_request( # noqa: PLR0915 forward_headers=forward_headers, ) - if merge_query_params: + # Apply default query parameters if provided, regardless of merge_query_params setting + if default_query_params or merge_query_params: + # Determine what to merge based on settings + request_params = dict(request.query_params) if merge_query_params else {} + # Create a new URL with the merged query params url = url.copy_with( query=urlencode( HttpPassThroughEndpointHelpers.get_merged_query_parameters( existing_url=url, - request_query_params=dict(request.query_params), + request_query_params=request_params, + default_query_params=default_query_params, ) ).encode("ascii") ) @@ -958,7 +967,7 @@ async def pass_through_request( # noqa: PLR0915 if isinstance(e, HTTPException): raise ProxyException( - message=getattr(e, "message", str(e.detail)), + message=getattr(e, "message", str(getattr(e, "detail", str(e)))), type=getattr(e, "type", "None"), param=getattr(e, "param", "None"), code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), @@ -1081,6 +1090,7 @@ def create_pass_through_route( custom_llm_provider: Optional[str] = None, is_streaming_request: Optional[bool] = False, query_params: Optional[dict] = None, + default_query_params: Optional[dict] = None, guardrails: Optional[Dict[str, Any]] = None, ): # check if target is an adapter.py or a url @@ -1145,7 +1155,7 @@ def create_pass_through_route( passthrough_params = ( InitPassThroughEndpointHelpers.get_registered_pass_through_route( - route=path + route=path, method=request.method ) ) target_params = { @@ -1173,6 +1183,7 @@ def create_pass_through_route( "cost_per_request", cost_per_request ) param_guardrails = target_params.get("guardrails", None) + param_default_query_params = target_params.get("default_query_params", None) # Construct the full target URL with subpath if needed full_target = ( @@ -1210,6 +1221,7 @@ def create_pass_through_route( forward_headers=cast(Optional[bool], param_forward_headers), merge_query_params=cast(Optional[bool], param_merge_query_params), query_params=final_query_params, + default_query_params=param_default_query_params, stream=is_streaming_request or stream, custom_body=final_custom_body, cost_per_request=cast(Optional[float], param_cost_per_request), @@ -1850,20 +1862,30 @@ class InitPassThroughEndpointHelpers: cost_per_request: Optional[float], endpoint_id: str, guardrails: Optional[dict] = None, + methods: Optional[List[str]] = None, + default_query_params: Optional[dict] = None, ): """Add exact path route for pass-through endpoint""" - route_key = f"{endpoint_id}:exact:{path}" + # Default to all methods if none specified (backward compatibility) + if methods is None or len(methods) == 0: + methods = ["GET", "POST", "PUT", "DELETE", "PATCH"] + + # Create route key that includes methods for uniqueness + methods_str = ",".join(sorted(methods)) + route_key = f"{endpoint_id}:exact:{path}:{methods_str}" # Check if this exact route is already registered if route_key in _registered_pass_through_routes: verbose_proxy_logger.debug( - "Updating duplicate exact pass through endpoint: %s (already registered)", + "Updating duplicate exact pass through endpoint: %s with methods %s (already registered)", path, + methods, ) verbose_proxy_logger.debug( - "adding exact pass through endpoint: %s, dependencies: %s", + "adding exact pass through endpoint: %s, methods: %s, dependencies: %s", path, + methods, dependencies, ) @@ -1879,9 +1901,10 @@ class InitPassThroughEndpointHelpers: merge_query_params, dependencies, cost_per_request=cost_per_request, + default_query_params=default_query_params, guardrails=guardrails, ), - methods=["GET", "POST", "PUT", "DELETE", "PATCH"], + methods=methods, dependencies=dependencies, ) @@ -1890,11 +1913,13 @@ class InitPassThroughEndpointHelpers: "endpoint_id": endpoint_id, "path": path, "type": "exact", + "methods": methods, "passthrough_params": { "target": target, "custom_headers": custom_headers, "forward_headers": forward_headers, "merge_query_params": merge_query_params, + "default_query_params": default_query_params, "dependencies": dependencies, "cost_per_request": cost_per_request, "guardrails": guardrails, @@ -1913,21 +1938,30 @@ class InitPassThroughEndpointHelpers: cost_per_request: Optional[float], endpoint_id: str, guardrails: Optional[dict] = None, + methods: Optional[List[str]] = None, + default_query_params: Optional[dict] = None, ): """Add wildcard route for sub-paths""" + # Default to all methods if none specified (backward compatibility) + if methods is None or len(methods) == 0: + methods = ["GET", "POST", "PUT", "DELETE", "PATCH"] + wildcard_path = f"{path}/{{subpath:path}}" - route_key = f"{endpoint_id}:subpath:{path}" + methods_str = ",".join(sorted(methods)) + route_key = f"{endpoint_id}:subpath:{path}:{methods_str}" # Check if this subpath route is already registered if route_key in _registered_pass_through_routes: verbose_proxy_logger.debug( - "Updating duplicate wildcard pass through endpoint: %s (already registered)", + "Updating duplicate wildcard pass through endpoint: %s with methods %s (already registered)", wildcard_path, + methods, ) verbose_proxy_logger.debug( - "adding wildcard pass through endpoint: %s, dependencies: %s", + "adding wildcard pass through endpoint: %s, methods: %s, dependencies: %s", wildcard_path, + methods, dependencies, ) @@ -1944,9 +1978,10 @@ class InitPassThroughEndpointHelpers: dependencies, include_subpath=True, cost_per_request=cost_per_request, + default_query_params=default_query_params, guardrails=guardrails, ), - methods=["GET", "POST", "PUT", "DELETE", "PATCH"], + methods=methods, dependencies=dependencies, ) @@ -1955,11 +1990,13 @@ class InitPassThroughEndpointHelpers: "endpoint_id": endpoint_id, "path": path, "type": "subpath", + "methods": methods, "passthrough_params": { "target": target, "custom_headers": custom_headers, "forward_headers": forward_headers, "merge_query_params": merge_query_params, + "default_query_params": default_query_params, "dependencies": dependencies, "cost_per_request": cost_per_request, "guardrails": guardrails, @@ -2026,11 +2063,12 @@ class InitPassThroughEndpointHelpers: return True # Fast path: check if any registered route key contains this path - # Keys are in format: "{endpoint_id}:exact:{path}" or "{endpoint_id}:subpath:{path}" + # Keys are in format: "{endpoint_id}:exact:{path}:{methods}" or "{endpoint_id}:subpath:{path}:{methods}" + # For backward compatibility, also support old format: "{endpoint_id}:exact:{path}" or "{endpoint_id}:subpath:{path}" # Extract unique paths from keys for quick checking for key in _registered_pass_through_routes.keys(): - parts = key.split(":", 2) # Split into [endpoint_id, type, path] - if len(parts) == 3: + parts = key.split(":", 3) # Split into [endpoint_id, type, path, methods?] + if len(parts) >= 3: route_type = parts[1] registered_path = ( InitPassThroughEndpointHelpers._build_full_path_with_root(parts[2]) @@ -2046,21 +2084,37 @@ class InitPassThroughEndpointHelpers: return False @staticmethod - def get_registered_pass_through_route(route: str) -> Optional[Dict[str, Any]]: - """Get passthrough params for a given route""" + def get_registered_pass_through_route( + route: str, method: Optional[str] = None + ) -> Optional[Dict[str, Any]]: + """Get passthrough params for a given route and optionally filter by HTTP method""" for key in _registered_pass_through_routes.keys(): - parts = key.split(":", 2) # Split into [endpoint_id, type, path] - if len(parts) == 3: + parts = key.split(":", 3) # Split into [endpoint_id, type, path, methods?] + if len(parts) >= 3: route_type = parts[1] registered_path = ( InitPassThroughEndpointHelpers._build_full_path_with_root(parts[2]) ) + # Get the methods for this route + route_methods = _registered_pass_through_routes[key].get("methods", []) + + # Check if path matches + path_matches = False if route_type == "exact" and route == registered_path: - return _registered_pass_through_routes[key] + path_matches = True elif route_type == "subpath": if route == registered_path or route.startswith( registered_path + "/" + ): + path_matches = True + + # If path matches and method filter is provided, check if method is allowed + if path_matches: + if ( + method is None + or not route_methods + or method in route_methods ): return _registered_pass_through_routes[key] @@ -2144,6 +2198,7 @@ async def initialize_pass_through_endpoints( ) _forward_headers = endpoint.get("forward_headers", None) _merge_query_params = endpoint.get("merge_query_params", None) + _default_query_params = endpoint.get("default_query_params", None) _auth = endpoint.get("auth", None) _dependencies = None if _auth is not None and str(_auth).lower() == "true": @@ -2162,6 +2217,9 @@ async def initialize_pass_through_endpoints( # Get guardrails config if present _guardrails = endpoint.get("guardrails", None) + # Get methods list if present (None means all methods for backward compatibility) + _methods = endpoint.get("methods", None) + # Add exact path route verbose_proxy_logger.debug( "Initializing pass through endpoint: %s (ID: %s)", _path, endpoint_id @@ -2177,9 +2235,16 @@ async def initialize_pass_through_endpoints( cost_per_request=endpoint.get("cost_per_request", None), endpoint_id=endpoint_id, guardrails=_guardrails, + methods=_methods, + default_query_params=_default_query_params, ) - visited_endpoints.add(f"{endpoint_id}:exact:{_path}") + # Generate route key with methods for tracking + methods_for_key = ( + _methods if _methods else ["GET", "POST", "PUT", "DELETE", "PATCH"] + ) + methods_str = ",".join(sorted(methods_for_key)) + visited_endpoints.add(f"{endpoint_id}:exact:{_path}:{methods_str}") # Add wildcard route for sub-paths if endpoint.get("include_subpath", False) is True: @@ -2194,9 +2259,11 @@ async def initialize_pass_through_endpoints( cost_per_request=endpoint.get("cost_per_request", None), endpoint_id=endpoint_id, guardrails=_guardrails, + methods=_methods, + default_query_params=_default_query_params, ) - visited_endpoints.add(f"{endpoint_id}:subpath:{_path}") + visited_endpoints.add(f"{endpoint_id}:subpath:{_path}:{methods_str}") verbose_proxy_logger.debug( "Added new pass through endpoint: %s (ID: %s)", _path, endpoint_id @@ -2509,6 +2576,8 @@ async def update_pass_through_endpoints( cost_per_request=updated_endpoint.cost_per_request, endpoint_id=updated_endpoint.id or endpoint_id or "", guardrails=getattr(updated_endpoint, "guardrails", None), + methods=updated_endpoint.methods, + default_query_params=updated_endpoint.default_query_params, ) else: InitPassThroughEndpointHelpers.add_exact_path_route( @@ -2522,6 +2591,8 @@ async def update_pass_through_endpoints( cost_per_request=updated_endpoint.cost_per_request, endpoint_id=updated_endpoint.id or endpoint_id or "", guardrails=getattr(updated_endpoint, "guardrails", None), + methods=updated_endpoint.methods, + default_query_params=updated_endpoint.default_query_params, ) return PassThroughEndpointResponse( @@ -2598,6 +2669,8 @@ async def create_pass_through_endpoints( cost_per_request=created_endpoint.cost_per_request, endpoint_id=created_endpoint.id or "", guardrails=getattr(created_endpoint, "guardrails", None), + methods=created_endpoint.methods, + default_query_params=created_endpoint.default_query_params, ) else: InitPassThroughEndpointHelpers.add_exact_path_route( @@ -2611,6 +2684,8 @@ async def create_pass_through_endpoints( cost_per_request=created_endpoint.cost_per_request, endpoint_id=created_endpoint.id or "", guardrails=getattr(created_endpoint, "guardrails", None), + methods=created_endpoint.methods, + default_query_params=created_endpoint.default_query_params, ) return PassThroughEndpointResponse(endpoints=[created_endpoint]) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_method_specific_routing.py b/tests/test_litellm/proxy/pass_through_endpoints/test_method_specific_routing.py new file mode 100644 index 00000000000..7e5b9ff6403 --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_method_specific_routing.py @@ -0,0 +1,155 @@ +""" +Test method-specific routing for pass-through endpoints. + +This test demonstrates the ability to configure different targets +for the same path but different HTTP methods. +""" + +import pytest + +from litellm.proxy._types import PassThroughGenericEndpoint + + +def test_pass_through_endpoint_with_methods(): + """Test creating pass-through endpoints with specific methods""" + + # Create endpoint for GET /azure/kb + get_endpoint = PassThroughGenericEndpoint( + id="get-azure-kb", + path="/azure/kb", + target="https://api1.example.com/knowledge-base", + methods=["GET"], + headers={"Authorization": "Bearer token1"}, + ) + + assert get_endpoint.path == "/azure/kb" + assert get_endpoint.methods == ["GET"] + assert get_endpoint.target == "https://api1.example.com/knowledge-base" + + # Create endpoint for POST /azure/kb + post_endpoint = PassThroughGenericEndpoint( + id="post-azure-kb", + path="/azure/kb", + target="https://api2.example.com/knowledge-base", + methods=["POST"], + headers={"Authorization": "Bearer token2"}, + ) + + assert post_endpoint.path == "/azure/kb" + assert post_endpoint.methods == ["POST"] + assert post_endpoint.target == "https://api2.example.com/knowledge-base" + + # These should be different endpoints despite same path + assert get_endpoint.id != post_endpoint.id + assert get_endpoint.target != post_endpoint.target + + +def test_pass_through_endpoint_multiple_methods(): + """Test creating endpoint with multiple methods""" + + endpoint = PassThroughGenericEndpoint( + id="multi-method", + path="/azure/kb", + target="https://api.example.com/kb", + methods=["GET", "POST", "PUT"], + headers={}, + ) + + assert len(endpoint.methods) == 3 + assert "GET" in endpoint.methods + assert "POST" in endpoint.methods + assert "PUT" in endpoint.methods + + +def test_pass_through_endpoint_no_methods_backward_compatibility(): + """Test that endpoints without methods field work (backward compatibility)""" + + # When methods is None, all methods should be supported + endpoint = PassThroughGenericEndpoint( + id="all-methods", + path="/azure/kb", + target="https://api.example.com/kb", + headers={}, + ) + + assert endpoint.methods is None # Default is None for backward compatibility + + +def test_pass_through_endpoint_serialization(): + """Test that endpoints with methods can be serialized/deserialized""" + + endpoint = PassThroughGenericEndpoint( + id="test-endpoint", + path="/test", + target="https://api.example.com", + methods=["GET", "POST"], + headers={"key": "value"}, + cost_per_request=0.5, + ) + + # Serialize to dict + endpoint_dict = endpoint.model_dump() + assert endpoint_dict["methods"] == ["GET", "POST"] + + # Deserialize from dict + restored_endpoint = PassThroughGenericEndpoint(**endpoint_dict) + assert restored_endpoint.methods == ["GET", "POST"] + assert restored_endpoint.path == "/test" + assert restored_endpoint.target == "https://api.example.com" + + +def test_route_key_generation_with_methods(): + """Test that route keys include methods for uniqueness""" + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + InitPassThroughEndpointHelpers, + ) + + # Simulate how route keys are generated + endpoint_id_1 = "endpoint-1" + path = "/azure/kb" + methods_1 = ["GET"] + methods_str_1 = ",".join(sorted(methods_1)) + route_key_1 = f"{endpoint_id_1}:exact:{path}:{methods_str_1}" + + endpoint_id_2 = "endpoint-2" + methods_2 = ["POST"] + methods_str_2 = ",".join(sorted(methods_2)) + route_key_2 = f"{endpoint_id_2}:exact:{path}:{methods_str_2}" + + # Keys should be different even though path is the same + assert route_key_1 != route_key_2 + assert route_key_1 == "endpoint-1:exact:/azure/kb:GET" + assert route_key_2 == "endpoint-2:exact:/azure/kb:POST" + + +def test_config_yaml_example(): + """ + Example configuration for config.yaml showing method-specific routing: + + general_settings: + pass_through_endpoints: + # GET endpoint for retrieving knowledge base + - id: "get-azure-kb" + path: "/azure/kb" + target: "https://read-api.example.com/kb" + methods: ["GET"] + headers: + Authorization: "bearer os.environ/READ_API_KEY" + + # POST endpoint for creating knowledge base entries + - id: "post-azure-kb" + path: "/azure/kb" + target: "https://write-api.example.com/kb" + methods: ["POST"] + headers: + Authorization: "bearer os.environ/WRITE_API_KEY" + + # PUT endpoint for updating knowledge base + - id: "put-azure-kb" + path: "/azure/kb" + target: "https://update-api.example.com/kb" + methods: ["PUT"] + headers: + Authorization: "bearer os.environ/UPDATE_API_KEY" + """ + pass diff --git a/ui/litellm-dashboard/src/components/add_pass_through.tsx b/ui/litellm-dashboard/src/components/add_pass_through.tsx index ff724e42eb8..af55a7919a9 100644 --- a/ui/litellm-dashboard/src/components/add_pass_through.tsx +++ b/ui/litellm-dashboard/src/components/add_pass_through.tsx @@ -23,6 +23,7 @@ import { ApiOutlined, } from "@ant-design/icons"; import KeyValueInput from "./key_value_input"; +import QueryParamInput from "./query_param_input"; import { passThroughItem } from "./pass_through_settings"; import RoutePreview from "./route_preview"; import NotificationsManager from "./molecules/notifications_manager"; @@ -30,6 +31,8 @@ import PassThroughSecuritySection from "./common_components/PassThroughSecurityS import PassThroughGuardrailsSection from "./common_components/PassThroughGuardrailsSection"; const { Option } = Select2; +const HTTP_METHODS = ["GET", "POST", "PUT", "DELETE", "PATCH"]; + interface AddFallbacksProps { // models: string[] | undefined; accessToken: string; @@ -52,12 +55,14 @@ const AddPassThroughEndpoint: React.FC = ({ const [targetValue, setTargetValue] = useState(""); const [includeSubpath, setIncludeSubpath] = useState(true); const [authEnabled, setAuthEnabled] = useState(false); + const [selectedMethods, setSelectedMethods] = useState([]); const [guardrails, setGuardrails] = useState>({}); const handleCancel = () => { form.resetFields(); setPathValue(""); setTargetValue(""); setIncludeSubpath(true); + setSelectedMethods([]); setGuardrails({}); setIsModalVisible(false); }; @@ -86,6 +91,11 @@ const AddPassThroughEndpoint: React.FC = ({ formValues.guardrails = guardrails; } + // Add methods to formValues (only if specific methods are selected) + if (selectedMethods && selectedMethods.length > 0) { + formValues.methods = selectedMethods; + } + console.log(`formValues: ${JSON.stringify(formValues)}`); const response = await createPassThroughEndpoint(accessToken, formValues); @@ -101,6 +111,7 @@ const AddPassThroughEndpoint: React.FC = ({ setPathValue(""); setTargetValue(""); setIncludeSubpath(true); + setSelectedMethods([]); setGuardrails({}); setIsModalVisible(false); } catch (error) { @@ -204,6 +215,41 @@ const AddPassThroughEndpoint: React.FC = ({ /> + + HTTP Methods (Optional) + + + + + } + name="methods" + extra={ +
+ {selectedMethods.length === 0 + ? "All HTTP methods supported (default)" + : `Only ${selectedMethods.join(", ")} requests will be routed to this endpoint`} +
+ } + className="mb-4" + > + + {HTTP_METHODS.map((method) => ( + + ))} + +
+
Include Subpaths
@@ -250,6 +296,34 @@ const AddPassThroughEndpoint: React.FC = ({ + {/* Default Query Parameters Section */} + + Default Query Parameters + + Add query parameters that will be automatically sent with every request to the target API + + + + Default Query Parameters (Optional) + + + + + } + name="default_query_params" + extra={ +
+
Parameters are sent with all GET, POST, PUT, PATCH requests
+
Client parameters override defaults. Examples: version=v1, format=json, key=default
+
+ } + > + +
+
+ {/* Security Section */} = ({ onChange, value, @@ -21,7 +26,7 @@ const PassThroughRoutesSelector: React.FC = ({ disabled = false, teamId, }) => { - const [passThroughRoutes, setPassThroughRoutes] = useState([]); + const [passThroughRoutes, setPassThroughRoutes] = useState>([]); const [loading, setLoading] = useState(false); useEffect(() => { @@ -32,7 +37,24 @@ const PassThroughRoutesSelector: React.FC = ({ try { const response = await getPassThroughEndpointsCall(accessToken, teamId); if (response.endpoints) { - const routes = response.endpoints.map((route: { path: string }) => route.path); + const routes = response.endpoints.flatMap((endpoint: PassThroughEndpoint) => { + const path = endpoint.path; + const methods = endpoint.methods; + + // If methods are specified, create one entry per method + if (methods && methods.length > 0) { + return methods.map((method) => ({ + label: `${method} ${path}`, + value: path, // Keep value as path for backward compatibility + })); + } + + // If no methods specified, show just the path (all methods supported) + return [{ + label: path, + value: path, + }]; + }); setPassThroughRoutes(routes); } } catch (error) { @@ -54,10 +76,7 @@ const PassThroughRoutesSelector: React.FC = ({ loading={loading} className={className} allowClear - options={passThroughRoutes.map((route) => ({ - label: route, - value: route, - }))} + options={passThroughRoutes} optionFilterProp="label" showSearch style={{ width: "100%" }} diff --git a/ui/litellm-dashboard/src/components/pass_through_info.tsx b/ui/litellm-dashboard/src/components/pass_through_info.tsx index fe3b204b507..9a8d2a3a07f 100644 --- a/ui/litellm-dashboard/src/components/pass_through_info.tsx +++ b/ui/litellm-dashboard/src/components/pass_through_info.tsx @@ -13,7 +13,7 @@ import { TabPanels, TextInput, } from "@tremor/react"; -import { Button, Form, Input, Switch, InputNumber } from "antd"; +import { Button, Form, Input, Switch, InputNumber, Select } from "antd"; import { updatePassThroughEndpoint, deletePassThroughEndpointsCall } from "./networking"; import { Eye, EyeOff } from "lucide-react"; import RoutePreview from "./route_preview"; @@ -21,6 +21,9 @@ import NotificationsManager from "./molecules/notifications_manager"; import PassThroughSecuritySection from "./common_components/PassThroughSecuritySection"; import PassThroughGuardrailsSection from "./common_components/PassThroughGuardrailsSection"; +const HTTP_METHODS = ["GET", "POST", "PUT", "DELETE", "PATCH"]; +const { Option } = Select; + export interface PassThroughInfoProps { endpointData: PassThroughEndpoint; onClose: () => void; @@ -38,6 +41,7 @@ interface PassThroughEndpoint { include_subpath?: boolean; cost_per_request?: number; auth?: boolean; + methods?: string[]; guardrails?: Record; } @@ -70,6 +74,7 @@ const PassThroughInfoView: React.FC = ({ const [loading, setLoading] = useState(false); const [isEditing, setIsEditing] = useState(false); const [authEnabled, setAuthEnabled] = useState(initialEndpointData?.auth || false); + const [selectedMethods, setSelectedMethods] = useState(initialEndpointData?.methods || []); const [guardrails, setGuardrails] = useState>( initialEndpointData?.guardrails || {} ); @@ -97,6 +102,7 @@ const PassThroughInfoView: React.FC = ({ include_subpath: values.include_subpath, cost_per_request: values.cost_per_request, auth: premiumUser ? values.auth : undefined, + methods: selectedMethods && selectedMethods.length > 0 ? selectedMethods : undefined, guardrails: guardrails && Object.keys(guardrails).length > 0 ? guardrails : undefined, }; @@ -191,6 +197,23 @@ const PassThroughInfoView: React.FC = ({ {endpointData.auth ? "Auth Required" : "No Auth"}
+ {endpointData.methods && endpointData.methods.length > 0 && ( +
+ HTTP Methods: +
+ {endpointData.methods.map((method) => ( + + {method} + + ))} +
+
+ )} + {(!endpointData.methods || endpointData.methods.length === 0) && ( +
+ All HTTP methods supported +
+ )} {endpointData.cost_per_request !== undefined && (
Cost per request: ${endpointData.cost_per_request} @@ -277,6 +300,7 @@ const PassThroughInfoView: React.FC = ({ include_subpath: endpointData.include_subpath || false, cost_per_request: endpointData.cost_per_request, auth: endpointData.auth || false, + methods: endpointData.methods || [], }} layout="vertical" > @@ -295,6 +319,31 @@ const PassThroughInfoView: React.FC = ({ /> + + + + diff --git a/ui/litellm-dashboard/src/components/pass_through_settings.tsx b/ui/litellm-dashboard/src/components/pass_through_settings.tsx index b37d3df7e41..d60caa54667 100644 --- a/ui/litellm-dashboard/src/components/pass_through_settings.tsx +++ b/ui/litellm-dashboard/src/components/pass_through_settings.tsx @@ -39,7 +39,9 @@ export interface passThroughItem { include_subpath?: boolean; cost_per_request?: number; auth?: boolean; + methods?: string[]; guardrails?: Record; + default_query_params?: Record; } // Password field component for headers @@ -147,6 +149,32 @@ const PassThroughSettings: React.FC = ({ accessToken, accessorKey: "target", cell: (info: any) => {info.getValue()}, }, + { + header: () => ( +
+ Methods + + + +
+ ), + accessorKey: "methods", + cell: (info: any) => { + const methods = info.getValue(); + if (!methods || methods.length === 0) { + return ALL; + } + return ( +
+ {methods.map((method: string) => ( + + {method} + + ))} +
+ ); + }, + }, { header: () => (
diff --git a/ui/litellm-dashboard/src/components/query_param_input.tsx b/ui/litellm-dashboard/src/components/query_param_input.tsx new file mode 100644 index 00000000000..4a98ab48bca --- /dev/null +++ b/ui/litellm-dashboard/src/components/query_param_input.tsx @@ -0,0 +1,57 @@ +import React, { useState } from "react"; +import { Button, Space } from "antd"; +import { MinusCircleOutlined, PlusOutlined } from "@ant-design/icons"; +import { TextInput } from "@tremor/react"; + +interface QueryParamInputProps { + value?: Record; + onChange?: (value: Record) => void; +} + +const QueryParamInput: React.FC = ({ value = {}, onChange }) => { + const [pairs, setPairs] = useState<[string, string][]>(Object.entries(value)); + + const handleAdd = () => { + setPairs([...pairs, ["", ""]]); + }; + + const handleRemove = (index: number) => { + const newPairs = pairs.filter((_, i) => i !== index); + setPairs(newPairs); + onChange?.(Object.fromEntries(newPairs)); + }; + + const handleChange = (index: number, key: string, val: string) => { + const newPairs = [...pairs]; + newPairs[index] = [key, val]; + setPairs(newPairs); + onChange?.(Object.fromEntries(newPairs)); + }; + + return ( +
+ {pairs.map(([key, val], index) => ( + + handleChange(index, e.target.value, val)} + /> + handleChange(index, key, e.target.value)} + /> +
+ handleRemove(index)} style={{ cursor: "pointer" }} /> +
+
+ ))} + +
+ ); +}; + +export default QueryParamInput; \ No newline at end of file