diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/ai_suggestion_modal.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/ai_suggestion_modal.tsx index 34204624702..89ce1f40105 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/ai_suggestion_modal.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/ai_suggestion_modal.tsx @@ -1,4 +1,4 @@ -import React, { useEffect, useMemo, useState } from "react"; +import React, { useEffect, useMemo, useRef, useState } from "react"; import { Modal, Spin, Checkbox, Select, Input, Typography, Tooltip } from "antd"; import { Button, Card } from "@tremor/react"; import { @@ -93,6 +93,14 @@ const AiSuggestionModal: React.FC = ({ } }, [visible]); + const enrichAbortRef = useRef(null); + + useEffect(() => { + if (!visible) enrichAbortRef.current?.abort(); + }, [visible]); + + useEffect(() => () => enrichAbortRef.current?.abort(), []); + const loadModels = async () => { if (!accessToken) return; setIsLoadingModels(true); @@ -287,6 +295,9 @@ const AiSuggestionModal: React.FC = ({ setIsEnriching(true); setEnrichStatusMessage(""); + enrichAbortRef.current?.abort(); + const controller = new AbortController(); + enrichAbortRef.current = controller; try { for (const template of templatesToEnrich) { const paramName = template.llm_enrichment.parameter; @@ -343,7 +354,7 @@ const AiSuggestionModal: React.FC = ({ (error) => { finish(() => reject(new Error(error))); }, - undefined, + { signal: controller.signal }, (status) => setEnrichStatusMessage(status), ).catch((error) => { finish(() => reject(error)); @@ -351,7 +362,9 @@ const AiSuggestionModal: React.FC = ({ }); } } catch (e) { - console.error("Failed to enrich templates:", e); + if (!controller.signal.aborted) { + console.error("Failed to enrich templates:", e); + } } finally { setIsEnriching(false); setEnrichStatusMessage(""); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/template_parameter_modal.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/template_parameter_modal.test.tsx new file mode 100644 index 00000000000..87ccd56dcc0 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/template_parameter_modal.test.tsx @@ -0,0 +1,92 @@ +import React from "react"; +import { screen, waitFor, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { renderWithProviders } from "@/../tests/test-utils"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import TemplateParameterModal from "./template_parameter_modal"; +import * as networking from "@/components/networking"; + +vi.mock("@/components/networking", () => ({ + modelHubCall: vi.fn(), + enrichPolicyTemplateStream: vi.fn(), +})); + +const template = { + id: "competitor-blocking", + title: "Competitor Blocking", + parameters: [ + { + name: "brand_name", + label: "Brand Name", + type: "string", + required: true, + placeholder: "e.g. Acme Airlines", + }, + ], + llm_enrichment: { parameter: "brand_name" }, +}; + +const defaultProps = { + visible: true, + template, + onConfirm: vi.fn(), + onCancel: vi.fn(), + accessToken: "test-token", +}; + +type UserEvent = ReturnType; + +const startEnrichment = async (user: UserEvent) => { + await user.type(screen.getByPlaceholderText("e.g. Acme Airlines"), "Acme Airlines"); + + const modelSection = screen.getByText("Select Model").closest("div") as HTMLElement; + const combobox = within(modelSection).getByRole("combobox"); + await user.click(combobox); + await user.click(await screen.findByTitle("gpt-4o")); + + await user.click(screen.getByRole("button", { name: /generate competitor names/i })); + + await waitFor(() => { + expect(networking.enrichPolicyTemplateStream).toHaveBeenCalledTimes(1); + }); + + const options = vi.mocked(networking.enrichPolicyTemplateStream).mock.calls[0][7]; + expect(options?.signal).toBeInstanceOf(AbortSignal); + return options!.signal!; +}; + +describe("TemplateParameterModal enrichment stream cancellation", () => { + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(networking.modelHubCall).mockResolvedValue({ + data: [{ model_group: "gpt-4o" }], + } as any); + vi.mocked(networking.enrichPolicyTemplateStream).mockReturnValue(new Promise(() => {})); + }); + + it("aborts an in-flight enrichment stream when the modal unmounts", async () => { + const user = userEvent.setup(); + const { unmount } = renderWithProviders(); + + const signal = await startEnrichment(user); + expect(signal.aborted).toBe(false); + + unmount(); + + expect(signal.aborted).toBe(true); + }); + + it("aborts an in-flight enrichment stream when the modal is closed", async () => { + const user = userEvent.setup(); + const { rerender } = renderWithProviders(); + + const signal = await startEnrichment(user); + expect(signal.aborted).toBe(false); + + rerender(); + + await waitFor(() => { + expect(signal.aborted).toBe(true); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/template_parameter_modal.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/template_parameter_modal.tsx index 600970c6a6b..eb9adc91794 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/template_parameter_modal.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/template_parameter_modal.tsx @@ -1,4 +1,4 @@ -import React, { useState, useEffect } from "react"; +import React, { useState, useEffect, useRef } from "react"; import { Modal, Spin, Radio, Select } from "antd"; import { Button, TextInput } from "@tremor/react"; import { modelHubCall, enrichPolicyTemplateStream } from "@/components/networking"; @@ -75,6 +75,14 @@ const TemplateParameterModal: React.FC = ({ } }, [visible, hasEnrichment, competitorMode]); + const enrichAbortRef = useRef(null); + + useEffect(() => { + if (!visible) enrichAbortRef.current?.abort(); + }, [visible]); + + useEffect(() => () => enrichAbortRef.current?.abort(), []); + const loadModels = async () => { if (!accessToken) return; setIsLoadingModels(true); @@ -100,6 +108,9 @@ const TemplateParameterModal: React.FC = ({ setCompetitorTags([]); setVariationsMap({}); setStatusMessage(""); + enrichAbortRef.current?.abort(); + const controller = new AbortController(); + enrichAbortRef.current = controller; try { await enrichPolicyTemplateStream( accessToken, @@ -121,11 +132,13 @@ const TemplateParameterModal: React.FC = ({ setIsGenerating(false); setStatusMessage(""); }, - undefined, + { signal: controller.signal }, (status) => setStatusMessage(status), ); } catch (error) { - console.error("Error generating competitor names:", error); + if (!controller.signal.aborted) { + console.error("Error generating competitor names:", error); + } setIsGenerating(false); } }; @@ -135,6 +148,9 @@ const TemplateParameterModal: React.FC = ({ setIsRefining(true); setStatusMessage(""); + enrichAbortRef.current?.abort(); + const controller = new AbortController(); + enrichAbortRef.current = controller; try { await enrichPolicyTemplateStream( accessToken, @@ -162,11 +178,14 @@ const TemplateParameterModal: React.FC = ({ { instruction: refinementInput.trim(), existingCompetitors: competitorTags, + signal: controller.signal, }, (status) => setStatusMessage(status), ); } catch (error) { - console.error("Error refining competitor names:", error); + if (!controller.signal.aborted) { + console.error("Error refining competitor names:", error); + } setIsRefining(false); } }; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index d6e9ba5665c..59003c942e0 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -4109,7 +4109,7 @@ export const enrichPolicyTemplateStream = async ( guardrailDefinitions: any[]; }) => void, onError?: (error: string) => void, - options?: { instruction?: string; existingCompetitors?: string[] }, + options?: { instruction?: string; existingCompetitors?: string[]; signal?: AbortSignal }, onStatus?: (message: string) => void, ) => { const url = proxyBaseUrl ? `${proxyBaseUrl}/policy/templates/enrich/stream` : `/policy/templates/enrich/stream`; @@ -4124,6 +4124,7 @@ export const enrichPolicyTemplateStream = async ( "Content-Type": "application/json", }, body: JSON.stringify(body), + signal: options?.signal, }); if (!response.ok) {