mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(ui): abort policy enrichment streams when their modals close or unmount
enrichPolicyTemplateStream had no AbortSignal parameter, so its while(true) reader loop could not be cancelled: closing the AI suggestion or template parameter modal mid-stream left the SSE connection draining until the server finished, with every event callback still calling setState on the unmounted modal. Thread an optional signal through the helper's options and abort from both callers when the modal is closed, unmounted, or a new enrichment run supersedes an in-flight one. Abort rejections are not reported as errors
This commit is contained in:
parent
2b2ae4ca49
commit
f4561c26c0
4 changed files with 133 additions and 8 deletions
|
|
@ -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<AiSuggestionModalProps> = ({
|
|||
}
|
||||
}, [visible]);
|
||||
|
||||
const enrichAbortRef = useRef<AbortController | null>(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<AiSuggestionModalProps> = ({
|
|||
|
||||
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<AiSuggestionModalProps> = ({
|
|||
(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<AiSuggestionModalProps> = ({
|
|||
});
|
||||
}
|
||||
} catch (e) {
|
||||
console.error("Failed to enrich templates:", e);
|
||||
if (!controller.signal.aborted) {
|
||||
console.error("Failed to enrich templates:", e);
|
||||
}
|
||||
} finally {
|
||||
setIsEnriching(false);
|
||||
setEnrichStatusMessage("");
|
||||
|
|
|
|||
|
|
@ -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<typeof userEvent.setup>;
|
||||
|
||||
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(<TemplateParameterModal {...defaultProps} />);
|
||||
|
||||
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(<TemplateParameterModal {...defaultProps} />);
|
||||
|
||||
const signal = await startEnrichment(user);
|
||||
expect(signal.aborted).toBe(false);
|
||||
|
||||
rerender(<TemplateParameterModal {...defaultProps} visible={false} />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(signal.aborted).toBe(true);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -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<TemplateParameterModalProps> = ({
|
|||
}
|
||||
}, [visible, hasEnrichment, competitorMode]);
|
||||
|
||||
const enrichAbortRef = useRef<AbortController | null>(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<TemplateParameterModalProps> = ({
|
|||
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<TemplateParameterModalProps> = ({
|
|||
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<TemplateParameterModalProps> = ({
|
|||
|
||||
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<TemplateParameterModalProps> = ({
|
|||
{
|
||||
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);
|
||||
}
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue