From 6b52abe6529e42ab0c4e899c63ac62c848f1af4a Mon Sep 17 00:00:00 2001 From: Hunter Wittenborn Date: Sat, 25 Apr 2026 07:08:48 -0500 Subject: [PATCH] fix oauth credential refresh and test coverage --- litellm/proxy/auth/user_api_key_auth.py | 33 ++++++++++------- .../proxy/credential_endpoints/endpoints.py | 35 ++++++++++++++++++- .../model_add/AddCredentialModal.test.tsx | 35 ++++++++++++++----- .../model_add/AddCredentialModal.tsx | 10 +++--- 4 files changed, 88 insertions(+), 25 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 9dd2ab18c8a..79af2e3e561 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -342,13 +342,14 @@ def get_rbac_role(jwt_handler: JWTHandler, scopes: List[str]) -> str: def get_api_key( custom_litellm_key_header: Optional[str], api_key: str, - azure_api_key_header: Optional[str], - anthropic_api_key_header: Optional[str], - google_ai_studio_api_key_header: Optional[str], - azure_apim_header: Optional[str], - pass_through_endpoints: Optional[List[dict]], - route: str, - request: Request, + azure_api_key_header: Optional[str] = None, + anthropic_api_key_header: Optional[str] = None, + google_ai_studio_api_key_header: Optional[str] = None, + azure_apim_header: Optional[str] = None, + pass_through_endpoints: Optional[List[dict]] = None, + route: str = "", + request: Optional[Request] = None, + AZURE_AI_API_KEY_header: Optional[str] = None, ) -> Tuple[str, Optional[str]]: """ Returns: @@ -360,6 +361,8 @@ def get_api_key( ) api_key = api_key + if azure_api_key_header is None and AZURE_AI_API_KEY_header is not None: + azure_api_key_header = AZURE_AI_API_KEY_header passed_in_key: Optional[str] = None if isinstance(custom_litellm_key_header, str): passed_in_key = custom_litellm_key_header @@ -519,12 +522,13 @@ async def _resolve_jwt_to_virtual_key( async def _user_api_key_auth_builder( # noqa: PLR0915 request: Request, api_key: str, - azure_api_key_header: str, - anthropic_api_key_header: Optional[str], - google_ai_studio_api_key_header: Optional[str], - azure_apim_header: Optional[str], - request_data: dict, + azure_api_key_header: Optional[str] = None, + anthropic_api_key_header: Optional[str] = None, + google_ai_studio_api_key_header: Optional[str] = None, + azure_apim_header: Optional[str] = None, + request_data: Optional[dict] = None, custom_litellm_key_header: Optional[str] = None, + AZURE_AI_API_KEY_header: Optional[str] = None, ) -> UserAPIKeyAuth: from litellm.proxy.proxy_server import ( general_settings, @@ -557,6 +561,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 pass_through_endpoints: Optional[List[dict]] = general_settings.get( "pass_through_endpoints", None ) + if azure_api_key_header is None and AZURE_AI_API_KEY_header is not None: + azure_api_key_header = AZURE_AI_API_KEY_header + if request_data is None: + request_data = {} + ## CHECK IF X-LITELM-API-KEY IS PASSED IN - supercedes Authorization header api_key, passed_in_key = get_api_key( custom_litellm_key_header=custom_litellm_key_header, diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 938987e0724..a75be374048 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -15,7 +15,10 @@ from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.litellm_logging import _get_masked_values from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_value_helper, + encrypt_value_helper, +) from litellm.proxy.utils import handle_exception_on_proxy, jsonify_object from litellm.types.utils import CreateCredentialItem, CredentialItem @@ -42,6 +45,23 @@ class CredentialHelperUtils: credential_info=credential.credential_info or {}, ) + @staticmethod + def decrypt_credential_values(credential: CredentialItem) -> CredentialItem: + """Decrypt values so in-memory credentials stay usable after DB updates.""" + decrypted_credential_values = {} + for key, value in (credential.credential_values or {}).items(): + decrypted_credential_values[key] = decrypt_value_helper( + value=value, + key=key, + return_original_value=True, + ) + + return CredentialItem( + credential_name=credential.credential_name, + credential_values=decrypted_credential_values, + credential_info=credential.credential_info or {}, + ) + @router.post( "/credentials", @@ -389,6 +409,19 @@ async def update_credential( "updated_by": user_api_key_dict.user_id, }, ) + if merged_credential.credential_name != credential_name: + litellm.credential_list = [ + cred + for cred in litellm.credential_list + if cred.credential_name != credential_name + ] + CredentialAccessor.upsert_credentials( + [ + CredentialHelperUtils.decrypt_credential_values( + CredentialItem(**merged_credential.model_dump()) + ) + ] + ) return {"success": True, "message": "Credential updated successfully"} except Exception as e: return handle_exception_on_proxy(e) diff --git a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.test.tsx b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.test.tsx index aee7a0cdd1d..2c8fdb211d6 100644 --- a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.test.tsx +++ b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.test.tsx @@ -1,14 +1,18 @@ +// @vitest-environment jsdom + import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { render, screen, waitFor } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; import { Providers } from "../provider_info_helpers"; import AddCredentialModal from "./AddCredentialModal"; -vi.mock("../networking", async () => { - const actual = await vi.importActual("../networking"); - return { - ...actual, - getProviderCreateMetadata: vi.fn().mockResolvedValue([ +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => ({ accessToken: "test-token" }), +})); + +vi.mock("@/app/(dashboard)/hooks/providers/useProviderFields", () => ({ + useProviderFields: () => ({ + data: [ { provider: "OpenAI", provider_display_name: Providers.OpenAI, @@ -43,9 +47,24 @@ vi.mock("../networking", async () => { }, ], }, - ]), - }; -}); + ], + isLoading: false, + error: null, + }), +})); + +vi.mock("@/hooks/useDeviceCodeFlow", () => ({ + useDeviceCodeFlow: () => ({ + state: { phase: "idle" }, + start: vi.fn(), + reset: vi.fn(), + renderUI: () => null, + }), +})); + +vi.mock("@/components/networking", () => ({ + credentialCreateCall: vi.fn(), +})); const createQueryClient = () => new QueryClient({ diff --git a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx index 83cbafe72f9..77eed496530 100644 --- a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx +++ b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx @@ -3,7 +3,6 @@ import { Select as AntdSelect, Button, Form, Modal, Tooltip, Typography } from " import type { UploadProps } from "antd/es/upload"; import React, { useCallback, useMemo, useState } from "react"; import { credentialCreateCall } from "@/components/networking"; -import NotificationsManager from "../molecules/notifications_manager"; import ProviderSpecificFields from "../add_model/provider_specific_fields"; import { Providers, providerLogoMap } from "../provider_info_helpers"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; @@ -16,11 +15,14 @@ interface AddCredentialsModalProps { onCancel: () => void; onAddCredential: (values: any) => void; uploadProps: UploadProps; - initialCredentialName?: string; - initialProvider?: string; } -const AddCredentialsModal: React.FC = ({ open, onCancel, onAddCredential, uploadProps, initialCredentialName, initialProvider }) => { +const AddCredentialsModal: React.FC = ({ + open, + onCancel, + onAddCredential, + uploadProps, +}) => { const [form] = Form.useForm(); const [selectedProvider, setSelectedProvider] = useState(Providers.OpenAI); const { accessToken } = useAuthorized();