diff --git a/ui/litellm-dashboard/src/components/policies/ai_suggestion_modal.tsx b/ui/litellm-dashboard/src/components/policies/ai_suggestion_modal.tsx index aff629ed4f4..623c291f079 100644 --- a/ui/litellm-dashboard/src/components/policies/ai_suggestion_modal.tsx +++ b/ui/litellm-dashboard/src/components/policies/ai_suggestion_modal.tsx @@ -1,7 +1,7 @@ -import React, { useState } from "react"; -import { Modal, Spin, Checkbox } from "antd"; -import { Button, TextInput } from "@tremor/react"; -import { suggestPolicyTemplates } from "../networking"; +import React, { useState, useEffect } from "react"; +import { Modal, Spin, Checkbox, Select } from "antd"; +import { Button } from "@tremor/react"; +import { suggestPolicyTemplates, modelHubCall } from "../networking"; interface SuggestedTemplate { template_id: string; @@ -31,6 +31,33 @@ const AiSuggestionModal: React.FC = ({ const [suggestions, setSuggestions] = useState(null); const [explanation, setExplanation] = useState(null); const [selectedIds, setSelectedIds] = useState>(new Set()); + const [selectedModel, setSelectedModel] = useState(undefined); + const [availableModels, setAvailableModels] = useState([]); + const [isLoadingModels, setIsLoadingModels] = useState(false); + + useEffect(() => { + if (visible && availableModels.length === 0) { + loadModels(); + } + }, [visible]); + + const loadModels = async () => { + if (!accessToken) return; + setIsLoadingModels(true); + try { + const fetchedModels = await modelHubCall(accessToken); + if (fetchedModels?.data?.length > 0) { + const models = fetchedModels.data + .map((item: any) => item.model_group as string) + .sort(); + setAvailableModels(models); + } + } catch (error) { + console.error("Failed to load models:", error); + } finally { + setIsLoadingModels(false); + } + }; const resetState = () => { setAttackExamples([""]); @@ -39,6 +66,7 @@ const AiSuggestionModal: React.FC = ({ setSuggestions(null); setExplanation(null); setSelectedIds(new Set()); + setSelectedModel(undefined); }; const handleCancel = () => { @@ -67,18 +95,18 @@ const AiSuggestionModal: React.FC = ({ description.trim().length > 0; const handleSuggest = async () => { - if (!accessToken || !hasInput) return; + if (!accessToken || !hasInput || !selectedModel) return; setIsLoading(true); try { const result = await suggestPolicyTemplates( accessToken, attackExamples, - description + description, + selectedModel ); setSuggestions(result.selected_templates || []); setExplanation(result.explanation || null); - // Pre-select all suggested templates setSelectedIds( new Set( (result.selected_templates || []).map( @@ -125,187 +153,296 @@ const AiSuggestionModal: React.FC = ({ return ( -

AI Policy Suggestion

-

- {showResults - ? "Select which templates to use" - : "Describe what you want to block and we'll suggest the best policy templates"} -

- - } + title={null} open={visible} onCancel={handleCancel} - width={600} - footer={ - showResults - ? [ - , - , - ] - : [ - , - , - ] - } + width={820} + footer={null} + styles={{ body: { padding: 0 } }} > + {/* Header */} +
+

+ AI Policy Suggestion +

+

+ {showResults + ? `${suggestions?.length || 0} template${(suggestions?.length || 0) !== 1 ? "s" : ""} matched your requirements` + : "Describe what you want to block and we'll suggest the best policy templates"} +

+
+ +
+ {!showResults ? ( -
+ /* ── Input phase ── */ +
+ {/* Model selector */}
-