fix oauth credential refresh and test coverage

This commit is contained in:
Hunter Wittenborn 2026-04-25 07:08:48 -05:00
parent 788da36fea
commit 6b52abe652
4 changed files with 88 additions and 25 deletions

View file

@ -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,

View file

@ -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)

View file

@ -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({

View file

@ -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<AddCredentialsModalProps> = ({ open, onCancel, onAddCredential, uploadProps, initialCredentialName, initialProvider }) => {
const AddCredentialsModal: React.FC<AddCredentialsModalProps> = ({
open,
onCancel,
onAddCredential,
uploadProps,
}) => {
const [form] = Form.useForm();
const [selectedProvider, setSelectedProvider] = useState<Providers>(Providers.OpenAI);
const { accessToken } = useAuthorized();