mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix oauth credential refresh and test coverage
This commit is contained in:
parent
788da36fea
commit
6b52abe652
4 changed files with 88 additions and 25 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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({
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue