diff --git a/litellm/proxy/config_resolvers/sso.py b/litellm/proxy/config_resolvers/sso.py index c61f5b6724d..5e6fd0ed5b0 100644 --- a/litellm/proxy/config_resolvers/sso.py +++ b/litellm/proxy/config_resolvers/sso.py @@ -37,6 +37,7 @@ SSO_DESCRIPTORS: Final[tuple[FieldDescriptor, ...]] = ( FieldDescriptor("generic_token_endpoint", "generic_token_endpoint", "GENERIC_TOKEN_ENDPOINT"), FieldDescriptor("generic_userinfo_endpoint", "generic_userinfo_endpoint", "GENERIC_USERINFO_ENDPOINT"), FieldDescriptor("generic_scope", "generic_scope", "GENERIC_SCOPE", default="openid email profile"), + FieldDescriptor("generic_authorization_params", "generic_authorization_params", "GENERIC_AUTHORIZATION_PARAMS"), FieldDescriptor("saml_idp_metadata_url", "saml_idp_metadata_url", "SAML_IDP_METADATA_URL"), FieldDescriptor("saml_idp_metadata_xml", "saml_idp_metadata_xml", "SAML_IDP_METADATA_XML"), FieldDescriptor("saml_sp_entity_id", "saml_sp_entity_id", "SAML_SP_ENTITY_ID"), diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 01807fefd78..e88f608c54b 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -33,7 +33,7 @@ from typing import ( cast, overload, ) -from urllib.parse import parse_qs, urlencode, urlparse +from urllib.parse import parse_qs, parse_qsl, urlencode, urlparse if TYPE_CHECKING: import httpx @@ -1420,6 +1420,21 @@ def _parse_generic_sso_headers() -> dict[str, str]: return result +_GENERIC_SSO_RESERVED_AUTHORIZE_PARAMS: Final = frozenset( + {"client_id", "redirect_uri", "response_type", "scope", "state", "code_challenge", "code_challenge_method"} +) + + +def _parse_generic_sso_authorization_params() -> dict[str, str]: + """Parse the query-string GENERIC_AUTHORIZATION_PARAMS env var (``resource=https://api.example.com``) into + extra parameters for the authorization request. Keys the OAuth flow sets itself are ignored.""" + raw: Final = os.getenv("GENERIC_AUTHORIZATION_PARAMS", "") + pairs: Final = parse_qsl(raw.strip(), keep_blank_values=False) + for key in sorted({key for key, _ in pairs if key in _GENERIC_SSO_RESERVED_AUTHORIZE_PARAMS}): + verbose_proxy_logger.warning("Ignoring GENERIC_AUTHORIZATION_PARAMS key %s, the SSO flow sets it", key) + return {key: value for key, value in pairs if key not in _GENERIC_SSO_RESERVED_AUTHORIZE_PARAMS} + + def _handle_generic_sso_error( e: Exception, generic_authorization_endpoint: str | None, @@ -3109,7 +3124,9 @@ class SSOAuthenticationHandler: state_only_params[key] = value # Get the redirect response from fastapi-sso with only state param - redirect_response: Final = await generic_sso.get_login_redirect(**state_only_params) + redirect_response: Final = await generic_sso.get_login_redirect( + params=_parse_generic_sso_authorization_params(), **state_only_params + ) # If PKCE is enabled, add PKCE parameters to the redirect URL if code_verifier and "state" in redirect_params: diff --git a/litellm/types/proxy/management_endpoints/ui_sso.py b/litellm/types/proxy/management_endpoints/ui_sso.py index b4691b9b08c..f20cb01f641 100644 --- a/litellm/types/proxy/management_endpoints/ui_sso.py +++ b/litellm/types/proxy/management_endpoints/ui_sso.py @@ -152,6 +152,13 @@ class SSOConfig(LiteLLMPydanticObjectBase): default=None, description="Space-separated OAuth scopes requested from the generic provider, e.g. 'openid email profile'", ) + generic_authorization_params: str | None = Field( + default=None, + description=( + "Extra query parameters for the generic provider's authorization request, in query-string form, " + "e.g. 'resource=https://litellm.example.com' so AD FS issues tokens for that Web API" + ), + ) # SAML SSO saml_idp_metadata_url: str | None = Field( diff --git a/tests/unit/proxy/config_resolvers/test_config_resolvers.py b/tests/unit/proxy/config_resolvers/test_config_resolvers.py index 9f91d9ee2c9..5d19e909277 100644 --- a/tests/unit/proxy/config_resolvers/test_config_resolvers.py +++ b/tests/unit/proxy/config_resolvers/test_config_resolvers.py @@ -58,6 +58,7 @@ def test_sso_descriptor_mapping_is_single_sourced(): # The write path and read path both consume this mapping; it must cover every # env-backed SSO field and map to the uppercase env var. assert SSO_FIELD_ENV_VARS["generic_client_id"] == "GENERIC_CLIENT_ID" + assert SSO_FIELD_ENV_VARS["generic_authorization_params"] == "GENERIC_AUTHORIZATION_PARAMS" assert SSO_SECRET_FIELDS == frozenset( {"google_client_secret", "microsoft_client_secret", "generic_client_secret"} ) diff --git a/tests/unit/proxy/management_endpoints/test_ui_sso.py b/tests/unit/proxy/management_endpoints/test_ui_sso.py index 8ff0b24982f..c0737002b2d 100644 --- a/tests/unit/proxy/management_endpoints/test_ui_sso.py +++ b/tests/unit/proxy/management_endpoints/test_ui_sso.py @@ -5,6 +5,7 @@ import os from contextlib import ExitStack, asynccontextmanager from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch +from urllib.parse import parse_qs, urlparse import httpx import pytest @@ -4730,6 +4731,132 @@ class TestGenericResponseConvertorUserRole: assert result.user_role is None +class TestGenericSSOAuthorizationParams: + """GENERIC_AUTHORIZATION_PARAMS adds extra query parameters (AD FS ``resource``) to the authorize redirect.""" + + @staticmethod + def _generic_sso(scope: list[str]) -> object: + from fastapi_sso.sso.generic import create_provider + + provider = create_provider( + name="oidc", + discovery_document={ + "authorization_endpoint": "https://idp.example.com/adfs/oauth2/authorize", + "token_endpoint": "https://idp.example.com/adfs/oauth2/token", + "userinfo_endpoint": "https://idp.example.com/adfs/userinfo", + }, + ) + return provider( + client_id="litellm-client", + client_secret="secret", + redirect_uri="http://localhost:4000/sso/callback", + allow_insecure_http=True, + scope=scope, + ) + + def test_parse_is_empty_when_unset(self, monkeypatch): + from litellm.proxy.management_endpoints.ui_sso import ( + _parse_generic_sso_authorization_params, + ) + + monkeypatch.delenv("GENERIC_AUTHORIZATION_PARAMS", raising=False) + assert _parse_generic_sso_authorization_params() == {} + monkeypatch.setenv("GENERIC_AUTHORIZATION_PARAMS", " ") + assert _parse_generic_sso_authorization_params() == {} + + def test_parse_reads_query_string_form(self, monkeypatch): + from litellm.proxy.management_endpoints.ui_sso import ( + _parse_generic_sso_authorization_params, + ) + + monkeypatch.setenv( + "GENERIC_AUTHORIZATION_PARAMS", + "resource=https://litellm.example.com/api&prompt=login", + ) + assert _parse_generic_sso_authorization_params() == { + "resource": "https://litellm.example.com/api", + "prompt": "login", + } + + def test_parse_drops_keys_the_flow_sets_itself(self, monkeypatch): + from litellm.proxy.management_endpoints.ui_sso import ( + _parse_generic_sso_authorization_params, + ) + + monkeypatch.setenv( + "GENERIC_AUTHORIZATION_PARAMS", + "resource=https://litellm.example.com/api&state=attacker&redirect_uri=https://evil.example" + "&client_id=other&scope=openid&response_type=token&code_challenge=x&code_challenge_method=plain", + ) + assert _parse_generic_sso_authorization_params() == {"resource": "https://litellm.example.com/api"} + + @pytest.mark.asyncio + async def test_redirect_location_carries_resource(self, monkeypatch): + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + monkeypatch.delenv("GENERIC_CLIENT_USE_PKCE", raising=False) + monkeypatch.setenv("GENERIC_AUTHORIZATION_PARAMS", "resource=https://litellm.example.com/api") + response = await SSOAuthenticationHandler.get_generic_sso_redirect_response( + generic_sso=self._generic_sso(["openid", "profile", "email", "allatclaims"]), + state="cli-state-123", + generic_authorization_endpoint="https://idp.example.com/adfs/oauth2/authorize", + ) + + assert response is not None + location = urlparse(str(response.headers["location"])) + query = parse_qs(location.query) + assert location.path == "/adfs/oauth2/authorize" + assert query["resource"] == ["https://litellm.example.com/api"] + assert query["state"] == ["cli-state-123"] + assert query["client_id"] == ["litellm-client"] + assert query["redirect_uri"] == ["http://localhost:4000/sso/callback"] + assert query["scope"] == ["openid profile email allatclaims"] + assert query["response_type"] == ["code"] + + @pytest.mark.asyncio + async def test_redirect_location_unchanged_when_unset(self, monkeypatch): + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + monkeypatch.delenv("GENERIC_CLIENT_USE_PKCE", raising=False) + monkeypatch.delenv("GENERIC_AUTHORIZATION_PARAMS", raising=False) + response = await SSOAuthenticationHandler.get_generic_sso_redirect_response( + generic_sso=self._generic_sso(["openid"]), + state="cli-state-123", + generic_authorization_endpoint="https://idp.example.com/adfs/oauth2/authorize", + ) + + assert response is not None + query = parse_qs(urlparse(str(response.headers["location"])).query) + assert "resource" not in query + assert set(query) == {"response_type", "client_id", "redirect_uri", "scope", "state"} + + @pytest.mark.asyncio + async def test_redirect_location_keeps_pkce_alongside_resource(self, monkeypatch): + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + monkeypatch.setenv("GENERIC_CLIENT_USE_PKCE", "true") + monkeypatch.setenv("GENERIC_AUTHORIZATION_PARAMS", "resource=https://litellm.example.com/api") + mock_cache = MagicMock(redis_cache=None) + mock_cache.async_set_cache = AsyncMock() + with ( + patch("litellm.proxy.proxy_server.redis_usage_cache", None), + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + ): + response = await SSOAuthenticationHandler.get_generic_sso_redirect_response( + generic_sso=self._generic_sso(["openid"]), + state="cli-state-123", + generic_authorization_endpoint="https://idp.example.com/adfs/oauth2/authorize", + ) + + assert response is not None + query = parse_qs(urlparse(str(response.headers["location"])).query) + assert query["resource"] == ["https://litellm.example.com/api"] + assert query["code_challenge_method"] == ["S256"] + assert len(query["code_challenge"]) == 1 + assert query["state"] == ["cli-state-123"] + mock_cache.async_set_cache.assert_called_once() + + class TestGetGenericSSORedirectParams: """Test _get_generic_sso_redirect_params state parameter priority handling""" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts index 83847261fe8..0082ac9e330 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts @@ -29,6 +29,7 @@ export interface SSOSettingsValues { saml_sp_entity_id: string | null; saml_allow_unsolicited: string | null; generic_scope: string | null; + generic_authorization_params: string | null; proxy_base_url: string | null; user_email: string | null; ui_access_mode: string | null; diff --git a/ui/litellm-dashboard/src/components/SSOModals.test.tsx b/ui/litellm-dashboard/src/components/SSOModals.test.tsx index bee5f13c626..ee8da4c8df1 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.test.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.test.tsx @@ -531,6 +531,7 @@ describe("SSOModals", () => { saml_sp_entity_id: null, saml_allow_unsolicited: null, generic_scope: null, + generic_authorization_params: null, proxy_base_url: null, user_email: null, sso_provider: null, diff --git a/ui/litellm-dashboard/src/components/SSOModals.tsx b/ui/litellm-dashboard/src/components/SSOModals.tsx index 609c0efcd99..b5a702b8ad1 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.tsx @@ -108,6 +108,7 @@ const SSOModals: React.FC = ({ generic_token_endpoint: ssoData.values.generic_token_endpoint, generic_userinfo_endpoint: ssoData.values.generic_userinfo_endpoint, generic_scope: ssoData.values.generic_scope, + generic_authorization_params: ssoData.values.generic_authorization_params, saml_idp_metadata_url: ssoData.values.saml_idp_metadata_url, saml_idp_metadata_xml: ssoData.values.saml_idp_metadata_xml, saml_sp_entity_id: ssoData.values.saml_sp_entity_id, @@ -220,6 +221,7 @@ const SSOModals: React.FC = ({ saml_sp_entity_id: null, saml_allow_unsolicited: null, generic_scope: null, + generic_authorization_params: null, proxy_base_url: null, user_email: null, sso_provider: null, diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx index eddc5716305..46c6fc143dc 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx @@ -289,13 +289,13 @@ describe("renderProviderFields", () => { it("should return fields for okta provider", () => { const result = renderProviderFields("okta"); expect(result).not.toBeNull(); - expect(result?.length).toBe(6); + expect(result?.length).toBe(7); }); it("should return fields for generic provider", () => { const result = renderProviderFields("generic"); expect(result).not.toBeNull(); - expect(result?.length).toBe(6); + expect(result?.length).toBe(7); }); it.each(["okta", "generic"])( @@ -308,6 +308,16 @@ describe("renderProviderFields", () => { }, ); + it.each(["okta", "generic"])( + "renders an optional generic_authorization_params field for %s mapped to GENERIC_AUTHORIZATION_PARAMS", + (provider) => { + const field = ssoProviderConfigs[provider].fields.find((f) => f.name === "generic_authorization_params"); + expect(field).toBeDefined(); + expect(field?.required).toBe(false); + expect(ssoProviderConfigs[provider].envVarMap.generic_authorization_params).toBe("GENERIC_AUTHORIZATION_PARAMS"); + }, + ); + it("submits generic_scope untouched, so saving an unrelated edit cannot clear GENERIC_SCOPE", async () => { // update_sso_settings clears the env var for any mapped field its payload // omits, and antd only submits mounted fields. So the Scopes field being @@ -356,6 +366,48 @@ describe("renderProviderFields", () => { }); }); + it("submits generic_authorization_params untouched, so saving an unrelated edit cannot clear GENERIC_AUTHORIZATION_PARAMS", async () => { + const handleSubmit = vi.fn(); + let form!: ReturnType; + const TestWrapper = () => { + const formInstance = useSSOSettingsForm("sso-settings"); + form = formInstance; + return ; + }; + + renderWithProviders(); + + const savedValues = { + ...emptySSOSettingsFormValues, + sso_provider: "generic", + generic_client_id: "client-id", + generic_client_secret: "client-secret", + generic_authorization_endpoint: "https://idp.example.com/authorize", + generic_token_endpoint: "https://idp.example.com/token", + generic_userinfo_endpoint: "https://idp.example.com/userinfo", + generic_authorization_params: "resource=https://litellm.example.com/api", + proxy_base_url: "https://gateway.example.com", + user_email: "admin@example.com", + }; + await act(async () => { + form.reset(savedValues); + }); + + await act(async () => { + form.setValue("generic_token_endpoint", "https://idp.example.com/token/v2"); + submitMountedSSOValues(form, "sso-settings", handleSubmit)(); + }); + + await waitFor(() => { + expect(handleSubmit).toHaveBeenCalledWith( + expect.objectContaining({ + generic_token_endpoint: "https://idp.example.com/token/v2", + generic_authorization_params: "resource=https://litellm.example.com/api", + }), + ); + }); + }); + it("renders provider logos in the dropdown and falls back to a letter avatar on load error", async () => { const TestWrapper = () => { const form = useSSOSettingsForm("sso-settings"); diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx index e4725294e05..46c397c6893 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx @@ -27,6 +27,7 @@ export interface SSOSettingsFormValues { generic_token_endpoint?: string; generic_userinfo_endpoint?: string; generic_scope?: string; + generic_authorization_params?: string; saml_idp_metadata_url?: string; saml_idp_metadata_xml?: string; saml_sp_entity_id?: string; @@ -91,6 +92,7 @@ export const ssoProviderConfigs: Record = { generic_token_endpoint: "GENERIC_TOKEN_ENDPOINT", generic_userinfo_endpoint: "GENERIC_USERINFO_ENDPOINT", generic_scope: "GENERIC_SCOPE", + generic_authorization_params: "GENERIC_AUTHORIZATION_PARAMS", }, fields: [ { label: "Generic Client ID", name: "generic_client_id" }, @@ -107,6 +109,12 @@ export const ssoProviderConfigs: Record = { placeholder: "https://your-domain/userinfo", }, { label: "Scopes", name: "generic_scope", placeholder: "openid email profile", required: false }, + { + label: "Extra Authorization Params", + name: "generic_authorization_params", + placeholder: "resource=https://your-api-identifier", + required: false, + }, ], }, generic: { @@ -117,6 +125,7 @@ export const ssoProviderConfigs: Record = { generic_token_endpoint: "GENERIC_TOKEN_ENDPOINT", generic_userinfo_endpoint: "GENERIC_USERINFO_ENDPOINT", generic_scope: "GENERIC_SCOPE", + generic_authorization_params: "GENERIC_AUTHORIZATION_PARAMS", }, fields: [ { label: "Generic Client ID", name: "generic_client_id" }, @@ -125,6 +134,12 @@ export const ssoProviderConfigs: Record = { { label: "Token Endpoint", name: "generic_token_endpoint" }, { label: "Userinfo Endpoint", name: "generic_userinfo_endpoint" }, { label: "Scopes", name: "generic_scope", placeholder: "openid email profile", required: false }, + { + label: "Extra Authorization Params", + name: "generic_authorization_params", + placeholder: "resource=https://your-api-identifier", + required: false, + }, ], }, saml: { diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.tsx index 58c03eda979..86b939c08f1 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.tsx @@ -45,6 +45,7 @@ export const toSSOFormValues = (values: SSOSettingsValues): SSOSettingsFormValue generic_token_endpoint: values.generic_token_endpoint ?? "", generic_userinfo_endpoint: values.generic_userinfo_endpoint ?? "", generic_scope: values.generic_scope ?? undefined, + generic_authorization_params: values.generic_authorization_params ?? undefined, saml_idp_metadata_url: values.saml_idp_metadata_url ?? undefined, saml_idp_metadata_xml: values.saml_idp_metadata_xml ?? undefined, saml_sp_entity_id: values.saml_sp_entity_id ?? undefined, diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx index 90d78c0ce72..eb9df343162 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx @@ -130,6 +130,10 @@ export default function SSOSettings() { render: (values: SSOSettingsValues) => , }, { label: "Scopes", render: (values: SSOSettingsValues) => renderSimpleValue(values.generic_scope) }, + { + label: "Extra Authorization Params", + render: (values: SSOSettingsValues) => renderSimpleValue(values.generic_authorization_params), + }, { label: "Proxy Base URL", render: (values: SSOSettingsValues) => renderSimpleValue(values.proxy_base_url) }, isTeamMappingsEnabled ? { label: "Team IDs JWT Field", render: (values: SSOSettingsValues) => renderTeamMappingsField(values) } @@ -160,6 +164,10 @@ export default function SSOSettings() { render: (values: SSOSettingsValues) => , }, { label: "Scopes", render: (values: SSOSettingsValues) => renderSimpleValue(values.generic_scope) }, + { + label: "Extra Authorization Params", + render: (values: SSOSettingsValues) => renderSimpleValue(values.generic_authorization_params), + }, { label: "Proxy Base URL", render: (values: SSOSettingsValues) => renderSimpleValue(values.proxy_base_url) }, isTeamMappingsEnabled ? { label: "Team IDs JWT Field", render: (values: SSOSettingsValues) => renderTeamMappingsField(values) } diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index ca6d3bdf269..a80f27a3b94 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -45147,6 +45147,11 @@ export interface components { * @description Authorization endpoint URL for generic OAuth provider */ generic_authorization_endpoint?: string | null; + /** + * Generic Authorization Params + * @description Extra query parameters for the generic provider's authorization request, in query-string form, e.g. 'resource=https://litellm.example.com' so AD FS issues tokens for that Web API + */ + generic_authorization_params?: string | null; /** * Generic Client Id * @description Generic OAuth Client ID for SSO authentication (used for Okta and other providers)