diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 8ecda9c47a1..05df3c2dcbb 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -6178,31 +6178,36 @@ async def test_generate_key_with_agent_id(): @pytest.mark.asyncio async def test_generate_key_helper_fn_agent_id(): - """Test that generate_key_helper_fn correctly stores agent_id in key_data.""" - from unittest.mock import AsyncMock, MagicMock, patch + """Test that generate_key_helper_fn passes agent_id into the insert_data call.""" + from unittest.mock import AsyncMock, MagicMock, call, patch import litellm.proxy.management_endpoints.key_management_endpoints as km mock_prisma_client = AsyncMock() - mock_insert = AsyncMock(return_value=MagicMock(token="sk-test", created_at=None, updated_at=None, litellm_budget_table=None)) + mock_insert = AsyncMock( + return_value=MagicMock( + token="sk-test", + created_at=None, + updated_at=None, + litellm_budget_table=None, + ) + ) mock_prisma_client.insert_data = mock_insert with patch.object(km, "prisma_client", mock_prisma_client): with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): - try: - result = await generate_key_helper_fn( - request_type="key", - agent_id="test-agent-456", - key_alias="test-agent-key", - models=[], - table_name="key", - ) - except Exception: - # We just verify the agent_id was passed to insert_data - pass + await generate_key_helper_fn( + request_type="key", + agent_id="test-agent-456", + key_alias="test-agent-key", + models=[], + table_name="key", + ) - # Verify insert_data was called with agent_id in key_data - if mock_insert.called: - call_kwargs = mock_insert.call_args - key_data = call_kwargs[1].get("data") or (call_kwargs[0][0] if call_kwargs[0] else {}) - assert key_data.get("agent_id") == "test-agent-456" + assert mock_insert.called, "insert_data was never called" + # insert_data is called as insert_data(data=key_data, ...) + call_kwargs = mock_insert.call_args.kwargs + key_data = call_kwargs.get("data", {}) + assert key_data.get("agent_id") == "test-agent-456", ( + f"Expected agent_id='test-agent-456' in key_data, got: {key_data.get('agent_id')}" + ) diff --git a/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx b/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx index 477ec9a3ace..360b9b7ee93 100644 --- a/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx +++ b/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx @@ -93,11 +93,7 @@ const AddAgentForm: React.FC = ({ const handleNext = async () => { try { if (currentStep === 0) { - const fieldsToValidate = - agentType === CUSTOM_AGENT_TYPE - ? ["agent_name"] - : ["agent_name"]; - await form.validateFields(fieldsToValidate); + await form.validateFields(["agent_name"]); const agentName = form.getFieldValue("agent_name"); if (agentName && !newKeyName) { setNewKeyName(`${agentName}-key`); @@ -184,7 +180,12 @@ const AddAgentForm: React.FC = ({ newKeyModels, ); setCreatedKeyValue(keyResponse.key || null); - } else if (keyAssignOption === "existing_key" && selectedExistingKey) { + } else if (keyAssignOption === "existing_key") { + if (!selectedExistingKey) { + message.error("Please select an existing key to assign"); + setIsSubmitting(false); + return; + } await keyUpdateCall(accessToken, { key: selectedExistingKey, agent_id: agentId, diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx index 8f3d62352b5..714507da1aa 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx @@ -298,6 +298,10 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { if (keyOwner === "you") { formValues.user_id = userID; } else if (keyOwner === "agent") { + if (!selectedAgentId) { + message.error("Please select an agent"); + return; + } formValues.agent_id = selectedAgentId; }