diff --git a/doc/code/framework.md b/doc/code/framework.md index 41b1c1e338..0b30cf34c3 100644 --- a/doc/code/framework.md +++ b/doc/code/framework.md @@ -403,6 +403,8 @@ See [message normalizers](./targets/11_message_normalizer) for capability behavi - If you are creating a component with user input (e.g. via config, REST, or automatically) it should always use the registry - If you are storing an instance of a component, it should always use the registry +- Construction does not require registration. The caller owns temporary instances + and their cleanup; the registry does not retain them. - The registry accepts only explicitly supported external inputs, permits opaque Python objects only for in-process callers, and leaves component validation to constructors. ## [Setup](./setup/0_setup) diff --git a/doc/code/registry/0_registry.md b/doc/code/registry/0_registry.md index 0d0fc93c01..a6c88f4440 100644 --- a/doc/code/registry/0_registry.md +++ b/doc/code/registry/0_registry.md @@ -75,6 +75,62 @@ creation can build more than one component for the same name, but only one registration succeeds unless replacement is requested. The backend does not need a separate registration lock. +### Construction without registration + +Use `create_instance(type_name, **params)` to build an object without adding it +to `.instances`. The caller owns that object and its cleanup. +`create_named_instance(...)` remains the build-and-register operation. + +The existing `POST /api/targets`, `POST /api/converters`, and `POST /api/scorers` +endpoints accept `register: false`. Registration defaults to `true`, so existing +clients keep their behavior. For example: + +```json +{ + "type": "CaesarConverter", + "register": false, + "params": { "caesar_offset": 3 } +} +``` + +The response contains an `identifier` and, for targets, `capabilities`. It does +not contain a registry name or a reusable object handle. The backend discards +the built object after returning its descriptor. + +Targets and converters can also use a registered source that supports +reconstruction: + +```json +{ + "type": "OpenAIChatTarget", + "register": false, + "source": { + "source_name": "objective", + "source_hash": "", + "params": { "temperature": 0.7 } + } +} +``` + +Use `source.params` for overrides; do not also supply `params`. The backend +keeps the source's other constructor inputs, including authentication, on the +server. It does not change the source. Missing, changed, or ambiguous sources +produce an error. If a source was renamed, a unique matching hash can resolve it. +An optional `effective_hash` checks the reconstructed object's identity. +Registered descriptors report `reconstructable` and, for targets, +`supports_temperature_override`. Unsupported inputs can also report +`reconstruction_error`. +Constructor inputs that are not in the registry's build contract are not +discarded silently: reconstruction is rejected. Such sources can still be used +as registered objects. +Reconstruction retains resolved defaults from the declared constructor chain, +including explicitly forwarded parent parameters. Arguments generated inside a +constructor are not treated as extra inputs from its caller. + +This REST flag applies to targets, converters, and scorers. `AttackRegistry` +remains a class catalog, not a store of attack instances, and has no new REST +endpoint. + Constructor annotations define parameter metadata and coercion. Enum parameters accept member names or values. Types that inherit `StructuredParameterValue` declare their allowed variants through `get_registry_input_variants()`; the registry diff --git a/doc/gui/0_gui.md b/doc/gui/0_gui.md index d77bd99e16..af79caff39 100644 --- a/doc/gui/0_gui.md +++ b/doc/gui/0_gui.md @@ -102,7 +102,7 @@ The Chat view is the primary workspace for running interactive attacks against c #### Sending Messages -For a new chat, your default objective target is preselected if it is available. Click the target badge in the shared toolbar beside the label controls to open the target dropdown. If no target is selected, click **Select a target** in the same place. Your choice applies to this chat without changing the default. Saved chats keep their original target. An attack saved without a target uses this dropdown until its first send binds the selected target. +For a new chat, your default objective target is preselected if it is available. Click the target badge in the chat ribbon below the common labels and above the objective to open the target dropdown. If no target is selected, click **Select a target** in the same place. Your choice applies to this chat without changing the default. Saved chats keep their original target. An attack saved without a target uses this dropdown until its first send binds the selected target. Clicking **Chat** while already in a new chat keeps its target and draft. Starting a new attack resets both. Default changes in another tab apply to the next new @@ -112,6 +112,22 @@ Type a message and press Enter (or click Send) to send it to the chat target. Th When you open a saved chat, CoPyRIT automatically selects the target originally used, if its registered identity still matches. This also applies to direct links, reloads, and browser Back/Forward navigation. You can continue the same conversation without selecting the target again. Opening a saved chat does not change your defaults. +#### Temperature + +Select a target, then set **Temperature** before the first send. Leave it empty +to keep the source setting. The supported range is 0 to 2. This creates a private +target configuration for the attack; it does not add or change a registered target. +Temperature is read-only after the attack is bound. In the conversation editor, +a temperature change requires **New attack**, not **Same attack**. +Hover over or click the read-only field to see why it cannot be changed. + +Saved attacks retain the source name, source identity, temperature, and effective +identity, but not credentials. After a restart, the backend tries to reconstruct +the target from a matching registered source. If reconstruction fails, sending +is blocked rather than using the source's default temperature. OpenAI-family +targets with a temperature parameter support this control. Externally owned HTTP +clients and temperature set through `extra_body_parameters` are not supported. + #### Repeating a Message Use **n=1** beside Send to choose **1 to 10** repetitions. Enter and Send use the @@ -161,6 +177,20 @@ messages endpoint and `send=false` context storage are unchanged. Open **Converters** and use the picker above the working input to add registered converters in the order you want them to run. +For a stage that supports reconstruction, open **... > Settings** to change its +constructor settings. Changes create a private converter for that stage. Other +stages and the registered source keep their settings. **Reset to registered +converter** removes the stage's overrides. **Use default / not set** clears a +structured setting's override and restores the registered source's value. +Changing settings invalidates that stage and its downstream results. + +Closing the converter pane discards temporary settings, but keeps content that +you already applied and the identifiers of the converters that produced it. +Temporary converters are built for each preview operation, not retained in a +second registry. A runtime change clears temporary settings and preview results. +Apply temporary converter results with **Add converted value** before repeating +a message. Independent repeated conversion supports registered stages only. + The top text box is an editable working copy: changing it does not change the original chat message. The top **Convert** button runs the active tab's configured pipeline and any configured inputs that do not have results yet. After every configured input @@ -249,6 +279,15 @@ represents a manual conversion. `request_converter_configurations` controls conversion of pieces without a preconverted value; it does not describe which converters already ran. +For temporary stages, preview requests provide `converter_specs`, aligned with +`converter_ids`; use `null` for a registered stage. Preview responses include the +actual temporary identifier and a signed `provenance` token. Pass these tokens +as `applied_converter_provenance`, aligned with `applied_converter_ids`, when +sending or saving the converted content. Tokens remain valid after the pane +closes, without retaining converter objects. A restart invalidates tokens for +unsaved content; convert it again. Stored messages retain their converter +identifiers and do not depend on those tokens. + #### Attachments Click the attachment button to add images, audio, video, or documents to your message. Supported types include `image/*`, `audio/*`, `video/*`, `.pdf`, `.doc`, `.docx`, and `.txt`. Attachments are displayed as chips below the input with type icons and file sizes. @@ -371,7 +410,7 @@ Export stays available for read-only historical conversations, and is disabled w The labels bar above the page content is available across the GUI, including scanner setup, Home, Chat, and History. It shows the active labels for future attacks and scans, not the attribution of a historical run you are viewing. Click the labels icon to open **Default Labels** and add, edit, or remove custom labels. The required `operator` and `operation` controls remain in the bar, outside this popover, and cannot be removed. A signed-in operator is read-only. -In Chat, the active target, Markdown toggle, export menu, conversations panel toggle, and **New Attack** button share the right side of this bar. They wrap below the labels on narrow screens. +In Chat, a separate ribbon below the labels contains the target, temperature, and **Edit Conversation** controls on the left. The Markdown toggle, export menu, conversations panel toggle, and **New Attack** button are on the right. Clicking the `operation` label opens a picker listing the operations already recorded in memory, so you can choose one without typing it from memory. Typing a name that doesn't exist yet offers to create it. Very long lists show the first 200 and say how many are left, so type to narrow them. On narrow screens, use the labels icon to view or edit labels that do not fit inline. diff --git a/frontend/src/App.labels.test.tsx b/frontend/src/App.labels.test.tsx index c2ad8fa90e..64fcab3150 100644 --- a/frontend/src/App.labels.test.tsx +++ b/frontend/src/App.labels.test.tsx @@ -247,20 +247,34 @@ describe('Shared new run labels', () => { })) }) - it('hosts Chat controls beside the labels and removes them when navigating away', async () => { + it('hosts Chat controls below the labels and above the objective, and removes them when navigating away', async () => { const user = userEvent.setup() + jest.mocked(targetsApi.listTargets).mockResolvedValue({ + items: [{ ...TARGET, supports_temperature_override: true }], + pagination: { limit: 200, has_more: false }, + }) renderApp('/chat') - const toolbar = within(currentLabels()).getByRole('group', { name: 'Chat controls' }) - const targetPicker = await within(toolbar).findByRole('combobox', { name: 'Chat target' }) + const toolbar = screen.getByRole('group', { name: 'Chat controls' }) + expect(currentLabels()).not.toContainElement(toolbar) + expect(currentLabels().compareDocumentPosition(toolbar) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy() + expect(toolbar.compareDocumentPosition(screen.getByTestId('objective-header')) & Node.DOCUMENT_POSITION_FOLLOWING) + .toBeTruthy() + const settings = within(toolbar).getByRole('group', { name: 'Conversation settings' }) + const actions = within(toolbar).getByRole('group', { name: 'Conversation actions' }) + const targetPicker = await within(settings).findByRole('combobox', { name: 'Chat target' }) await waitFor(() => expect(targetPicker).toBeEnabled()) await user.selectOptions(targetPicker, 'test_target') expect(targetPicker).toHaveValue('test_target') + const temperature = within(settings).getByRole('spinbutton', { name: 'Temperature' }) + await user.type(temperature, '0.7') + expect(temperature).toHaveValue(0.7) + expect(within(settings).getByRole('button', { name: 'Edit Conversation' })).toBeEnabled() expect(within(screen.getByTestId('chat-area')).queryByRole('group', { name: 'Chat controls' })) .not.toBeInTheDocument() - expect(within(toolbar).getByRole('button', { name: 'Export conversation' })).toBeDisabled() - expect(within(toolbar).getByRole('button', { name: 'Toggle conversations panel' })).toBeDisabled() - expect(within(toolbar).getByRole('button', { name: 'New Attack' })).toBeDisabled() - const markdown = within(toolbar).getByRole('switch') + expect(within(actions).getByRole('button', { name: 'Export conversation' })).toBeDisabled() + expect(within(actions).getByRole('button', { name: 'Toggle conversations panel' })).toBeDisabled() + expect(within(actions).getByRole('button', { name: 'New Attack' })).toBeDisabled() + const markdown = within(actions).getByRole('switch') await user.click(markdown) expect(readUserPreferences('local').chatMarkdown).toBe(true) @@ -271,7 +285,7 @@ describe('Shared new run labels', () => { expect(screen.queryByRole('group', { name: 'Chat controls' })).not.toBeInTheDocument() await user.click(screen.getByRole('button', { name: 'Chat', exact: true })) expect(screen.getAllByRole('group', { name: 'Chat controls' })).toHaveLength(1) - expect(within(currentLabels()).getByRole('switch')).toBeChecked() + expect(within(screen.getByRole('group', { name: 'Conversation actions' })).getByRole('switch')).toBeChecked() }) it('preserves an edit made before the backend defaults arrive', async () => { @@ -330,7 +344,7 @@ describe('Shared new run labels', () => { expect(attacksApi.createAttack).not.toHaveBeenCalled() expect(attacksApi.submitMessageSend).not.toHaveBeenCalled() - const toolbar = within(currentLabels()).getByRole('group', { name: 'Chat controls' }) + const toolbar = screen.getByRole('group', { name: 'Chat controls' }) expect(within(toolbar).getByLabelText('Active target: test_target')).toBeInTheDocument() await waitFor(() => expect(within(toolbar).getByRole('button', { name: 'Export conversation' })).toBeEnabled()) await user.click(within(toolbar).getByRole('button', { name: 'Export conversation' })) diff --git a/frontend/src/App.tsx b/frontend/src/App.tsx index e110b0e8c1..c92fe26d65 100644 --- a/frontend/src/App.tsx +++ b/frontend/src/App.tsx @@ -499,6 +499,7 @@ function AppContent({ operatorAlias }: { operatorAlias: string | null }) { endpoint: targetEndpoint(createdTarget), model_name: targetModelName(createdTarget), identifier_hash: targetIdentifierHash(createdTarget), + binding: createdTarget.binding, } : null skipNextLoadForAttackId.current = arId diff --git a/frontend/src/components/Chat/ChatWindow.styles.ts b/frontend/src/components/Chat/ChatWindow.styles.ts index 57756b3b78..e3c6a525ca 100644 --- a/frontend/src/components/Chat/ChatWindow.styles.ts +++ b/frontend/src/components/Chat/ChatWindow.styles.ts @@ -90,15 +90,34 @@ export const useChatWindowStyles = makeStyles({ flexGrow: 1, flexWrap: 'wrap', gap: tokens.spacingHorizontalM, + minWidth: 0, + maxWidth: '100%', + }, + conversationControls: { + display: 'flex', + alignItems: 'center', + flexWrap: 'wrap', + gap: tokens.spacingHorizontalM, + marginRight: 'auto', + minWidth: 0, maxWidth: '100%', }, editActions: { display: 'flex', flexWrap: 'wrap', alignItems: 'center', - marginRight: 'auto', gap: tokens.spacingHorizontalXS, }, + temperatureField: { + gridTemplateColumns: 'max-content 5rem', + columnGap: tokens.spacingHorizontalNone, + alignItems: 'center', + flexShrink: 0, + }, + temperatureInput: { + width: '5rem', + minWidth: 0, + }, sharedTarget: { maxWidth: '240px', [NARROW_VIEWPORT_QUERY]: { @@ -123,6 +142,7 @@ export const useChatWindowStyles = makeStyles({ flexWrap: 'wrap', justifyContent: 'flex-end', minWidth: 0, + marginLeft: 'auto', }, ribbonAction: { ...mobileTouchTarget, diff --git a/frontend/src/components/Chat/ChatWindow.test.tsx b/frontend/src/components/Chat/ChatWindow.test.tsx index 8adfbacbb9..c5d9b7688a 100644 --- a/frontend/src/components/Chat/ChatWindow.test.tsx +++ b/frontend/src/components/Chat/ChatWindow.test.tsx @@ -26,7 +26,7 @@ import { TargetInstance, TargetResponseStatus, } from "../../types"; -import { attacksApi, convertersApi, scoresApi } from "../../services/api"; +import { attacksApi, convertersApi, scoresApi, targetsApi } from "../../services/api"; import * as messageMapper from "../../utils/messageMapper"; const buildCapabilities = ( @@ -73,6 +73,7 @@ jest.mock("../../services/api", () => ({ deleteConverter: jest.fn(), previewConversion: jest.fn(), }, + targetsApi: { buildTarget: jest.fn() }, labelsApi: { getLabels: jest.fn().mockImplementation(() => new Promise(() => {})), }, @@ -1710,6 +1711,119 @@ describe("ChatWindow Integration", () => { // First message → create attack + send // ----------------------------------------------------------------------- + it("builds a private temperature setting before creating the attack", async () => { + const user = userEvent.setup(); + const source = { ...mockTarget, supports_temperature_override: true }; + const onConversationCreated = jest.fn(); + jest.mocked(targetsApi.buildTarget).mockResolvedValue({ + identifier: { ...source.identifier, hash: "private-hash", temperature: 0.8 }, + capabilities: source.capabilities, + }); + mockedMapper.backendMessagesToFrontend.mockReturnValue([]); + mockedMapper.buildMessagePieces.mockResolvedValue([{ data_type: "text", original_value: "Hello" }]); + mockSendResult.mockResolvedValue({ + ...makeTextResponse("Reply"), + attack: { attack_result_id: "default-created-attack", conversation_id: "default-created-conversation" }, + } as never); + render(); + expect(screen.getByText("Temperature:")).toBeInTheDocument(); + expect(screen.queryByText("Empty keeps the source setting.")).not.toBeInTheDocument(); + await user.type(screen.getByLabelText("Temperature"), "0.8"); + await user.type(screen.getByRole("textbox"), "Hello"); + await user.click(screen.getByRole("button", { name: /send/i })); + await waitFor(() => expect(mockedAttacksApi.createAttack).toHaveBeenCalledTimes(1)); + const binding = { + version: 1, source_name: source.target_registry_name, source_hash: source.identifier.hash, + temperature: 0.8, effective_hash: "private-hash", + }; + expect(targetsApi.buildTarget).toHaveBeenCalledWith("OpenAIChatTarget", { + source_name: source.target_registry_name, source_hash: source.identifier.hash, + params: { temperature: 0.8 }, + }); + expect(mockedAttacksApi.createAttack.mock.calls[0][0].target_binding).toEqual(binding); + expect(onConversationCreated.mock.calls[0][3].binding).toEqual(binding); + expect(source.identifier.hash).not.toBe("private-hash"); + }); + + it.each([ + { ready: true, generation: "gen-1", defaultsReady: false }, + { ready: true, generation: "gen-2", defaultsReady: true }, + { ready: false, generation: "gen-1", defaultsReady: true }, + ])("preserves the first-send draft when launch state becomes %j during temperature construction", async ( + next: { ready: boolean; generation: string; defaultsReady: boolean }, + ) => { + const user = userEvent.setup(); + const runtime = jest.spyOn(runtimeHooks, "useRuntime").mockReturnValue({ + ready: true, state: "ready", generation: "gen-1", + }); + const source = { ...mockTarget, supports_temperature_override: true }; + const built = { + identifier: { ...source.identifier, hash: "private-hash", temperature: 0.8 }, + capabilities: source.capabilities, + }; + let finish: (value: typeof built) => void = () => {}; + jest.mocked(targetsApi.buildTarget).mockImplementationOnce(() => new Promise((resolve) => { finish = resolve; })); + mockedMapper.buildMessagePieces.mockResolvedValue([{ data_type: "text", original_value: "Prepared prompt" }]); + const rendered = render(); + await user.type(screen.getByLabelText("Temperature"), "0.8"); + const input = screen.getByRole("textbox"); + await user.type(input, "Prepared prompt"); + await user.click(screen.getByRole("button", { name: "Send message" })); + await waitFor(() => expect(targetsApi.buildTarget).toHaveBeenCalledTimes(1)); + runtime.mockReturnValue({ + ready: next.ready, state: next.ready ? "ready" : "unavailable", generation: next.generation, + }); + rendered.rerender(); + await act(async () => { finish(built); }); + expect(await screen.findByText(/Runtime or default labels changed while preparing this message/)).toBeInTheDocument(); + expect(input).toHaveValue("Prepared prompt"); + expect(screen.getByLabelText("Temperature")).toHaveValue(0.8); + expect(mockedAttacksApi.createAttack).not.toHaveBeenCalled(); + expect(mockedAttacksApi.submitMessageSend).not.toHaveBeenCalled(); + }); + + it("shows saved temperature as read-only", () => { + const target = { + ...mockTarget, supports_temperature_override: true, + identifier: { ...mockTarget.identifier, temperature: 0.8 }, + }; + render(); + expect(screen.getByLabelText("Temperature")).toHaveAttribute("readonly"); + expect(screen.getByLabelText("Temperature")).toHaveValue(0.8); + expect(screen.queryByText("Read-only for this attack.")).not.toBeInTheDocument(); + expect(screen.queryByRole("tooltip")).not.toBeInTheDocument(); + }); + + it.each([ + { interaction: "hover", supportsOverride: true }, + { interaction: "click", supportsOverride: true }, + { interaction: "hover", supportsOverride: false }, + { interaction: "click", supportsOverride: false }, + ])("explains read-only temperature on $interaction with override support $supportsOverride", async ({ + interaction, supportsOverride, + }: { interaction: string; supportsOverride: boolean }) => { + const user = userEvent.setup(); + const target = { + ...mockTarget, supports_temperature_override: supportsOverride, + identifier: { ...mockTarget.identifier, temperature: 0.8 }, + }; + render(); + const input = screen.getByLabelText("Temperature"); + expect(screen.queryByText("Read-only for this attack.")).not.toBeInTheDocument(); + if (interaction === "hover") await user.hover(input); + else await user.click(input); + expect(await screen.findByRole("tooltip")).toHaveTextContent( + "Temperature can only be modified in a new attack" + ); + expect(input).toHaveAttribute("readonly"); + expect(input).toHaveValue(0.8); + expect(targetsApi.buildTarget).not.toHaveBeenCalled(); + }); + it("should create attack and send text message on first message", async () => { const user = userEvent.setup(); const onConversationCreated = jest.fn(); @@ -5547,6 +5661,29 @@ describe("ChatWindow Integration", () => { }); }); + it("should show only the registered name and LLM badge in the picker and keep stage settings accessible", async () => { + const user = userEvent.setup(); + mockedConvertersApi.listConverters.mockResolvedValue({ + items: [{ + ...makeConverterInstance("translation_spanish", "TranslationConverter"), + is_llm_based: true, + reconstructable: true, + }], + }); + render(); + await user.click(screen.getByRole("button", { name: "Toggle converter panel" })); + await user.click(await screen.findByRole("combobox", { name: "Add converter" })); + const option = await screen.findByRole("option", { name: /translation_spanish.*TranslationConverter/ }); + expect(within(option).getByText("translation_spanish")).toBeInTheDocument(); + expect(within(option).getByText("LLM")).toBeInTheDocument(); + expect(within(option).queryByText("TranslationConverter")).not.toBeInTheDocument(); + await user.click(option); + expect(screen.getByRole("button", { name: "Reorder converter translation_spanish" })).toBeEnabled(); + expect(screen.getByRole("button", { name: "Remove converter translation_spanish" })).toBeEnabled(); + await user.click(screen.getByRole("button", { name: "Settings for converter translation_spanish" })); + expect(screen.getByRole("menuitem", { name: "Settings", exact: true })).not.toHaveAttribute("aria-disabled", "true"); + }); + it("should not show constructor parameters for a registered converter", async () => { mockedConvertersApi.listConverters.mockResolvedValue({ items: [makeConverterInstance("base64-configured", "Base64Converter")], diff --git a/frontend/src/components/Chat/ChatWindow.tsx b/frontend/src/components/Chat/ChatWindow.tsx index 157098f6f9..3d9b0766c1 100644 --- a/frontend/src/components/Chat/ChatWindow.tsx +++ b/frontend/src/components/Chat/ChatWindow.tsx @@ -3,6 +3,8 @@ import type { ChangeEvent } from 'react' import { createPortal } from 'react-dom' import { Button, + Field, + Input, Breadcrumb, BreadcrumbDivider, BreadcrumbItem, @@ -44,6 +46,7 @@ import type { PieceConversion } from './converterTypes' import { useChatConverters } from '@/hooks/useChatConverters' import { useRuntime } from '@/hooks/useRuntime' import { useUserPreferences } from '@/hooks/useUserPreferences' +import { buildTemperatureTarget } from '@/services/targetRegistry' import { basenameFromValue, applyConvertedValues, @@ -339,6 +342,13 @@ export default function ChatWindow({ const restoreFocusSourceAttributes = useRestoreFocusSource() const [messages, setMessages] = useState([]) const [pendingObjective, setPendingObjective] = useState('') + const [temperature, setTemperature] = useState('') + const [temperatureExplanationAttackId, setTemperatureExplanationAttackId] = useState(null) + const [temperatureAttackId, setTemperatureAttackId] = useState(attackResultId) + if (temperatureAttackId !== attackResultId) { + setTemperatureAttackId(attackResultId) + setTemperature('') + } const currentObjective = attackResultId ? objective : pendingObjective const runtime = useRuntime() const newAttackContext: NewAttackContext = { generation: runtime.generation, ready: runtime.ready && defaultsReady, labels } @@ -675,6 +685,13 @@ export default function ChatWindow({ !operation.controller.signal.aborted && pendingSendsRef.current.get(operation.conversationId) === operation ) + const checkSendPreparation = (operation: PendingSend, requireDefaults: boolean): void => { + const current = launchStateRef.current + if (current.generation !== operation.converterGeneration || !current.ready + || (requireDefaults && !current.defaultsReady)) { + throw new Error('Runtime or default labels changed while preparing this message. Your draft is preserved. Retry after default labels finish loading.') + } + } const isViewingSend = (operation: PendingSend): boolean => ( viewedAttackRef.current === operation.attackResultId && ( @@ -1064,6 +1081,9 @@ export default function ChatWindow({ pieceIds, conversions, ) + const requestConfigurations = count > 1 && converterMode === 'per_branch' ? buildRequestConverterConfigurations( + buildConverterInputs(originalValue, attachments), pieceIds, pipelines ?? {}, conversions, + ) : [] if (!isCurrentSend(operation)) { return { status: 'non_retryable_failure', clearDraft: false } } // Create attack lazily on first message @@ -1071,13 +1091,16 @@ export default function ChatWindow({ let currentConversationId = conversationId let currentActiveConversationId = activeConversationId if (!currentAttackResultId) { - const currentLaunchState = launchStateRef.current - if (currentLaunchState.generation !== operation.converterGeneration - || !currentLaunchState.ready || !currentLaunchState.defaultsReady) { - throw new Error('Runtime or default labels changed while preparing this message. Your draft is preserved. Retry after default labels finish loading.') + checkSendPreparation(operation, true) + let selectedTarget = activeTarget + if (temperature.trim()) { + selectedTarget = await buildTemperatureTarget(activeTarget, Number(temperature)) + if (!isCurrentSend(operation)) { return { status: 'non_retryable_failure', clearDraft: false } } } + checkSendPreparation(operation, true) const createRequest: CreateAttackRequest = { target_registry_name: activeTarget.target_registry_name, + ...(selectedTarget.binding ? { target_binding: selectedTarget.binding } : {}), name: pendingObjective || undefined, // TODO(PyRIT 1.4): Pass only dedicated attribution after legacy label aliases are removed. // The create-attack API normalizes these aliases through _AttackAttributionInput. @@ -1104,7 +1127,11 @@ export default function ChatWindow({ operation.conversationId = currentConversationId pendingSendsRef.current.set(currentConversationId, operation) if (navigationRevisionRef.current === submittedNavigationRevision) { - onConversationCreated(currentAttackResultId, currentConversationId, pendingObjective || undefined) + if (selectedTarget.binding) { + onConversationCreated(currentAttackResultId, currentConversationId, pendingObjective || undefined, selectedTarget) + } else { + onConversationCreated(currentAttackResultId, currentConversationId, pendingObjective || undefined) + } viewedAttackRef.current = currentAttackResultId viewedConvRef.current = currentConversationId } @@ -1124,9 +1151,6 @@ export default function ChatWindow({ if (!currentAttackResultId || !effectiveConvId) { throw new Error('Message send is missing an attack or conversation ID.') } - const requestConfigurations = count > 1 && converterMode === 'per_branch' ? buildRequestConverterConfigurations( - buildConverterInputs(originalValue, attachments), pieceIds, pipelines ?? {}, conversions, - ) : [] const addMessageRequest: MessageSendRequest = { role: 'user', pieces, @@ -1140,6 +1164,12 @@ export default function ChatWindow({ } : {}), ...(requestConfigurations.length ? { request_converter_configurations: requestConfigurations } : {}), } + if (targetResolutionStatus === 'unbound' && temperature.trim()) { + const selectedTarget = await buildTemperatureTarget(activeTarget, Number(temperature)) + if (!isCurrentSend(operation)) { return { status: 'non_retryable_failure', clearDraft: false } } + addMessageRequest.target_binding = selectedTarget.binding + } + checkSendPreparation(operation, false) submissionAttempted = true operation.progress = await attacksApi.submitMessageSend(currentAttackResultId, addMessageRequest) if (!isCurrentSend(operation)) { return { status: 'non_retryable_failure', clearDraft: false } } @@ -1476,7 +1506,7 @@ export default function ChatWindow({ } } - const handleEditorSaved = (response: AddMessageResponse): void => { + const handleEditorSaved = (response: AddMessageResponse, savedTarget?: TargetInstance | null): void => { editor.discard() setEditorNotice('Conversation saved.') setMessages(backendMessagesToFrontend(response.messages.messages)) @@ -1484,7 +1514,7 @@ export default function ChatWindow({ if (response.attack.attack_result_id === attackResultId) { onSelectConversation(response.messages.conversation_id) } else { - onConversationCreated(response.attack.attack_result_id, response.messages.conversation_id, response.attack.objective, editorTarget) + onConversationCreated(response.attack.attack_result_id, response.messages.conversation_id, response.attack.objective, savedTarget ?? editorTarget) } onAttackChange?.(response.attack) @@ -1534,6 +1564,7 @@ export default function ChatWindow({ } const sameAttackDisabledReason = !attackResultId ? 'No saved attack exists yet.' + : editDraft?.temperature.trim() ? 'Choose New attack to change the temperature.' : attackOperator && attackOperator !== currentOperator ? 'This attack belongs to another operator.' : attackTarget && (!editorTarget || !targetInfoMatchesTarget(attackTarget, editorTarget)) ? 'The selected target differs from this attack. Choose New attack.' @@ -1545,6 +1576,8 @@ export default function ChatWindow({ ? 'Default labels are not ready. Retry after default labels finish loading.' : undefined const editorDataTypes = draftDataTypes(editDraft?.messages ?? []) + const temperatureTarget = editDraft !== null ? editorTarget : activeTarget + const temperatureReadOnly = editDraft === null && Boolean(attackResultId) && targetResolutionStatus !== 'unbound' const singleTurnLimitReached = activeTarget?.capabilities?.supports_multi_turn === false && messages.some(m => m.role === 'user') const recoverableProcessingErrorIndex = recoverableSend?.conversationId === viewedConversationId @@ -1593,35 +1626,74 @@ export default function ChatWindow({ role="group" aria-label="Chat controls" > -
- {(!attackResultId || targetResolutionStatus === 'unbound' || editDraft !== null) && !isLoadingAttack ? ( - editorTargetDisabledReason(target, editorDataTypes) : undefined} - /> - ) : activeTarget ? ( - - ) : ( - - No target selected - +
+
+ {(!attackResultId || targetResolutionStatus === 'unbound' || editDraft !== null) && !isLoadingAttack ? ( + { + setTemperature('') + if (editDraft !== null) editor.changeTarget(target) + else onSelectTarget(target) + }} + disabledReason={editDraft !== null + ? (target: TargetInstance) => editorTargetDisabledReason(target, editorDataTypes) : undefined} + /> + ) : activeTarget ? ( + + ) : ( + + No target selected + + )} +
+ {temperatureTarget && ( + + { + setTemperatureExplanationAttackId(temperatureReadOnly && data.visible ? attackResultId : null) + }}> + { + if (temperatureReadOnly) setTemperatureExplanationAttackId(attackResultId) + }}> + { + if (editDraft !== null) editor.changeTemperature(data.value) + else setTemperature(data.value) + }} /> + + + )} +
+ + {editDraft !== null && } +
-
- - {editDraft !== null && } -
-
+
sameAttackDisabledReason?: string newAttackDisabledReason?: string - onSaved: (response: AddMessageResponse) => void + onSaved: (response: AddMessageResponse, target?: TargetInstance | null) => void } export interface ConversationEditorHandle { diff --git a/frontend/src/components/Chat/ConverterPanel/ConverterPanel.styles.ts b/frontend/src/components/Chat/ConverterPanel/ConverterPanel.styles.ts index 26d4e83c32..c92736c28a 100644 --- a/frontend/src/components/Chat/ConverterPanel/ConverterPanel.styles.ts +++ b/frontend/src/components/Chat/ConverterPanel/ConverterPanel.styles.ts @@ -120,15 +120,11 @@ export const useConverterPanelStyles = makeStyles({ optionHeader: { display: 'flex', alignItems: 'baseline', - justifyContent: 'space-between', + justifyContent: 'flex-start', flexWrap: 'wrap', minWidth: 0, gap: `${tokens.spacingVerticalXXS} ${tokens.spacingHorizontalS}`, }, - optionType: { - color: tokens.colorNeutralForeground3, - fontFamily: tokens.fontFamilyMonospace, - }, converterCard: { display: 'flex', flexDirection: 'column', @@ -141,7 +137,7 @@ export const useConverterPanelStyles = makeStyles({ converterCardHeader: { display: 'flex', alignItems: 'center', - justifyContent: 'space-between', + justifyContent: 'flex-start', gap: tokens.spacingHorizontalS, cursor: 'grab', userSelect: 'none', @@ -157,6 +153,22 @@ export const useConverterPanelStyles = makeStyles({ cursor: 'grabbing', }, }, + converterCardFooter: { + display: 'flex', + alignItems: 'center', + flexWrap: 'wrap', + gap: tokens.spacingHorizontalS, + }, + removeConverterButton: { + ...mobileTouchTarget, + marginLeft: 'auto', + flexShrink: 0, + }, + settingsButton: { + ...mobileTouchTarget, + marginLeft: 'auto', + flexShrink: 0, + }, previewButton: { ...mobileTouchTarget, alignSelf: 'flex-start', @@ -248,11 +260,6 @@ export const useConverterPanelStyles = makeStyles({ fontSize: tokens.fontSizeBase100, fontWeight: tokens.fontWeightSemibold as unknown as string, verticalAlign: 'middle', - }, - optionBadges: { - display: 'flex', - alignItems: 'center', - gap: '2px', flexShrink: 0, }, touchTarget: { diff --git a/frontend/src/components/Chat/ConverterPanel/ConverterPanel.tsx b/frontend/src/components/Chat/ConverterPanel/ConverterPanel.tsx index b69cd1311b..052b8178b8 100644 --- a/frontend/src/components/Chat/ConverterPanel/ConverterPanel.tsx +++ b/frontend/src/components/Chat/ConverterPanel/ConverterPanel.tsx @@ -3,6 +3,11 @@ import type { DragEvent, KeyboardEvent, ReactNode } from 'react' import { Button, + Menu, + MenuTrigger, + MenuPopover, + MenuList, + MenuItem, MessageBar, MessageBarBody, Spinner, @@ -15,6 +20,7 @@ import { OpenRegular, PlayRegular, ReOrderDotsVerticalRegular, + MoreHorizontalRegular, } from '@fluentui/react-icons' import CreateConverterDialog from '@/components/Registry/CreateConverterDialog' @@ -23,6 +29,7 @@ import { convertersApi } from '@/services/api' import { toApiError } from '@/services/errors' import type { ChatConverterController, ConverterInputPiece, ConverterInstance, ConverterPipelineStage, ConverterStageResult, + SourceInstanceSpec, ConverterIdentifier, } from '@/types' import { @@ -175,6 +182,18 @@ export default function ConverterPanel({ const [isLoading, setIsLoading] = useState(true) const [error, setError] = useState(null) const [createDialogOpen, setCreateDialogOpen] = useState(false) + const [settings, setSettings] = useState<{ + stageId: string + pieceType: string + editing: { converter: ConverterInstance; spec?: SourceInstanceSpec } + } | null>(null) + const [settingsGeneration, setSettingsGeneration] = useState(generation) + if (settingsGeneration !== generation) { + setSettingsGeneration(generation) + setSettings(null) + } + const discardTemporarySettings = controller.discardTemporarySettings + useEffect(() => () => discardTemporarySettings?.(), [discardTemporarySettings]) const [panelWidth, setPanelWidth] = useState(DEFAULT_PANEL_WIDTH) const isResizing = useRef(false) const draggedConverterIndex = useRef(null) @@ -252,7 +271,7 @@ export default function ConverterPanel({ const selectedConverters = useMemo( () => selectedStages.flatMap((stage: ConverterPipelineStage): SelectedConverter[] => { const converter = converters.find((candidate: ConverterInstance) => candidate.converter_id === stage.converterId) - return converter ? [{ ...converter, stageId: stage.id }] : [] + return converter ? [{ ...converter, identifier: stage.identifier ?? converter.identifier, stageId: stage.id }] : [] }), [converters, selectedStages], ) @@ -522,6 +541,7 @@ export default function ConverterPanel({ {converter.converter_id} {converter.is_llm_based && LLM} + {selectedStages[index]?.temporary && Temporary}
{converter.identifier.class_name !== converter.converter_id && ( @@ -565,22 +585,53 @@ export default function ConverterPanel({ ? (value: string) => controller.editStageOutput(input.id, converter.stageId, value) : undefined} /> - {hasRemaining && ( - - )}
) })} +
+ {!isBatch && index < selectedConverters.length - 1 && activeInputs.map((input: ConverterInputPiece) => ( + + ))} + + + +
) })} @@ -635,6 +686,18 @@ export default function ConverterPanel({ void loadConverters(converterId, effectiveActiveTab) }} /> + {settings && setSettings(null)} + onCreated={() => { throw new Error('Settings must not register a converter.') }} + onTemporary={(spec: SourceInstanceSpec, identifier: ConverterIdentifier) => { + setPipeline(settings.pieceType, (stages: ConverterPipelineStage[]) => stages.map((stage) => + stage.id === settings.stageId ? { ...stage, temporary: spec, identifier } : stage, + )) + setSettings(null) + }} + />} ) } diff --git a/frontend/src/components/Chat/ConverterPanel/SelectConverterInput.tsx b/frontend/src/components/Chat/ConverterPanel/SelectConverterInput.tsx index e34809b77b..4808b1eeed 100644 --- a/frontend/src/components/Chat/ConverterPanel/SelectConverterInput.tsx +++ b/frontend/src/components/Chat/ConverterPanel/SelectConverterInput.tsx @@ -85,12 +85,7 @@ export default function SelectConverterInput({
{converter.converter_id} -
- - {converter.identifier.class_name} - - {converter.is_llm_based && LLM} -
+ {converter.is_llm_based && LLM}
{description} diff --git a/frontend/src/components/Chat/converterTypes.test.ts b/frontend/src/components/Chat/converterTypes.test.ts index ce8c4df487..98267b7921 100644 --- a/frontend/src/components/Chat/converterTypes.test.ts +++ b/frontend/src/components/Chat/converterTypes.test.ts @@ -40,6 +40,21 @@ function makeConversion( } describe('converter draft mapping', () => { + it('requires applying private settings before independent repeated conversion', () => { + const inputs = buildConverterInputs('text', []) + const pipelines = { + text: [{ + id: 'private', converterId: 'registered', + temporary: { source_name: 'registered', source_hash: 'source', params: { language: 'French' } }, + }], + } + expect(() => buildRequestConverterConfigurations(inputs, ['text'], pipelines, {})) + .toThrow('Apply temporary converter settings with Add converted value') + expect(buildRequestConverterConfigurations(inputs, ['text'], pipelines, { + text: makeConversion('text', ['private']), + })).toEqual([]) + }) + it('builds repeat pipelines in piece and stage order without rerunning applied previews', () => { const inputs = buildConverterInputs('text', [ { draftId: 'first', type: 'image', name: 'first.png', url: 'first.png', mimeType: 'image/png' }, diff --git a/frontend/src/components/Chat/converterTypes.ts b/frontend/src/components/Chat/converterTypes.ts index f063e325cb..76370868ce 100644 --- a/frontend/src/components/Chat/converterTypes.ts +++ b/frontend/src/components/Chat/converterTypes.ts @@ -69,6 +69,7 @@ export function applyConvertedValues( converted_value: conversion.convertedValue, converted_value_data_type: conversion.convertedDataType, applied_converter_ids: conversion.converterInstanceIds, + ...(conversion.converterProvenance ? { applied_converter_provenance: conversion.converterProvenance } : {}), } : piece }) } @@ -85,6 +86,9 @@ export function buildRequestConverterConfigurations( const input = inputs.find((candidate: ConverterInputPiece) => candidate.id === pieceId) if (!input) throw new Error('Message piece has no matching converter input.') const stages = pipelines[input.pieceType] ?? [] + if (stages.some((stage: ConverterPipelineStage) => stage.temporary)) { + throw new Error('Apply temporary converter settings with Add converted value before repeating a message.') + } return stages.length ? [{ converter_ids: stages.map((stage: ConverterPipelineStage) => stage.converterId), indexes_to_apply: [index], diff --git a/frontend/src/components/Layout/MainLayout.styles.ts b/frontend/src/components/Layout/MainLayout.styles.ts index 94098fd183..40393860ff 100644 --- a/frontend/src/components/Layout/MainLayout.styles.ts +++ b/frontend/src/components/Layout/MainLayout.styles.ts @@ -147,8 +147,15 @@ export const useMainLayoutStyles = makeStyles({ }, }, toolbarSlot: { + display: 'flex', + alignItems: 'center', + flexShrink: 0, + minWidth: 0, + minHeight: '48px', maxWidth: '100%', - marginLeft: 'auto', + padding: `${tokens.spacingVerticalXS} ${tokens.spacingHorizontalL}`, + borderBottom: `1px solid ${tokens.colorNeutralStroke1}`, + backgroundColor: tokens.colorNeutralBackground3, ':empty': { display: 'none', }, diff --git a/frontend/src/components/Layout/MainLayout.tsx b/frontend/src/components/Layout/MainLayout.tsx index e45b611f91..dcb276e82c 100644 --- a/frontend/src/components/Layout/MainLayout.tsx +++ b/frontend/src/components/Layout/MainLayout.tsx @@ -137,9 +137,9 @@ export default function MainLayout({
-
+
{children}
diff --git a/frontend/src/components/Registry/CreateConverterDialog.test.tsx b/frontend/src/components/Registry/CreateConverterDialog.test.tsx index cb1852f591..4cf6969944 100644 --- a/frontend/src/components/Registry/CreateConverterDialog.test.tsx +++ b/frontend/src/components/Registry/CreateConverterDialog.test.tsx @@ -12,6 +12,7 @@ jest.mock('@/services/api', () => ({ listConverterTypes: jest.fn(), listConverters: jest.fn(), createConverter: jest.fn(), + buildConverter: jest.fn(), }, targetsApi: { listTargets: jest.fn(), @@ -134,6 +135,35 @@ describe('CreateConverterDialog', () => { jest.restoreAllMocks() }) + it('builds private stage settings without changing or registering the source', async () => { + const user = userEvent.setup() + const onTemporary = jest.fn() + const converter = { + converter_id: 'source', + identifier: { + class_name: 'CaesarConverter', class_module: 'pyrit.converter', + pyrit_version: 'test', hash: 'source-hash', caesar_offset: 1, + }, + } + mockedConvertersApi.buildConverter.mockResolvedValue({ + identifier: { ...converter.identifier, hash: 'private-hash', caesar_offset: 3 }, + }) + renderDialog({ editing: { converter }, onTemporary }) + const offset = await screen.findByLabelText(/caesar_offset/i) + expect(offset).toHaveValue('1') + expect(screen.queryByLabelText(/registry name/i)).not.toBeInTheDocument() + await user.clear(offset) + await user.type(offset, '3') + await user.click(screen.getByRole('button', { name: 'Apply Settings' })) + await waitFor(() => expect(onTemporary).toHaveBeenCalledTimes(1)) + expect(mockedConvertersApi.buildConverter).toHaveBeenCalledWith('CaesarConverter', { + source_name: 'source', source_hash: 'source-hash', params: { caesar_offset: 3 }, + }) + expect(mockedConvertersApi.createConverter).not.toHaveBeenCalled() + expect(converter.identifier.caesar_offset).toBe(1) + expect(onTemporary.mock.calls[0][1].hash).toBe('private-hash') + }) + it('loads converter classes from registry type metadata', async () => { const user = userEvent.setup() renderDialog() @@ -146,6 +176,109 @@ describe('CreateConverterDialog', () => { expect(mockedConvertersApi.listConverterTypes).toHaveBeenCalledTimes(1) }) + it('should clear a prior structured override without clearing untouched settings', async () => { + const user = userEvent.setup() + mockConverterParameters([wordSelectionParameter, { + name: 'prefix', type_name: 'str', required: false, + }]) + const converter = { + converter_id: 'source', + identifier: { + class_name: 'TextConverter', class_module: 'pyrit.converter', + pyrit_version: 'test', hash: 'source-hash', + }, + } + const params = { + word_selection_strategy: { type: 'random', parameters: { proportion: 0.3, seed: 42 } }, + prefix: 'Keep this override', + } + const onTemporary = jest.fn() + mockedConvertersApi.buildConverter.mockResolvedValue({ identifier: converter.identifier }) + renderDialog({ + editing: { converter, spec: { source_name: 'source', source_hash: 'source-hash', params } }, + onTemporary, + }) + const strategy = await screen.findByRole('combobox', { name: 'word_selection_strategy' }) + expect(strategy).toHaveValue('random') + await user.selectOptions(strategy, '') + await user.click(screen.getByRole('button', { name: 'Apply Settings' })) + await waitFor(() => expect(onTemporary).toHaveBeenCalledTimes(1)) + expect(mockedConvertersApi.buildConverter).toHaveBeenCalledWith('TextConverter', { + source_name: 'source', source_hash: 'source-hash', params: { prefix: 'Keep this override' }, + }) + expect(params.word_selection_strategy.type).toBe('random') + expect(mockedConvertersApi.createConverter).not.toHaveBeenCalled() + }) + + it('should never register settings when the temporary apply handler is missing', async () => { + const user = userEvent.setup() + renderDialog({ + editing: { + converter: { + converter_id: 'source', + identifier: { + class_name: 'CaesarConverter', class_module: 'pyrit.converter', + pyrit_version: 'test', hash: 'source-hash', caesar_offset: 1, + }, + }, + }, + }) + await screen.findByLabelText(/caesar_offset/i) + await user.click(screen.getByRole('button', { name: 'Apply Settings' })) + expect(await screen.findByRole('alert')).toHaveTextContent(/need an apply handler/i) + expect(mockedConvertersApi.createConverter).not.toHaveBeenCalled() + expect(mockedConvertersApi.buildConverter).not.toHaveBeenCalled() + }) + + it.each([false, true])( + 'should apply an upload only to its original settings dialog (closed: %s)', + async (closed) => { + const user = userEvent.setup() + mockConverterParameters([ + { name: 'source', type_name: 'Path', required: true }, + ], 'PathConverter') + const converter = { + converter_id: 'source', + identifier: { + class_name: 'PathConverter', class_module: 'pyrit.converter', + pyrit_version: 'test', hash: 'source-hash', source: 'original.txt', + }, + } + const editing = { converter } + const onTemporary = jest.fn() + mockedConvertersApi.buildConverter.mockResolvedValue({ identifier: converter.identifier }) + const click = jest.spyOn(HTMLInputElement.prototype, 'click').mockImplementation(() => undefined) + const read = jest.spyOn(FileReader.prototype, 'readAsDataURL').mockImplementation(() => undefined) + const { rerender } = renderDialog({ editing, onTemporary }) + await screen.findByLabelText('source *') + await user.click(screen.getByRole('button', { name: 'Upload' })) + const fileInput = click.mock.contexts[0] + if (!(fileInput instanceof HTMLInputElement)) throw new Error('File picker did not open') + screen.getByRole('dialog').appendChild(fileInput) + await user.upload(fileInput, new File(['uploaded'], 'input.txt', { type: 'text/plain' })) + fileInput.remove() + const reader = read.mock.contexts[0] + if (!(reader instanceof FileReader)) throw new Error('File read did not start') + if (closed) { + rerender(dialogTree({ open: false, editing, onTemporary })) + rerender(dialogTree({ open: true, editing, onTemporary })) + await screen.findByLabelText('source *') + } + const upload = 'data:text/plain;base64,dXBsb2FkZWQ=' + act(() => { + Object.defineProperty(reader, 'result', { value: upload }) + reader.dispatchEvent(new ProgressEvent('load')) + }) + expect(screen.getByLabelText('source *')).toHaveValue(closed ? 'original.txt' : upload) + await user.click(screen.getByRole('button', { name: 'Apply Settings' })) + await waitFor(() => expect(onTemporary).toHaveBeenCalledTimes(1)) + expect(mockedConvertersApi.buildConverter).toHaveBeenCalledWith('PathConverter', { + source_name: 'source', source_hash: 'source-hash', params: closed ? {} : { source: upload }, + }) + expect(mockedConvertersApi.createConverter).not.toHaveBeenCalled() + }, + ) + it('hides converter types with required parameters the form cannot configure', async () => { mockedConvertersApi.listConverterTypes.mockResolvedValue({ items: [ diff --git a/frontend/src/components/Registry/CreateConverterDialog.tsx b/frontend/src/components/Registry/CreateConverterDialog.tsx index 9d6c56ff23..53e93f39c3 100644 --- a/frontend/src/components/Registry/CreateConverterDialog.tsx +++ b/frontend/src/components/Registry/CreateConverterDialog.tsx @@ -23,7 +23,9 @@ import { import { convertersApi, targetsApi } from '@/services/api' import { toApiError } from '@/services/errors' -import type { ConverterInstance, ConverterTypeEntry, Parameter, TargetInstance } from '@/types' +import type { + ConverterInstance, ConverterTypeEntry, Parameter, TargetInstance, SourceInstanceSpec, ConverterIdentifier, +} from '@/types' import ParameterField from '@/components/Parameters/ParameterField' import { buildParametersFromForm, @@ -67,6 +69,8 @@ interface CreateConverterDialogProps { open: boolean onClose: () => void onCreated: (converterId: string) => void + editing?: { converter: ConverterInstance; spec?: SourceInstanceSpec } + onTemporary?: (spec: SourceInstanceSpec, identifier: ConverterIdentifier) => void } interface ParameterInputProps { @@ -250,6 +254,8 @@ export default function CreateConverterDialog({ open, onClose, onCreated, + editing, + onTemporary, }: CreateConverterDialogProps) { const styles = useCreateConverterDialogStyles() const [converterTypes, setConverterTypes] = useState([]) @@ -259,6 +265,7 @@ export default function CreateConverterDialog({ const [registryName, setRegistryName] = useState('') const [nameEdited, setNameEdited] = useState(false) const [parameterValues, setParameterValues] = useState>({}) + const [changedParameters, setChangedParameters] = useState>(new Set()) const [loading, setLoading] = useState(false) const [submitting, setSubmitting] = useState(false) const [showValidation, setShowValidation] = useState(false) @@ -292,11 +299,19 @@ export default function CreateConverterDialog({ if (!responses) return const [response, targetResponse, converterResponse] = responses if (!cancelled) { - setConverterTypes( - response.items.filter(canConfigureConverterType), - ) + setConverterTypes(editing ? response.items : response.items.filter(canConfigureConverterType)) setTargets(targetResponse.items) setConverters(converterResponse.items) + if (editing) { + const entry = response.items.find((item) => item.converter_type === editing.converter.identifier.class_name) + setSelectedType(editing.converter.identifier.class_name) + const params: Record = { + ...editing.converter.identifier, + ...editing.spec?.params, + } + setParameterValues(getInitialFormValues(entry?.parameters ?? [], params, { prefillDefaults: false })) + setChangedParameters(new Set()) + } } }) .catch((err) => { @@ -310,8 +325,8 @@ export default function CreateConverterDialog({ .finally(() => { if (!cancelled) setLoading(false) }) - return () => { cancelled = true } - }, [open]) + return () => { cancelled = true; openEpochRef.current += 1 } + }, [open, editing]) // Hand the keyboard to the failure once React has committed it. A frame // callback can run before the render that adds the message bar, and focusing @@ -362,6 +377,7 @@ export default function CreateConverterDialog({ setRegistryName('') setNameEdited(false) setParameterValues({}) + setChangedParameters(new Set()) setShowValidation(false) setError(null) } @@ -390,13 +406,17 @@ export default function CreateConverterDialog({ } const browse = (parameterName: string) => { + const epoch = openEpochRef.current const input = document.createElement('input') input.type = 'file' input.onchange = () => { + if (openEpochRef.current !== epoch) return const file = input.files?.[0] if (!file) return const reader = new FileReader() reader.onload = () => { + if (openEpochRef.current !== epoch) return + setChangedParameters((current) => new Set([...current, parameterName])) setParameterValues((current) => ({ ...current, [parameterName]: String(reader.result ?? ''), @@ -413,12 +433,14 @@ export default function CreateConverterDialog({ && !parameter.default && !formValueIsSet(parameterValues[parameter.name]), ) - if (!selectedType || !registryName.trim() || missingParameters) { + if (!selectedType || (!editing && (!registryName.trim() || missingParameters))) { setShowValidation(true) return } - const parameters = selectedConverterType?.parameters ?? [] + const parameters = (selectedConverterType?.parameters ?? []).filter( + (parameter) => !editing || changedParameters.has(parameter.name), + ) const params: Record = Object.fromEntries( parameters .filter((parameter) => !sendsTypedValue(parameter)) @@ -442,6 +464,21 @@ export default function CreateConverterDialog({ setSubmitting(true) setError(null) try { + if (editing) { + if (!onTemporary) throw new Error('Temporary converter settings need an apply handler.') + const spec: SourceInstanceSpec = { + source_name: editing.spec?.source_name ?? editing.converter.converter_id, + source_hash: editing.spec?.source_hash ?? editing.converter.identifier.hash, + params: { + ...Object.fromEntries(Object.entries(editing.spec?.params ?? {}) + .filter(([name]: [string, unknown]) => !changedParameters.has(name))), + ...params, + }, + } + const response = await convertersApi.buildConverter(selectedType, spec) + if (openEpochRef.current === epoch) onTemporary(spec, response.identifier) + return + } const response = await convertersApi.createConverter({ name: registryName.trim(), type: selectedType, @@ -478,7 +515,7 @@ export default function CreateConverterDialog({ { if (!data.open) close() }}> - Add Converter + {editing ? 'Converter Settings' : 'Add Converter'}
)} {loading && } + {editing && + Only changed fields replace source settings. These settings apply to this stage, + not the registered converter. Closing the converter pane discards them. + } {!loading && converterTypes.length === 0 && !error && ( No converter types are available. )} @@ -505,6 +546,7 @@ export default function CreateConverterDialog({ validationMessage={showValidation && !selectedType ? 'Select a converter type' : undefined} >
)} - - + } + {editing && Changes apply only to this stage. The registered converter stays unchanged.}
{selectedConverterType?.parameters.map((parameter) => (
@@ -598,10 +641,10 @@ export default function CreateConverterDialog({ disabled={submitting} showRequiredError={showValidation && parameter.required} testIdPrefix="structured" - onChange={(name, value) => setParameterValues((current) => ({ - ...current, - [name]: value, - }))} + onChange={(name, value) => { + setChangedParameters((current) => new Set([...current, name])) + setParameterValues((current) => ({ ...current, [name]: value })) + }} /> ) : setParameterValues((current) => ({ - ...current, - [parameter.name]: value, - }))} + onChange={(value) => { + setChangedParameters((current) => new Set([...current, parameter.name])) + setParameterValues((current) => ({ ...current, [parameter.name]: value })) + }} onBrowse={() => browse(parameter.name)} />}
@@ -634,7 +677,7 @@ export default function CreateConverterDialog({ disabledFocusable={submitDisabled} onClick={() => void submit()} > - {submitting ? 'Adding...' : 'Add Converter'} + {submitting ? editing ? 'Saving...' : 'Adding...' : editing ? 'Apply Settings' : 'Add Converter'} diff --git a/frontend/src/hooks/useAttackTargetResolution.test.tsx b/frontend/src/hooks/useAttackTargetResolution.test.tsx index 44f0d661d7..b92413006e 100644 --- a/frontend/src/hooks/useAttackTargetResolution.test.tsx +++ b/frontend/src/hooks/useAttackTargetResolution.test.tsx @@ -12,7 +12,7 @@ jest.mock('@/hooks/useRuntime', () => ({ })) jest.mock('@/services/api', () => ({ - targetsApi: { getTarget: jest.fn() }, + targetsApi: { getTarget: jest.fn(), buildTarget: jest.fn() }, })) jest.mock('@/services/targetRegistry', () => ({ @@ -40,6 +40,36 @@ describe('useAttackTargetResolution', () => { jest.mocked(targetsApi.getTarget).mockResolvedValue(replacementTarget) }) + it('reconstructs a saved temperature without selecting the registered default', async () => { + const binding = { + version: 1 as const, source_name: 'target', source_hash: 'source-hash', + temperature: 0.8, effective_hash: 'effective-hash', + } + jest.mocked(targetsApi.buildTarget).mockResolvedValue({ + identifier: { class_name: 'OpenAIChatTarget', hash: 'effective-hash', temperature: 0.8 }, + }) + const savedTarget: TargetInfo = { + ...targetInfo, target_type: 'OpenAIChatTarget', identifier_hash: 'effective-hash', binding, + } + const { result, rerender } = renderHook(() => useAttackTargetResolution({ + attackId: 'saved', attackLoadSequence: 1, + attackTarget: savedTarget, + attackTargetSource: 'persisted', + })) + await waitFor(() => expect(result.current.activeTarget?.binding).toEqual(binding)) + expect(targetsApi.getTarget).not.toHaveBeenCalled() + expect(targetsApi.buildTarget).toHaveBeenCalledWith('OpenAIChatTarget', { + source_name: 'target', source_hash: 'source-hash', params: { temperature: 0.8 }, + effective_hash: 'effective-hash', + }) + jest.mocked(targetsApi.buildTarget).mockRejectedValue(new Error('Source missing')) + mockRuntimeGeneration = 'generation-2' + rerender() + expect(result.current.activeTarget).toBeNull() + await waitFor(() => expect(result.current.resolutionStatus).toBe('error')) + expect(result.current.activeTarget).toBeNull() + }) + it('resolves a created attack from the new registry after the runtime generation changes', async () => { const { result, rerender } = renderHook(() => useAttackTargetResolution({ attackId: 'attack-id', diff --git a/frontend/src/hooks/useAttackTargetResolution.ts b/frontend/src/hooks/useAttackTargetResolution.ts index 935b8af2dc..8ea572e510 100644 --- a/frontend/src/hooks/useAttackTargetResolution.ts +++ b/frontend/src/hooks/useAttackTargetResolution.ts @@ -18,6 +18,7 @@ import type { TargetHashResolution } from '@/utils/targetIdentity' interface RegistryResolution { attackId: string | null attackLoadSequence: number + generation: string status: 'idle' | 'resolved' | 'unavailable' | 'ambiguous' | 'error' target?: TargetInstance } @@ -47,6 +48,21 @@ function hasCompleteIdentifier(target: TargetInfo | null): target is TargetInfo } async function resolvePersistedTarget(target: TargetInfo): Promise { + if (target.binding) { + const source = await targetsApi.buildTarget(target.target_type, { + source_name: target.binding.source_name, + source_hash: target.binding.source_hash, + params: { temperature: target.binding.temperature }, + effective_hash: target.binding.effective_hash, + }) + return { + status: 'resolved', + target: { + ...source, target_registry_name: target.binding.source_name, + binding: target.binding, reconstructable: true, supports_temperature_override: true, + }, + } + } if (target.target_registry_name) { try { const namedTarget = await targetsApi.getTarget(target.target_registry_name) @@ -80,6 +96,7 @@ export function useAttackTargetResolution({ const [registryResolution, setRegistryResolution] = useState({ attackId: null, attackLoadSequence: 0, + generation, status: 'idle', }) const [resolutionAttempt, setResolutionAttempt] = useState(0) @@ -98,15 +115,16 @@ export function useAttackTargetResolution({ setRegistryResolution({ attackId, attackLoadSequence, + generation, status: 'resolved', target: resolution.target, }) return } - setRegistryResolution({ attackId, attackLoadSequence, status: resolution.status }) + setRegistryResolution({ attackId, attackLoadSequence, generation, status: resolution.status }) } catch { if (cancelled) return - setRegistryResolution({ attackId, attackLoadSequence, status: 'error' }) + setRegistryResolution({ attackId, attackLoadSequence, generation, status: 'error' }) } } @@ -128,6 +146,7 @@ export function useAttackTargetResolution({ if ( registryResolution.attackId !== attackId || registryResolution.attackLoadSequence !== attackLoadSequence + || registryResolution.generation !== generation ) return 'loading' return registryResolution.status } @@ -137,9 +156,9 @@ export function useAttackTargetResolution({ : null const retryResolution = useCallback((): void => { - setRegistryResolution({ attackId: null, attackLoadSequence: 0, status: 'idle' }) + setRegistryResolution({ attackId: null, attackLoadSequence: 0, generation, status: 'idle' }) setResolutionAttempt((attempt) => attempt + 1) - }, []) + }, [generation]) return { activeTarget, diff --git a/frontend/src/hooks/useChatConverters.test.tsx b/frontend/src/hooks/useChatConverters.test.tsx index fec26dba79..fa6825e778 100644 --- a/frontend/src/hooks/useChatConverters.test.tsx +++ b/frontend/src/hooks/useChatConverters.test.tsx @@ -50,6 +50,64 @@ function makePreviewResponse() { } describe('useChatConverters across runtime generation changes', () => { + const privateSpec = { + source_name: 'base64-default', source_hash: 'source-hash', params: { example: 3 }, + } + + it('sends aligned private recipes and preserves applied provenance after the pane closes', async () => { + mockedPreview.mockResolvedValue({ + ...makePreviewResponse(), + converted_value: 'private output', converted_value_data_type: 'text', + steps: [{ ...makePreviewResponse().steps[0], output_value: 'private output', provenance: 'signed-evidence' }], + }) + const { result } = renderHook(() => useChatConverters('original text', NO_ATTACHMENTS)) + act(() => result.current.setPipeline('text', () => [ + { id: 'first', converterId: 'base64-default', temporary: privateSpec }, + ])) + await act(async () => { await result.current.convert('text') }) + expect(mockedPreview.mock.calls[0][0].converter_specs).toEqual([privateSpec]) + act(() => result.current.apply()) + expect(result.current.applied.text.converterProvenance).toEqual(['signed-evidence']) + act(() => result.current.discardTemporarySettings?.()) + expect(result.current.pipelines.text).toEqual([{ id: 'first', converterId: 'base64-default' }]) + expect(result.current.applied.text.convertedValue).toBe('private output') + expect(result.current.applied.text.converterProvenance).toEqual(['signed-evidence']) + expect(result.current.stageResults.text).toEqual([]) + }) + + it('discards private settings on runtime replacement and rejects stale preview results', async () => { + let finish: (value: ReturnType) => void = () => {} + mockedPreview.mockImplementation(() => new Promise((resolve) => { finish = resolve })) + const { result, rerender } = renderHook(() => useChatConverters('original text', NO_ATTACHMENTS)) + act(() => result.current.setPipeline('text', () => [ + { id: 'first', converterId: 'base64-default', temporary: privateSpec }, + ])) + let conversion: Promise | undefined + act(() => { conversion = result.current.convert('text') }) + runtime.generation = 'gen-2' + rerender() + await act(async () => { finish(makePreviewResponse()); await conversion }) + expect(result.current.pipelines.text).toEqual([{ id: 'first', converterId: 'base64-default' }]) + expect(result.current.stageResults.text).toBeUndefined() + expect(result.current.isConverting).toBe(false) + }) + + it('changes only the selected duplicate stage and invalidates its downstream results', async () => { + mockedPreview.mockResolvedValue({ + ...makePreviewResponse(), steps: [makePreviewResponse().steps[0], makePreviewResponse().steps[0]], + }) + const { result } = renderHook(() => useChatConverters('original text', NO_ATTACHMENTS)) + act(() => result.current.setPipeline('text', () => [ + { id: 'first', converterId: 'base64-default' }, { id: 'second', converterId: 'base64-default' }, + ])) + await act(async () => { await result.current.convert('text') }) + act(() => result.current.setPipeline('text', (stages) => stages.map((stage) => + stage.id === 'second' ? { ...stage, temporary: privateSpec } : stage))) + expect(result.current.pipelines.text[0].temporary).toBeUndefined() + expect(result.current.pipelines.text[1].temporary).toEqual(privateSpec) + expect(result.current.stageResults.text).toHaveLength(1) + }) + it('restores repeat pipelines and exact applied values without retaining unrelated selections', () => { const { result } = renderHook(() => useChatConverters('original text', NO_ATTACHMENTS)) const pipelines = { diff --git a/frontend/src/hooks/useChatConverters.ts b/frontend/src/hooks/useChatConverters.ts index 62fdeee3a9..7287f76717 100644 --- a/frontend/src/hooks/useChatConverters.ts +++ b/frontend/src/hooks/useChatConverters.ts @@ -91,6 +91,7 @@ function changePipeline(state: ConversionState, pieceType: string, stages: Conve prefixLength < previous.length && prefixLength < stages.length && previous[prefixLength].id === stages[prefixLength].id && previous[prefixLength].converterId === stages[prefixLength].converterId + && previous[prefixLength].temporary === stages[prefixLength].temporary ) prefixLength++ if (prefixLength === previous.length && prefixLength === stages.length) return state @@ -183,6 +184,9 @@ export function usePieceConverters(inputs: ConverterInputPiece[], scopeKey?: str if (state.scopeKey !== scopeKey) { setState({ ...reconcileInputs(state, inputs), scopeKey, generation, stageResults: {}, errors: {}, applied: {}, + pipelines: Object.fromEntries(Object.entries(state.pipelines).map(([pieceType, stages]) => [ + pieceType, stages.map((stage: ConverterPipelineStage) => ({ id: stage.id, converterId: stage.converterId })), + ])), workingInputs: {}, runId: state.runId + 1, isConverting: false, }) } else if (state.generation !== generation) { @@ -194,6 +198,9 @@ export function usePieceConverters(inputs: ConverterInputPiece[], scopeKey?: str setState({ ...next, generation, + pipelines: Object.fromEntries(Object.entries(next.pipelines).map(([pieceType, stages]) => [ + pieceType, stages.map((stage: ConverterPipelineStage) => ({ id: stage.id, converterId: stage.converterId })), + ])), stageResults: {}, errors: {}, applied: {}, @@ -232,6 +239,19 @@ export function usePieceConverters(inputs: ConverterInputPiece[], scopeKey?: str }) }, []) + const discardTemporarySettings = useCallback((): void => { + activeRun.current = null + setState((current: ConversionState) => { + let next = current + for (const [pieceType, stages] of Object.entries(current.pipelines)) { + next = changePipeline(next, pieceType, stages.map((stage: ConverterPipelineStage) => ({ + id: stage.id, converterId: stage.converterId, + }))) + } + return { ...next, applied: current.applied, runId: current.runId + 1, isConverting: false } + }) + }, []) + const editInput = useCallback((pieceId: string, value: string): void => { setState((current: ConversionState) => { const input = current.inputs.find((candidate: VersionedInput) => candidate.id === pieceId) @@ -316,6 +336,9 @@ export function usePieceConverters(inputs: ConverterInputPiece[], scopeKey?: str original_value: requestValue, original_value_data_type: dataType, converter_ids: remaining.map((stage: ConverterPipelineStage) => stage.converterId), + ...(remaining.some((stage: ConverterPipelineStage) => stage.temporary) ? { + converter_specs: remaining.map((stage: ConverterPipelineStage) => stage.temporary ?? null), + } : {}), }) if (response.steps.length !== remaining.length || response.steps.some( (step: ConverterPreviewStep, index: number) => step.converter_id !== remaining[index].converterId, @@ -420,6 +443,7 @@ export function usePieceConverters(inputs: ConverterInputPiece[], scopeKey?: str addConverter, setPipeline, retainConverters, + discardTemporarySettings, convert: (pieceType: string) => runConversion({ pieceType, includeIncomplete: true }), convertRemaining: (pieceId: string, stageId: string) => runConversion({ pieceId, afterStageId: stageId }), editInput, diff --git a/frontend/src/hooks/useConversationDraft.test.ts b/frontend/src/hooks/useConversationDraft.test.ts index 05cb293967..6cea16e9b0 100644 --- a/frontend/src/hooks/useConversationDraft.test.ts +++ b/frontend/src/hooks/useConversationDraft.test.ts @@ -1,12 +1,14 @@ import { act, renderHook } from '@testing-library/react' -import { attacksApi } from '@/services/api' +import { attacksApi, targetsApi } from '@/services/api' import { makeTarget } from '@/test-utils/targetFixtures' -import type { AddMessageResponse, ConversationSaveInput } from '@/types' +import type { AddMessageResponse, ConversationSaveInput, NewAttackContext, TargetInstance } from '@/types' import { useConversationDraft } from './useConversationDraft' -jest.mock('@/services/api', () => ({ attacksApi: { saveConversation: jest.fn() } })) +jest.mock('@/services/api', () => ({ + attacksApi: { saveConversation: jest.fn() }, targetsApi: { buildTarget: jest.fn() }, +})) const initial: ConversationSaveInput = { messages: [{ id: 'message', role: 'user', pieces: [{ draftId: 'piece', data_type: 'text', original_value: 'Prompt' }] }], @@ -17,6 +19,107 @@ const initial: ConversationSaveInput = { describe('useConversationDraft', () => { beforeEach(() => jest.resetAllMocks()) + const temperatureTarget: TargetInstance = { + ...makeTarget({ + target_registry_name: 'source', target_type: 'OpenAIChatTarget', temperature: 0.2, + capabilities: { supports_multi_turn: true, supports_editable_history: true, supported_input_modalities: ['text'] }, + }), + supports_temperature_override: true, + } + + it('saves a new attack with a private temperature and returns its target to chat', async () => { + jest.mocked(targetsApi.buildTarget).mockResolvedValue({ + identifier: { ...temperatureTarget.identifier, hash: 'effective-hash', temperature: 0.8 }, + capabilities: temperatureTarget.capabilities, + }) + jest.mocked(attacksApi.saveConversation).mockResolvedValue({ + attack: { + attack_result_id: 'saved', conversation_id: 'saved', attack_type: 'ManualAttack', objective: 'Objective', + converters: [], message_count: 1, related_conversation_ids: [], labels: {}, created_at: '', updated_at: '', + }, + messages: { conversation_id: 'saved', messages: [], target_response_status: null }, + }) + const onSaved = jest.fn() + const { result } = renderHook(() => useConversationDraft()) + act(() => result.current.begin({ ...initial, target: temperatureTarget })) + act(() => result.current.changeTemperature('0.8')) + await act(async () => { await result.current.save('new_attack', onSaved) }) + const binding = { + version: 1, source_name: 'source', source_hash: temperatureTarget.identifier.hash, + temperature: 0.8, effective_hash: 'effective-hash', + } + expect(jest.mocked(attacksApi.saveConversation).mock.calls[0][0].target_binding).toEqual(binding) + expect(onSaved.mock.calls[0][1].binding).toEqual(binding) + expect(temperatureTarget.identifier.temperature).toBe(0.2) + }) + + it('rejects temperature changes in the same attack and retains the draft', async () => { + const { result } = renderHook(() => useConversationDraft()) + act(() => result.current.begin({ ...initial, target: temperatureTarget })) + act(() => result.current.changeTemperature('0.8')) + await act(async () => { await result.current.save('same_attack', jest.fn()) }) + expect(result.current.error).toBe('Choose New attack to change the temperature.') + expect(targetsApi.buildTarget).not.toHaveBeenCalled() + expect(attacksApi.saveConversation).not.toHaveBeenCalled() + expect(result.current.draft?.temperature).toBe('0.8') + }) + + it.each([ + { generation: 'gen-1', ready: false }, + { generation: 'gen-2', ready: true }, + ])('preserves the draft when context becomes %j during temperature construction', async (next: NewAttackContext) => { + const built = { + identifier: { ...temperatureTarget.identifier, hash: 'effective-hash', temperature: 0.8 }, + capabilities: temperatureTarget.capabilities, + } + let finish: (value: typeof built) => void = () => {} + jest.mocked(targetsApi.buildTarget).mockImplementationOnce(() => new Promise((resolve) => { finish = resolve })) + jest.mocked(targetsApi.buildTarget).mockResolvedValue(built) + const onSaved = jest.fn() + const { result, rerender } = renderHook( + (context: NewAttackContext) => useConversationDraft(context), + { initialProps: { generation: 'gen-1', ready: true, labels: { operation: 'original' } } }, + ) + act(() => { + result.current.begin({ ...initial, target: temperatureTarget }) + result.current.changeTemperature('0.8') + }) + let pending: Promise + await act(async () => { pending = result.current.save('new_attack', onSaved) }) + rerender({ ...next, labels: { operation: 'replacement' } }) + await act(async () => { finish(built); await pending }) + expect(result.current.error).toContain('Runtime or default labels changed') + expect(result.current.draft?.temperature).toBe('0.8') + expect(result.current.saving).toBe(false) + expect(attacksApi.saveConversation).not.toHaveBeenCalled() + expect(onSaved).not.toHaveBeenCalled() + rerender({ generation: 'gen-2', ready: true, labels: { operation: 'replacement' } }) + jest.mocked(attacksApi.saveConversation).mockResolvedValue({ + attack: { + attack_result_id: 'saved', conversation_id: 'saved', attack_type: 'ManualAttack', objective: 'Objective', + converters: [], message_count: 1, related_conversation_ids: [], labels: {}, created_at: '', updated_at: '', + }, + messages: { conversation_id: 'saved', messages: [], target_response_status: null }, + }) + await act(async () => { await result.current.save('new_attack', onSaved) }) + expect(attacksApi.saveConversation).toHaveBeenCalledWith(expect.objectContaining({ + labels: { operation: 'replacement' }, target_binding: expect.objectContaining({ temperature: 0.8 }), + })) + expect(onSaved).toHaveBeenCalledTimes(1) + }) + + it('does not build a temperature target while defaults are loading', async () => { + const { result } = renderHook(() => useConversationDraft({ generation: 'gen-1', ready: false })) + act(() => { + result.current.begin({ ...initial, target: temperatureTarget }) + result.current.changeTemperature('0.8') + }) + await act(async () => { await result.current.save('new_attack', jest.fn()) }) + expect(result.current.error).toContain('Default labels are not ready') + expect(targetsApi.buildTarget).not.toHaveBeenCalled() + expect(attacksApi.saveConversation).not.toHaveBeenCalled() + }) + it.each(['same_attack', 'new_attack'] as const)( 'blocks %s for unsupported media and permits a targetless new attack', async (destination) => { diff --git a/frontend/src/hooks/useConversationDraft.ts b/frontend/src/hooks/useConversationDraft.ts index b7ac3c38a9..0edf86507b 100644 --- a/frontend/src/hooks/useConversationDraft.ts +++ b/frontend/src/hooks/useConversationDraft.ts @@ -1,6 +1,7 @@ import { useCallback, useEffect, useRef, useState } from 'react' import { toApiError } from '@/services/errors' +import { buildTemperatureTarget } from '@/services/targetRegistry' import type { AddMessageResponse, ConversationDraftMessage, ConversationDraftPiece, ConversationSaveInput, MessageAttachment, NewAttackContext, PieceConversion, SaveConversationRequest, TargetInstance, @@ -23,18 +24,21 @@ interface DraftState extends ConversationSaveInput { id: string baselineMessages: ConversationDraftMessage[] baselineTarget: string | undefined + temperature: string } export function useConversationDraft(newAttackContext?: NewAttackContext) { const [draft, setDraft] = useState(null) const [error, setError] = useState(null) + const [preparing, setPreparing] = useState(false) + const savePending = useRef(false) const activeId = useRef(null) const saved = useRef(false) const urls = useRef(new Set()) const workflow = useConversationSave(newAttackContext) const dirty = draft !== null && ( draft.messages !== draft.baselineMessages || draft.objective !== draft.initialObjective - || targetKey(draft.target) !== draft.baselineTarget + || targetKey(draft.target) !== draft.baselineTarget || draft.temperature.trim() !== '' ) const validationError = draft ? validateDraft(draft.messages) : null const targetError = draft?.target @@ -52,7 +56,7 @@ export function useConversationDraft(newAttackContext?: NewAttackContext) { activeId.current = id saved.current = false setError(null) - setDraft({ ...input, id, baselineMessages: input.messages, baselineTarget: targetKey(input.target) }) + setDraft({ ...input, id, baselineMessages: input.messages, baselineTarget: targetKey(input.target), temperature: '' }) }, [releaseUrls]) const discard = useCallback((): void => { activeId.current = null @@ -127,24 +131,41 @@ export function useConversationDraft(newAttackContext?: NewAttackContext) { ...piece, converted_value: result.convertedValue, converted_value_data_type: result.convertedDataType, applied_converter_ids: result.converterInstanceIds, + applied_converter_provenance: result.converterProvenance, } : piece }), }))) } const save = async ( - destination: SaveConversationRequest['destination'], onSaved: (response: AddMessageResponse) => void, + destination: SaveConversationRequest['destination'], onSaved: (response: AddMessageResponse, target?: TargetInstance | null) => void, ): Promise => { - if (!draft || workflow.saving) return + if (!draft || workflow.saving || savePending.current) return + savePending.current = true + setPreparing(true) const id = draft.id setError(null) try { if (validationError || targetError) throw new Error(validationError ?? targetError) - const response = await workflow.save(draft, destination) + if (destination === 'new_attack' && newAttackContext && !newAttackContext.ready) { + throw new Error('Default labels are not ready. Retry after default labels finish loading.') + } + let target = draft.target + if (draft.temperature.trim()) { + if (destination !== 'new_attack') throw new Error('Choose New attack to change the temperature.') + if (!target) throw new Error('Select a target before setting temperature.') + target = await buildTemperatureTarget(target, Number(draft.temperature)) + } + if (activeId.current !== id) return + const response = await workflow.save({ ...draft, target }, destination, newAttackContext) if (activeId.current !== id) return saved.current = true - onSaved(response) + if (draft.temperature.trim()) onSaved(response, target) + else onSaved(response) } catch (cause) { if (activeId.current === id) setError(toApiError(cause).detail) + } finally { + savePending.current = false + setPreparing(false) } } @@ -153,12 +174,15 @@ export function useConversationDraft(newAttackContext?: NewAttackContext) { }, []) return { - draft, dirty, error, validationError, targetError, saving: workflow.saving, + draft, dirty, error, validationError, targetError, saving: preparing || workflow.saving, begin, discard, save, changeMessages, changePiece, changeAttachments, addPiece, applyConversions, - shouldBlock: (): boolean => !saved.current && (dirty || workflow.saving), + shouldBlock: (): boolean => !saved.current && (dirty || preparing || workflow.saving), changeObjective, + changeTemperature: (temperature: string): void => { + setDraft((current: DraftState | null) => current ? { ...current, temperature } : current) + }, changeTarget: (target: TargetInstance | null): void => { - setDraft((current: DraftState | null) => current ? { ...current, target } : current) + setDraft((current: DraftState | null) => current ? { ...current, target, temperature: '' } : current) }, insert: (index: number): void => { const message = newDraftMessage() diff --git a/frontend/src/hooks/useConversationSave.ts b/frontend/src/hooks/useConversationSave.ts index 4fcc47c487..7d1b9b7a31 100644 --- a/frontend/src/hooks/useConversationSave.ts +++ b/frontend/src/hooks/useConversationSave.ts @@ -20,12 +20,13 @@ export function useConversationSave(newAttackContext?: NewAttackContext) { const save = useCallback(async ( input: ConversationSaveInput, destination: SaveConversationRequest['destination'], + preparationContext?: NewAttackContext, ): Promise => { if (pending.current) throw new Error('A conversation save is already in progress.') pending.current = true setSaving(true) try { - const initialContext = destination === 'new_attack' ? contextRef.current : undefined + const initialContext = destination === 'new_attack' ? preparationContext ?? contextRef.current : undefined if (initialContext && !initialContext.ready) { throw new Error('Default labels are not ready. Retry after default labels finish loading.') } @@ -45,6 +46,7 @@ export function useConversationSave(newAttackContext?: NewAttackContext) { expected_objective: updatesObjective ? input.initialObjective : undefined, objective: destination === 'new_attack' || updatesObjective ? objective : undefined, target_registry_name: input.target?.target_registry_name, + ...(input.target?.binding ? { target_binding: input.target.binding } : {}), operator: labels?.operator, operation: labels?.operation, labels, diff --git a/frontend/src/services/api.ts b/frontend/src/services/api.ts index 4e4cde69df..d5aeb96ed0 100644 --- a/frontend/src/services/api.ts +++ b/frontend/src/services/api.ts @@ -6,6 +6,9 @@ import { compatibility, COMPATIBILITY_HEADER } from './compatibility' import { getGraphScopes } from '../auth/msalConfig' import type { TargetInstance, + UnregisteredTarget, + UnregisteredConverter, + SourceInstanceSpec, TargetListResponse, TargetTypeListResponse, ConverterTypeListResponse, @@ -253,6 +256,10 @@ export const targetsApi = { const response = await apiClient.post('/targets', request) return response.data }, + buildTarget: async (type: string, source: SourceInstanceSpec): Promise => { + const response = await apiClient.post('/targets', { type, source, register: false }) + return response.data + }, } export const convertersApi = { @@ -275,6 +282,10 @@ export const convertersApi = { const response = await apiClient.post('/converters', request) return response.data }, + buildConverter: async (type: string, source: SourceInstanceSpec): Promise => { + const response = await apiClient.post('/converters', { type, source, register: false }) + return response.data + }, deleteConverter: async (converterId: string): Promise => { await apiClient.delete(`/converters/${encodeURIComponent(converterId)}`) diff --git a/frontend/src/services/targetRegistry.ts b/frontend/src/services/targetRegistry.ts index 7fb6eba0b5..6e5866c4bf 100644 --- a/frontend/src/services/targetRegistry.ts +++ b/frontend/src/services/targetRegistry.ts @@ -4,6 +4,27 @@ import type { TargetInstance } from '@/types' const TARGET_PAGE_SIZE = 200 const TARGET_MAX_PAGES = 100 +export async function buildTemperatureTarget(target: TargetInstance, temperature: number): Promise { + if (!Number.isFinite(temperature) || temperature < 0 || temperature > 2) { + throw new Error('Temperature must be between 0 and 2.') + } + if (!target.supports_temperature_override) { + throw new Error(target.reconstruction_error ?? 'This target does not support a separate temperature setting.') + } + const sourceName = target.binding?.source_name ?? target.target_registry_name + const sourceHash = target.binding?.source_hash ?? target.identifier.hash + const built = await targetsApi.buildTarget(target.identifier.class_name, { + source_name: sourceName, source_hash: sourceHash, params: { temperature }, + }) + return { + ...target, ...built, + binding: { + version: 1, source_name: sourceName, source_hash: sourceHash, + temperature, effective_hash: built.identifier.hash, + }, + } +} + /** Target identity resolution requires the complete registry, not a partial page set. */ export async function listRegisteredTargets(): Promise { const targets = new Map() diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index ea5333de68..e122ff4695 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -66,6 +66,7 @@ export interface PieceConversion { pieceId: string pieceType: string converterInstanceIds: string[] + converterProvenance?: Array convertedValue: string originalValue: string convertedDataType: string @@ -74,6 +75,8 @@ export interface PieceConversion { export interface ConverterPipelineStage { readonly id: string readonly converterId: string + readonly temporary?: SourceInstanceSpec + readonly identifier?: ConverterIdentifier } export interface ConverterStageResult { @@ -95,6 +98,7 @@ export interface ChatConverterController { addConverter: (pieceType: string, converterId: string) => void setPipeline: (pieceType: string, update: (stages: ConverterPipelineStage[]) => ConverterPipelineStage[]) => void retainConverters: (availableIds: Set) => void + discardTemporarySettings?: () => void convert: (pieceType: string) => Promise convertRemaining: (pieceId: string, stageId: string) => Promise editInput: (pieceId: string, value: string) => void @@ -273,6 +277,10 @@ export interface TargetIdentifier { export interface TargetInstance { target_registry_name: string + reconstructable?: boolean + reconstruction_error?: string | null + supports_temperature_override?: boolean + binding?: TargetBinding /** Typed identity: class name, endpoint, model name, generation params, content hash. */ identifier: TargetIdentifier capabilities?: TargetCapabilities | null @@ -351,6 +359,8 @@ export interface ConverterIdentifier { export interface ConverterInstance { converter_id: string + reconstructable?: boolean + reconstruction_error?: string | null identifier: ConverterIdentifier is_llm_based?: boolean description?: string | null @@ -366,6 +376,30 @@ export interface CreateConverterRequest { params?: Record } +export interface SourceInstanceSpec { + source_name: string + source_hash: string + params: Record + effective_hash?: string +} + +export interface TargetBinding { + version: 1 + source_name: string + source_hash: string + temperature: number + effective_hash: string +} + +export interface UnregisteredTarget { + identifier: TargetIdentifier + capabilities: TargetCapabilities +} + +export interface UnregisteredConverter { + identifier: ConverterIdentifier +} + export interface Parameter { name: string type_name: string @@ -402,6 +436,7 @@ export interface ConverterTypeListResponse { export interface ConverterPreviewRequest { original_value: string converter_ids: string[] + converter_specs?: Array original_value_data_type?: string start_token?: string end_token?: string @@ -414,6 +449,9 @@ export interface ConverterPreviewStep { converter_id: string input_data_type: string output_value: string output_data_type: string + source?: SourceInstanceSpec | null + identifier?: ConverterIdentifier | null + provenance?: string | null } export interface ConverterPreviewResponse { @@ -443,6 +481,7 @@ export interface TargetInfo { endpoint?: string | null model_name?: string | null identifier_hash: string + binding?: TargetBinding | null } export type AttackTargetResolutionStatus = @@ -489,6 +528,7 @@ export interface AttackSummary { export interface CreateAttackRequest { target_registry_name?: string + target_binding?: TargetBinding name?: string operator?: string operation?: string @@ -551,6 +591,7 @@ export interface SaveConversationRequest { expected_objective?: string objective?: string target_registry_name?: string + target_binding?: TargetBinding operator?: string operation?: string labels?: Record @@ -655,6 +696,7 @@ export interface MessagePieceRequest { converted_value?: string converted_value_data_type?: string applied_converter_ids?: string[] + applied_converter_provenance?: Array mime_type?: string original_prompt_id?: string source_piece_id?: string @@ -682,6 +724,7 @@ export interface ConverterConfigurationRequest { export interface AddMessageRequest extends MessageRequest { send: boolean target_registry_name?: string + target_binding?: TargetBinding converter_ids?: string[] request_converter_configurations?: ConverterConfigurationRequest[] response_converter_configurations?: ConverterConfigurationRequest[] diff --git a/frontend/src/utils/conversationDraft.ts b/frontend/src/utils/conversationDraft.ts index db4ad261b7..3d6efff79c 100644 --- a/frontend/src/utils/conversationDraft.ts +++ b/frontend/src/utils/conversationDraft.ts @@ -169,6 +169,7 @@ export async function serializeDraft(messages: ConversationDraftMessage[]): Prom converted_value: piece.converted_value, converted_value_data_type: piece.converted_value_data_type, applied_converter_ids: piece.converted_value === undefined ? undefined : piece.applied_converter_ids, + applied_converter_provenance: piece.converted_value === undefined ? undefined : piece.applied_converter_provenance, source_piece_id: piece.source_piece_id, mime_type: piece.mime_type, prompt_metadata: piece.prompt_metadata, diff --git a/frontend/src/utils/conversionResults.ts b/frontend/src/utils/conversionResults.ts index e87807ccad..be991048e4 100644 --- a/frontend/src/utils/conversionResults.ts +++ b/frontend/src/utils/conversionResults.ts @@ -13,6 +13,9 @@ export function buildAppliedConversions( pieceId: input.id, pieceType: input.pieceType, converterInstanceIds: result.steps.map((step: ConverterPreviewStep) => step.converter_id), + ...(result.steps.some((step: ConverterPreviewStep) => step.provenance) ? { + converterProvenance: result.steps.map((step: ConverterPreviewStep) => step.provenance ?? null), + } : {}), originalValue: input.value, convertedValue: result.converted_value, convertedDataType: result.converted_value_data_type, diff --git a/pyrit/backend/mappers/target_mappers.py b/pyrit/backend/mappers/target_mappers.py index c3b78792a4..46a97fa6de 100644 --- a/pyrit/backend/mappers/target_mappers.py +++ b/pyrit/backend/mappers/target_mappers.py @@ -18,6 +18,9 @@ from pyrit.models.catalog.target import TargetInstance from pyrit.prompt_target import PromptTarget from pyrit.prompt_target.common.target_capabilities import CapabilityName +from pyrit.registry import TargetRegistry +from pyrit.registry.registry import Reconstructable +from pyrit.registry.resolution import derive_parameters # Capability flag names that should never be surfaced as identifier-level params: # they are sourced from `target_obj.capabilities` instead. @@ -69,9 +72,24 @@ def target_object_to_instance(target_registry_name: str, target_obj: PromptTarge TargetInstance DTO with metadata derived from the object. """ target_identifier = TargetIdentifier.from_component_identifier(target_obj.get_identifier()) + reconstructable = isinstance(target_obj, Reconstructable) + reconstruction_error = None + if reconstructable: + try: + TargetRegistry.get_registry_singleton().get_reconstruction_parameters(target_obj) + except ValueError as exc: + reconstructable = False + reconstruction_error = str(exc) return TargetInstance( target_registry_name=target_registry_name, + reconstructable=reconstructable, + reconstruction_error=reconstruction_error, + supports_temperature_override=reconstructable + and any( + parameter.name == "temperature" + for parameter in derive_parameters(cls=type(target_obj), identifier_type=TargetIdentifier) + ), identifier=target_identifier, capabilities=target_obj.capabilities, target_specific_params=_target_specific_params(target_identifier), diff --git a/pyrit/backend/models/attacks.py b/pyrit/backend/models/attacks.py index 3b280c09d6..3b1bf039d1 100644 --- a/pyrit/backend/models/attacks.py +++ b/pyrit/backend/models/attacks.py @@ -32,6 +32,7 @@ PromptResponseError, Score, ) +from pyrit.models.component_spec import TargetBinding from pyrit.models.results.attack_result import normalize_legacy_attack_attribution @@ -43,6 +44,7 @@ class TargetInfo(BaseModel): endpoint: str | None = Field(None, description="Target endpoint URL") model_name: str | None = Field(None, description="Model or deployment name") identifier_hash: str = Field(..., description="Canonical target identifier hash") + binding: TargetBinding | None = None class ScoreView(Score): @@ -290,6 +292,7 @@ def target(self) -> TargetInfo | None: endpoint=cast("str | None", target_id.params.get("endpoint") or None), model_name=cast("str | None", target_id.params.get("model_name") or None), identifier_hash=target_id.hash, + binding=TargetBinding.from_metadata(self.metadata), ) @computed_field # type: ignore[prop-decorator] @@ -384,6 +387,7 @@ class MessagePieceRequest(BaseModel): description="Registry IDs of converters already applied, in execution order, including duplicates. " "Requires converted_value. Use an empty list for manual edits.", ) + applied_converter_provenance: list[str | None] | None = Field(None, max_length=MAX_ITEMS) mime_type: IdentifierStr | None = Field(None, description="MIME type for media content") prompt_metadata: dict[IdentifierStr, Any] | None = Field( None, @@ -414,6 +418,11 @@ def _validate_converted_value_data_type(self) -> "MessagePieceRequest": raise ValueError("converted_value_data_type requires converted_value") if self.applied_converter_ids is not None and self.converted_value is None: raise ValueError("applied_converter_ids requires converted_value") + if self.applied_converter_provenance is not None and ( + self.applied_converter_ids is None + or len(self.applied_converter_provenance) != len(self.applied_converter_ids) + ): + raise ValueError("applied_converter_provenance must match applied_converter_ids") return self @@ -481,6 +490,7 @@ class CreateAttackRequest(_AttackAttributionInput): """ name: TextStr | None = Field(None, description="Attack name/label") + target_binding: TargetBinding | None = None target_registry_name: IdentifierStr | None = Field( None, description="Target registry name, or None for a saved unbound attack" ) @@ -601,6 +611,7 @@ class SaveConversationRequest(_AttackAttributionInput): expected_objective: str | None = None objective: str | None = Field(None, description="Omit to keep the destination's objective unchanged") target_registry_name: str | None = None + target_binding: TargetBinding | None = None messages: list[MessageRequest] = Field(default_factory=list) @model_validator(mode="after") @@ -683,6 +694,7 @@ class AddMessageRequest(MessageRequest): None, description="Target registry name. Required when send=True so the backend knows which target to use.", ) + target_binding: TargetBinding | None = None converter_ids: list[IdentifierStr] | None = Field( None, max_length=MAX_ITEMS, diff --git a/pyrit/backend/models/converters.py b/pyrit/backend/models/converters.py index 28c5a46f40..afee9bf4f0 100644 --- a/pyrit/backend/models/converters.py +++ b/pyrit/backend/models/converters.py @@ -9,10 +9,11 @@ from typing import Any -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator from pyrit.backend.models.common import MAX_ITEMS, REGISTRY_INSTANCE_NAME_PATTERN, IdentifierStr from pyrit.models import ConverterIdentifier, Parameter, PromptDataType +from pyrit.models.component_spec import SourceInstanceSpec __all__ = [ "ConverterInstance", @@ -72,6 +73,8 @@ class ConverterInstance(BaseModel): identifier: ConverterIdentifier = Field(..., description="The converter's identity/configuration projection") is_llm_based: bool = Field(False, description="Whether this converter requires an LLM target") description: str | None = Field(None, description="Short description of the converter type") + reconstructable: bool = False + reconstruction_error: str | None = None class ConverterInstanceListResponse(BaseModel): @@ -83,12 +86,14 @@ class ConverterInstanceListResponse(BaseModel): class CreateConverterRequest(BaseModel): """Request to create a new converter instance.""" - name: str = Field( - ..., + name: str | None = Field( + None, min_length=1, pattern=REGISTRY_INSTANCE_NAME_PATTERN, description="Unique registry name for the converter instance", ) + register: bool = Field(True, description="Register the object; false returns only its descriptor") + source: SourceInstanceSpec | None = None type: IdentifierStr = Field(..., description="Converter type (e.g., 'Base64Converter')") params: dict[IdentifierStr, Any] = Field( default_factory=dict, @@ -96,6 +101,22 @@ class CreateConverterRequest(BaseModel): description="Converter constructor parameters", ) + @model_validator(mode="after") + def _validate_registration(self) -> "CreateConverterRequest": + if self.register and not self.name: + raise ValueError("name is required when register=true") + if self.source is not None and self.register: + raise ValueError("source requires register=false") + if self.source is not None and self.params: + raise ValueError("Use source.params for reconstruction overrides") + return self + + +class UnregisteredConverter(BaseModel): + """A constructed descriptor, without a registry ID.""" + + identifier: ConverterIdentifier + # ============================================================================ # Converter Preview @@ -111,6 +132,9 @@ class PreviewStep(BaseModel): input_data_type: PromptDataType = Field(..., description="Input data type") output_value: str = Field(..., description="Output from this converter") output_data_type: PromptDataType = Field(..., description="Output data type") + source: SourceInstanceSpec | None = None + identifier: ConverterIdentifier | None = None + provenance: str | None = None class ConverterPreviewRequest(BaseModel): @@ -119,6 +143,14 @@ class ConverterPreviewRequest(BaseModel): original_value: str = Field(..., description="Text to convert") original_value_data_type: PromptDataType = Field(default="text", description="Data type of original value") converter_ids: list[IdentifierStr] = Field(..., max_length=MAX_ITEMS, description="Converter instance IDs to apply") + converter_specs: list[SourceInstanceSpec | None] | None = Field(None, max_length=MAX_ITEMS) + + @model_validator(mode="after") + def _validate_specs(self) -> "ConverterPreviewRequest": + if self.converter_specs is not None and len(self.converter_specs) != len(self.converter_ids): + raise ValueError("converter_specs must match converter_ids in order") + return self + start_token: str = Field(default="⟪", min_length=1, description="Opening marker for selected text regions") end_token: str = Field(default="⟫", min_length=1, description="Closing marker for selected text regions") diff --git a/pyrit/backend/models/scorers.py b/pyrit/backend/models/scorers.py index 2da03a09d8..0787128aed 100644 --- a/pyrit/backend/models/scorers.py +++ b/pyrit/backend/models/scorers.py @@ -3,10 +3,10 @@ """Request and response models for scorer registry endpoints.""" -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator from pyrit.backend.models.common import REGISTRY_INSTANCE_NAME_PATTERN, PaginationInfo -from pyrit.models import JSONValue, Parameter +from pyrit.models import JSONValue, Parameter, ScorerIdentifier from pyrit.models.catalog.scorer import ScorerInstance @@ -35,6 +35,21 @@ class ScorerListResponse(BaseModel): class CreateScorerRequest(BaseModel): """Request to construct and register a named scorer.""" - name: str = Field(..., min_length=1, pattern=REGISTRY_INSTANCE_NAME_PATTERN, description="Unique registry name") + name: str | None = Field( + None, min_length=1, pattern=REGISTRY_INSTANCE_NAME_PATTERN, description="Unique registry name" + ) + register: bool = Field(True, description="Register the object; false returns only its descriptor") type: str = Field(..., description="Scorer class name") params: dict[str, JSONValue] = Field(default_factory=dict, description="Scorer constructor parameters") + + @model_validator(mode="after") + def _validate_registration(self) -> "CreateScorerRequest": + if self.register and not self.name: + raise ValueError("name is required when register=true") + return self + + +class UnregisteredScorer(BaseModel): + """A scorer descriptor without a registry key.""" + + identifier: ScorerIdentifier diff --git a/pyrit/backend/models/targets.py b/pyrit/backend/models/targets.py index 10b87ee9e0..378d0bbed0 100644 --- a/pyrit/backend/models/targets.py +++ b/pyrit/backend/models/targets.py @@ -10,11 +10,13 @@ from typing import Literal -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator from pyrit.backend.models.common import MAX_ITEMS, REGISTRY_INSTANCE_NAME_PATTERN, IdentifierStr, PaginationInfo -from pyrit.models import JSONValue, Parameter +from pyrit.models import JSONValue, Parameter, TargetIdentifier from pyrit.models.catalog.target import TargetInstance +from pyrit.models.component_spec import SourceInstanceSpec +from pyrit.models.target.target_capabilities import TargetCapabilities __all__ = [ "CreateTargetRequest", @@ -67,6 +69,8 @@ class CreateTargetRequest(BaseModel): pattern=REGISTRY_INSTANCE_NAME_PATTERN, description="Unique registry name; omitted only for legacy UI compatibility", ) + register: bool = Field(True, description="Register the object; false returns only its descriptor") + source: SourceInstanceSpec | None = None type: IdentifierStr = Field(..., description="Target type (e.g., 'OpenAIChatTarget')") params: dict[IdentifierStr, JSONValue] = Field( default_factory=dict, max_length=MAX_ITEMS, description="Target constructor parameters" @@ -81,3 +85,16 @@ class CreateTargetRequest(BaseModel): "AzureBlobStorageTarget, and PromptShieldTarget." ), ) + + @model_validator(mode="after") + def _validate_source(self) -> "CreateTargetRequest": + if self.source is not None and (self.register or self.params): + raise ValueError("source requires register=false and overrides in source.params") + return self + + +class UnregisteredTarget(BaseModel): + """A constructed target descriptor with no registry key or object handle.""" + + identifier: TargetIdentifier + capabilities: TargetCapabilities diff --git a/pyrit/backend/routes/converters.py b/pyrit/backend/routes/converters.py index 7423b1021f..2c1efe9641 100644 --- a/pyrit/backend/routes/converters.py +++ b/pyrit/backend/routes/converters.py @@ -18,6 +18,7 @@ ConverterPreviewResponse, ConverterTypeResponse, CreateConverterRequest, + UnregisteredConverter, ) from pyrit.backend.services.converter_service import get_converter_service @@ -58,21 +59,23 @@ async def list_converter_types() -> ConverterTypeResponse: # pyrit-async-suffix @router.post( "", - response_model=ConverterInstance, + response_model=ConverterInstance | UnregisteredConverter, status_code=status.HTTP_201_CREATED, responses={ 400: {"model": ProblemDetail, "description": "Invalid converter type or parameters"}, }, ) -async def create_converter(request: CreateConverterRequest) -> ConverterInstance: # pyrit-async-suffix-exempt +async def create_converter( + request: CreateConverterRequest, +) -> ConverterInstance | UnregisteredConverter: # pyrit-async-suffix-exempt """ Create a new converter instance. - Instantiates a converter with the given type and parameters. + Instantiates a converter with the given type and parameters, with optional registration. Supports nested converters via converter_id references in params. Returns: - ConverterInstance: The created converter instance details. + ConverterInstance | UnregisteredConverter: A named instance or an unregistered descriptor. """ service = get_converter_service() diff --git a/pyrit/backend/routes/scorers.py b/pyrit/backend/routes/scorers.py index 6635f10eff..fb99ad0c7b 100644 --- a/pyrit/backend/routes/scorers.py +++ b/pyrit/backend/routes/scorers.py @@ -10,6 +10,7 @@ CreateScorerRequest, ScorerListResponse, ScorerTypeResponse, + UnregisteredScorer, ) from pyrit.backend.services.scorer_service import get_scorer_service from pyrit.models.catalog.scorer import ScorerInstance @@ -44,16 +45,18 @@ async def list_scorers( @router.post( "", - response_model=ScorerInstance, + response_model=ScorerInstance | UnregisteredScorer, status_code=status.HTTP_201_CREATED, responses={400: {"model": ProblemDetail, "description": "Invalid scorer type, parameters, or name"}}, ) -async def create_scorer(request: CreateScorerRequest) -> ScorerInstance: # pyrit-async-suffix-exempt +async def create_scorer( + request: CreateScorerRequest, +) -> ScorerInstance | UnregisteredScorer: # pyrit-async-suffix-exempt """ - Construct a scorer through ScorerRegistry and register it under its name. + Construct a scorer through ScorerRegistry, with optional registration. Returns: - ScorerInstance: The registered scorer and complete identifier. + ScorerInstance | UnregisteredScorer: A named instance or an unregistered descriptor. """ try: return await get_scorer_service().create_scorer_async(request=request) diff --git a/pyrit/backend/routes/targets.py b/pyrit/backend/routes/targets.py index 050b4ed0c7..5c11f400fa 100644 --- a/pyrit/backend/routes/targets.py +++ b/pyrit/backend/routes/targets.py @@ -15,6 +15,7 @@ CreateTargetRequest, TargetListResponse, TargetTypeResponse, + UnregisteredTarget, ) from pyrit.backend.services.target_service import get_target_service from pyrit.models.catalog.target import TargetInstance @@ -65,7 +66,7 @@ async def list_target_types() -> TargetTypeResponse: # pyrit-async-suffix-exemp @router.post( "", - response_model=TargetInstance, + response_model=TargetInstance | UnregisteredTarget, status_code=status.HTTP_201_CREATED, responses={ 400: { @@ -76,17 +77,17 @@ async def list_target_types() -> TargetTypeResponse: # pyrit-async-suffix-exemp ) async def create_target( request: CreateTargetRequest, -) -> TargetInstance: # pyrit-async-suffix-exempt +) -> TargetInstance | UnregisteredTarget: # pyrit-async-suffix-exempt """ Create a new target instance. - Instantiates a target with the given type and parameters. - The target becomes available for use in attacks. + Instantiates a target with the given type and parameters. Registration is + optional; unregistered responses contain no reusable object handle. Note: Sensitive parameters (API keys, tokens) are filtered from the response. Returns: - CreateTargetResponse: The created target instance details. + TargetInstance | UnregisteredTarget: A named instance or an unregistered descriptor. """ service = get_target_service() diff --git a/pyrit/backend/services/attack_service.py b/pyrit/backend/services/attack_service.py index 14d4dfb2ec..845ecd256d 100644 --- a/pyrit/backend/services/attack_service.py +++ b/pyrit/backend/services/attack_service.py @@ -54,6 +54,7 @@ ) from pyrit.backend.models.common import PaginationInfo from pyrit.backend.models.message_sends import MessageSendRequest, MessageSendStatus +from pyrit.backend.services.component_lifecycle import release_component_async from pyrit.backend.services.media_persistence import persist_message_pieces_async from pyrit.backend.services.message_send_service import ( MessageSendService, @@ -87,6 +88,7 @@ MessagePiece, TargetIdentifier, ) +from pyrit.models.component_spec import TargetBinding from pyrit.models.messages.tool_content import validate_tool_conversation from pyrit.prompt_target import PromptTarget @@ -399,7 +401,9 @@ async def create_attack_async(self, *, request: CreateAttackRequest) -> CreateAt Raises: ValueError: If the target is not found. """ - target_identifier = await self._get_save_target_async(request.target_registry_name) + target_identifier = await self._get_save_target_async( + registry_name=request.target_registry_name, binding=request.target_binding + ) copied: Sequence[MessagePiece] = [] if request.source_conversation_id is not None and request.cutoff_index is not None: conversation, copied = await self._prepare_conversation_up_to_async( @@ -495,6 +499,7 @@ def _new_manual_attack( "created_at": now.isoformat(), "target_unbound": conversation.target_identifier is None, **({"target_registry_name": request.target_registry_name} if request.target_registry_name else {}), + **(request.target_binding.to_metadata() if request.target_binding else {}), }, ) @@ -569,6 +574,7 @@ async def save_conversation_async(self, *, request: SaveConversationRequest) -> new_attack: AttackResult | None = None expected_fields: dict[str, Any] = {} update_fields: dict[str, Any] = {} + binding = request.target_binding if same_attack: attack_result_id = str(request.attack_result_id) results = await self._memory.get_attack_results_async(attack_result_ids=[attack_result_id]) @@ -579,8 +585,14 @@ async def save_conversation_async(self, *, request: SaveConversationRequest) -> raise PermissionError("Cannot save to an attack owned by another operator") identifier = attack.get_attack_strategy_identifier() target_identifier = identifier.get_child("objective_target") if identifier else None + saved_binding = TargetBinding.from_metadata(attack.metadata) + if request.target_binding is not None and request.target_binding != saved_binding: + raise ValueError("Same attack must keep its temperature. Choose New attack to change it.") + binding = saved_binding if request.target_registry_name: - selected_target = await self._get_save_target_async(request.target_registry_name) + selected_target = await self._get_save_target_async( + registry_name=request.target_registry_name, binding=binding + ) if ( target_identifier is None or selected_target is None @@ -597,7 +609,9 @@ async def save_conversation_async(self, *, request: SaveConversationRequest) -> update_fields = self._objective_update_fields(old=attack.objective, new=request.objective) else: attack_result_id = str(uuid.uuid5(request.save_id, "attack")) - target_identifier = await self._get_save_target_async(request.target_registry_name) + target_identifier = await self._get_save_target_async( + registry_name=request.target_registry_name, binding=request.target_binding + ) new_attack = self._new_manual_attack( request=request, attack_result_id=attack_result_id, @@ -607,6 +621,7 @@ async def save_conversation_async(self, *, request: SaveConversationRequest) -> target = await self._validate_editor_target_async( target_identifier=target_identifier, registry_name=request.target_registry_name, + binding=binding, ) persisted_paths: list[str] = [] inserted = False @@ -642,6 +657,8 @@ async def save_conversation_async(self, *, request: SaveConversationRequest) -> inserted = await save_task raise finally: + if binding and target: + await release_component_async(target) if not inserted: await self._cleanup_saved_media_async(persisted_paths) return await self._saved_conversation_response_async( @@ -700,13 +717,21 @@ async def _get_save_source_async(self, *, request: SaveConversationRequest) -> d ) return {piece.id: piece for piece in pieces} - async def _get_save_target_async(self, registry_name: str | None) -> TargetIdentifier | None: + async def _get_save_target_async( + self, *, registry_name: str | None, binding: TargetBinding | None = None + ) -> TargetIdentifier | None: """ Resolve an optional target without calling it. Returns: The target identity, or None for an unbound draft. """ + if binding is not None: + target = await get_target_service().resolve_binding_async(binding) + try: + return TargetIdentifier.from_component_identifier(target.get_identifier()) + finally: + await release_component_async(target) if registry_name is None: return None service = get_target_service() @@ -795,6 +820,7 @@ async def _validate_editor_target_async( *, target_identifier: ComponentIdentifier | None, registry_name: str | None, + binding: TargetBinding | None = None, ) -> PromptTarget | None: """ Resolve a registered target that supports editable, multi-turn history. @@ -808,6 +834,16 @@ async def _validate_editor_target_async( if target_identifier is None: return None service = get_target_service() + if binding is not None: + target_object = await service.resolve_binding_async(binding) + if ( + target_object.get_identifier().hash != target_identifier.hash + or not target_object.capabilities.supports_editable_history + or not target_object.capabilities.supports_multi_turn + ): + await release_component_async(target_object) + raise ValueError("The saved target does not support this conversation") + return target_object target = await service.get_target_async(target_registry_name=registry_name) if registry_name else None if not registry_name: cursor = None @@ -1087,9 +1123,13 @@ async def _bind_requested_target_async(self, *, attack_result_id: str, request: if results and results[0].metadata.get("target_unbound") is True: if request.target_conversation_id not in results[0].get_active_conversation_ids(): raise ValueError(f"Conversation '{request.target_conversation_id}' is not part of this attack") - await self._bind_manual_target_async(attack=results[0], registry_name=request.target_registry_name) + await self._bind_manual_target_async( + attack=results[0], registry_name=request.target_registry_name, binding=request.target_binding + ) - async def _bind_manual_target_async(self, *, attack: AttackResult, registry_name: str | None) -> AttackResult: + async def _bind_manual_target_async( + self, *, attack: AttackResult, registry_name: str | None, binding: TargetBinding | None = None + ) -> AttackResult: """ Bind an explicitly unbound manual attack before its first send. @@ -1098,7 +1138,7 @@ async def _bind_manual_target_async(self, *, attack: AttackResult, registry_name """ if not registry_name: raise ValueError("Select a target before sending") - target = await self._get_save_target_async(registry_name) + target = await self._get_save_target_async(registry_name=registry_name, binding=binding) conversations = { conversation_id: await self._memory.get_conversation_messages_async(conversation_id=conversation_id) for conversation_id in attack.get_active_conversation_ids() @@ -1106,10 +1146,15 @@ async def _bind_manual_target_async(self, *, attack: AttackResult, registry_name target_object = await self._validate_editor_target_async( target_identifier=target, registry_name=registry_name, + binding=binding, ) - if target_object: - for conversation in conversations.values(): - target_object.validate_history(conversation) + try: + if target_object: + for conversation in conversations.values(): + target_object.validate_history(conversation) + finally: + if binding and target_object: + await release_component_async(target_object) atomic = AtomicAttackIdentifier.build( attack_identifier=AttackIdentifier( class_name="ManualAttack", @@ -1117,7 +1162,11 @@ async def _bind_manual_target_async(self, *, attack: AttackResult, registry_name objective_target=target, ) ) - metadata = {"target_unbound": False, "target_registry_name": registry_name} + metadata = { + "target_unbound": False, + "target_registry_name": registry_name, + **(binding.to_metadata() if binding else {}), + } await self._memory.update_attack_result_conditionally_async( attack_result_id=attack.attack_result_id, expected_fields={ diff --git a/pyrit/backend/services/component_lifecycle.py b/pyrit/backend/services/component_lifecycle.py new file mode 100644 index 0000000000..b11d87a1d5 --- /dev/null +++ b/pyrit/backend/services/component_lifecycle.py @@ -0,0 +1,66 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Cleanup of objects owned by one backend operation.""" + +import asyncio +import logging +from collections.abc import Callable +from typing import ParamSpec, Protocol, TypeVar, runtime_checkable + +_Params = ParamSpec("_Params") +_Component = TypeVar("_Component") +logger = logging.getLogger(__name__) + + +@runtime_checkable +class TemporaryResource(Protocol): + """A component with operation-owned resources.""" + + async def cleanup_target_async(self) -> None: + """Release owned resources, not registered dependencies.""" + ... + + +async def release_component_async(component: object) -> None: + """Release only a temporary object's explicitly owned resources.""" + if isinstance(component, TemporaryResource): + task = asyncio.create_task(component.cleanup_target_async()) + cancelled = False + while not task.done(): + try: + await asyncio.shield(task) + except asyncio.CancelledError: + cancelled = True + task.result() + if cancelled: + raise asyncio.CancelledError + + +async def construct_component_async( + factory: Callable[_Params, _Component], *args: _Params.args, **kwargs: _Params.kwargs +) -> _Component: + """ + Finish thread construction before cancellation releases ownership. + + Returns: + _Component: The constructed component, owned by the caller. + + Raises: + asyncio.CancelledError: After the cancelled operation releases its component. + """ + task = asyncio.create_task(asyncio.to_thread(factory, *args, **kwargs)) + try: + return await asyncio.shield(task) + except asyncio.CancelledError: + while not task.done(): + try: + await asyncio.shield(task) + except asyncio.CancelledError: + continue + except Exception: + logger.exception("Component construction failed while the request was cancelled") + break + if not task.cancelled() and task.exception() is None: + await release_component_async(task.result()) + raise diff --git a/pyrit/backend/services/converter_service.py b/pyrit/backend/services/converter_service.py index 2f31db3a1f..f5024419bb 100644 --- a/pyrit/backend/services/converter_service.py +++ b/pyrit/backend/services/converter_service.py @@ -15,7 +15,10 @@ import asyncio import base64 import binascii +import hashlib +import hmac import mimetypes +import secrets import uuid from contextlib import suppress from functools import lru_cache @@ -36,13 +39,16 @@ ConverterTypeResponse, CreateConverterRequest, PreviewStep, + UnregisteredConverter, ) +from pyrit.backend.services.component_lifecycle import construct_component_async from pyrit.backend.services.media_persistence import persist_media_value_async from pyrit.common.azure_storage import is_azure_blob_uri from pyrit.memory import data_serializer_factory -from pyrit.models import MessagePiece, PromptDataType +from pyrit.models import ConverterIdentifier, MessagePiece, PromptDataType from pyrit.prompt_normalizer import ConverterConfiguration, PromptNormalizer from pyrit.registry.components import ConverterRegistry +from pyrit.registry.registry import Reconstructable _OWNED_ARTIFACT_PATHS_KEY = "owned_artifact_paths" _DEFAULT_UPLOAD_EXTENSION = ".bin" @@ -61,6 +67,7 @@ def __init__(self) -> None: self._registry = ConverterRegistry.get_registry_singleton() self._upload_directory = TemporaryDirectory(prefix="pyrit-registry-uploads-") self._upload_path = Path(self._upload_directory.name).resolve() + self._provenance_key = secrets.token_bytes(32) def _build_instance_from_object(self, *, converter_id: str, converter_obj: Any) -> ConverterInstance: """ @@ -73,12 +80,19 @@ def _build_instance_from_object(self, *, converter_id: str, converter_obj: Any) """ metadata = self._registry.get_registered_class_metadata(converter_obj.__class__.__name__) description = metadata.class_description or None if metadata else None - return converter_object_to_instance( + result = converter_object_to_instance( converter_id=converter_id, converter_obj=converter_obj, is_llm_based=metadata.is_llm_based if metadata else False, description=description, ) + if isinstance(converter_obj, Reconstructable): + try: + self._registry.get_reconstruction_parameters(converter_obj) + result.reconstructable = True + except ValueError as exc: + result.reconstruction_error = str(exc) + return result # ======================================================================== # Public API Methods @@ -174,7 +188,9 @@ async def delete_converter_async(self, *, converter_id: str) -> bool: await self._remove_owned_artifacts_async(paths=owned_paths) return self._registry.instances.unregister(converter_id, expected_entry=entry) is not None - async def create_converter_async(self, *, request: CreateConverterRequest) -> ConverterInstance: + async def create_converter_async( + self, *, request: CreateConverterRequest + ) -> ConverterInstance | UnregisteredConverter: """ Create a new converter instance from API request. @@ -194,15 +210,48 @@ async def create_converter_async(self, *, request: CreateConverterRequest) -> Co """ if request.type not in self._registry: raise ValueError(f"Converter type '{request.type}' not found") - self._registry.instances.validate_name_available(request.name) + if request.register: + if request.name is None: + raise ValueError("name is required when register=true") + self._registry.instances.validate_name_available(request.name) + source = ( + self._registry.resolve_source(name=request.source.source_name, identifier_hash=request.source.source_hash) + if request.source + else None + ) + if source is not None and type(source) is not self._registry.get_class(request.type): + raise ValueError("The source converter type does not match the requested type") params, owned_paths = await self._persist_data_uri_params_async( converter_type=request.type, - params=request.params, + params=request.source.params if request.source else request.params, ) try: # Uploads may have yielded to another request that took the name. - self._registry.instances.validate_name_available(request.name) - converter_obj = self._registry.create_instance_from_external_input(request.type, params=params) + if request.register: + if request.name is None: + raise ValueError("name is required when register=true") + self._registry.instances.validate_name_available(request.name) + converter_obj = ( + await construct_component_async( + self._registry.recreate_instance, source=source, params=params, external_input=True + ) + if source is not None + else await construct_component_async( + self._registry.create_instance_from_external_input, request.type, params=params + ) + ) + if not request.register: + if ( + request.source + and request.source.effective_hash + and (converter_obj.get_identifier().hash != request.source.effective_hash) + ): + raise ValueError("Temporary converter configuration has changed") + return UnregisteredConverter( + identifier=ConverterIdentifier.from_component_identifier(converter_obj.get_identifier()) + ) + if request.name is None: + raise ValueError("name is required when register=true") converter = self._build_instance_from_object(converter_id=request.name, converter_obj=converter_obj) self._registry.instances.register( converter_obj, @@ -210,8 +259,12 @@ async def create_converter_async(self, *, request: CreateConverterRequest) -> Co metadata={_OWNED_ARTIFACT_PATHS_KEY: [str(path) for path in owned_paths]}, ) except (Exception, asyncio.CancelledError): - await self._remove_owned_artifacts_async(paths=owned_paths) + if request.register: + await self._remove_owned_artifacts_async(paths=owned_paths) raise + finally: + if not request.register: + await self._remove_owned_artifacts_async(paths=owned_paths) return converter @@ -243,14 +296,41 @@ async def preview_conversion_async(self, *, request: ConverterPreviewRequest) -> ) original_value = result.value - converters = self._gather_converters(converter_ids=request.converter_ids) - steps, final_value, final_type = await self._apply_converters_async( - converters=converters, - initial_value=original_value, - initial_type=data_type, - start_token=request.start_token, - end_token=request.end_token, - ) + converters: list[tuple[str, str, Any]] = [] + owned_paths: list[Path] = [] + specs = request.converter_specs or [None] * len(request.converter_ids) + try: + for converter_id, spec in zip(request.converter_ids, specs, strict=True): + if spec is None: + converters.extend(self._gather_converters(converter_ids=[converter_id])) + continue + if spec.source_name != converter_id: + raise ValueError("A converter specification must match its pipeline source") + source = self._registry.resolve_source(name=spec.source_name, identifier_hash=spec.source_hash) + params, paths = await self._persist_data_uri_params_async( + converter_type=type(source).__name__, params=spec.params + ) + owned_paths.extend(paths) + obj = await construct_component_async( + self._registry.recreate_instance, source=source, params=params, external_input=True + ) + if spec.effective_hash and obj.get_identifier().hash != spec.effective_hash: + raise ValueError("Temporary converter configuration has changed") + converters.append((converter_id, type(obj).__name__, obj)) + steps, final_value, final_type = await self._apply_converters_async( + converters=converters, + initial_value=original_value, + initial_type=data_type, + start_token=request.start_token, + end_token=request.end_token, + ) + for step, spec, (_, _, obj) in zip(steps, specs, converters, strict=True): + if spec is not None: + step.source = spec + step.identifier = ConverterIdentifier.from_component_identifier(obj.get_identifier()) + step.provenance = self._sign_identifier(step.identifier) + finally: + await self._remove_owned_artifacts_async(paths=owned_paths) return ConverterPreviewResponse( original_value=request.original_value, @@ -260,6 +340,28 @@ async def preview_conversion_async(self, *, request: ConverterPreviewRequest) -> steps=steps, ) + def _sign_identifier(self, identifier: ConverterIdentifier) -> str: + payload = identifier.model_dump_json().encode() + signature = hmac.new(self._provenance_key, payload, hashlib.sha256).hexdigest() + return f"{signature}.{base64.urlsafe_b64encode(payload).decode()}" + + def read_provenance(self, token: str) -> ConverterIdentifier: + """ + Validate preview-issued evidence without retaining a converter object. + + Returns: + ConverterIdentifier: The preview's actual component identity. + """ + try: + signature, encoded = token.split(".", 1) + payload = base64.b64decode(encoded, altchars=b"-_", validate=True) + except (ValueError, binascii.Error) as exc: + raise ValueError("Invalid converter provenance") from exc + expected = hmac.new(self._provenance_key, payload, hashlib.sha256).hexdigest() + if not hmac.compare_digest(signature, expected): + raise ValueError("Converter provenance is invalid or belongs to an earlier runtime; convert again") + return ConverterIdentifier.model_validate_json(payload) + def get_converter_objects_for_ids(self, *, converter_ids: list[str]) -> list[Any]: """ Get converter objects for a list of IDs. diff --git a/pyrit/backend/services/message_send_service.py b/pyrit/backend/services/message_send_service.py index 57e2064233..dbf0084f9f 100644 --- a/pyrit/backend/services/message_send_service.py +++ b/pyrit/backend/services/message_send_service.py @@ -27,6 +27,7 @@ MessageSendStatus, RequestConverterMode, ) +from pyrit.backend.services.component_lifecycle import release_component_async from pyrit.backend.services.converter_service import get_converter_service from pyrit.backend.services.manual_send_scheduler import ( ManualSendConflictError, @@ -49,6 +50,7 @@ Message, MessagePiece, ) +from pyrit.models.component_spec import TargetBinding from pyrit.prompt_normalizer import ConverterConfiguration, PromptNormalizer from pyrit.prompt_target import PromptTarget from pyrit.prompt_target.common.target_send_context import TargetSendContext @@ -66,6 +68,8 @@ class _ValidatedMessage: request_configurations: list[ConverterConfiguration] response_configurations: list[ConverterConfiguration] applied_identifiers: dict[int, list[ConverterIdentifier]] + owns_target: bool = False + target_binding: TargetBinding | None = None @dataclass(kw_only=True) @@ -111,16 +115,32 @@ def resolve_applied_converter_identifiers( Returns: Registry-validated converter identifiers by message piece index, preserving order and duplicates. """ - return { - index: [ - ConverterIdentifier.from_component_identifier(converter.get_identifier()) - for converter in get_converter_service().get_converter_objects_for_ids( - converter_ids=piece.applied_converter_ids + result: dict[int, list[ConverterIdentifier]] = {} + service = get_converter_service() + for index, piece in enumerate(pieces): + if not piece.applied_converter_ids: + continue + tokens = piece.applied_converter_provenance or [None] * len(piece.applied_converter_ids) + if all(token is None for token in tokens): + result[index] = [ + ConverterIdentifier.from_component_identifier(converter.get_identifier()) + for converter in service.get_converter_objects_for_ids(converter_ids=piece.applied_converter_ids) + ] + continue + registered = iter( + service.get_converter_objects_for_ids( + converter_ids=[ + name for name, token in zip(piece.applied_converter_ids, tokens, strict=True) if token is None + ] ) + ) + result[index] = [ + service.read_provenance(token) + if token is not None + else ConverterIdentifier.from_component_identifier(next(registered).get_identifier()) + for token in tokens ] - for index, piece in enumerate(pieces) - if piece.applied_converter_ids - } + return result class MessageSendService: @@ -280,13 +300,17 @@ async def _validate_message_async(self, *, attack_result_id: str, request: AddMe raise ValueError(f"Attack '{attack_result_id}' not found") ar = results[0] + binding = TargetBinding.from_metadata(ar.metadata) + if request.target_binding is not None and request.target_binding != binding: + raise ValueError("This attack's temperature is read-only. Create a new attack to change it.") target_registry_name = request.target_registry_name target = ( get_target_service().get_target_object(target_registry_name=target_registry_name) - if request.send and target_registry_name + if binding is None and request.send and target_registry_name else None ) - self._validate_target_match(attack_identifier=ar.get_attack_strategy_identifier(), target=target) + if binding is None: + self._validate_target_match(attack_identifier=ar.get_attack_strategy_identifier(), target=target) msg_conversation_id = request.target_conversation_id @@ -301,14 +325,23 @@ async def _validate_message_async(self, *, attack_result_id: str, request: AddMe response_converter_configs = self._resolve_converter_configs( configurations=request.response_converter_configurations ) - if request.send and target is None: + applied_identifiers = resolve_applied_converter_identifiers(request.pieces) + if binding is not None and request.send: + target = await get_target_service().resolve_binding_async(binding) + try: + self._validate_target_match(attack_identifier=ar.get_attack_strategy_identifier(), target=target) + finally: + await release_component_async(target) + target = None + if request.send and target is None and binding is None: raise ValueError(f"Target object for '{target_registry_name}' not found") return _ValidatedMessage( target=target, request_configurations=request_converter_configs, response_configurations=response_converter_configs, - applied_identifiers=resolve_applied_converter_identifiers(request.pieces), + applied_identifiers=applied_identifiers, + target_binding=binding if request.send else None, ) async def _execute_validated_message_async( @@ -320,30 +353,38 @@ async def _execute_validated_message_async( progress: MessageSendConversation | None = None, prepared_message: Message | None = None, ) -> None: - async with self._scheduler.operation_async(): - if progress is not None: - progress.state = MessageSendState.PREPARING - with attack_result_id_scope(attack_result_id=attack_result_id): - await self._complete_memory_write_async( - partial( - self._memory.add_conversation_to_memory_async, - conversation=Conversation( - conversation_id=request.target_conversation_id, - target_identifier=validated.target.get_identifier() if validated.target else None, - attack_result_id=attack_result_id, - ), + try: + async with self._scheduler.operation_async(): + if validated.target_binding is not None: + validated.target = await get_target_service().resolve_binding_async(validated.target_binding) + validated.owns_target = True + if progress is not None: + progress.state = MessageSendState.PREPARING + with attack_result_id_scope(attack_result_id=attack_result_id): + await self._complete_memory_write_async( + partial( + self._memory.add_conversation_to_memory_async, + conversation=Conversation( + conversation_id=request.target_conversation_id, + target_identifier=validated.target.get_identifier() if validated.target else None, + attack_result_id=attack_result_id, + ), + ) ) - ) - await self._execute_message_async( - attack_result_id=attack_result_id, - request=request, - target=validated.target, - request_converter_configurations=validated.request_configurations, - response_converter_configurations=validated.response_configurations, - applied_converter_identifiers=validated.applied_identifiers, - progress=progress, - prepared_message=prepared_message, - ) + await self._execute_message_async( + attack_result_id=attack_result_id, + request=request, + target=validated.target, + request_converter_configurations=validated.request_configurations, + response_converter_configurations=validated.response_configurations, + applied_converter_identifiers=validated.applied_identifiers, + progress=progress, + prepared_message=prepared_message, + ) + finally: + if validated.owns_target and validated.target: + await release_component_async(validated.target) + validated.owns_target = False async def _run_send_async( self, *, operation: _Send, request: MessageSendRequest, validated: _ValidatedMessage @@ -368,9 +409,19 @@ async def _run_repeated_send_async( async with self._scheduler.operation_async(): operation.status.state = MessageSendState.PREPARING with attack_result_id_scope(attack_result_id=operation.status.attack_result_id): - messages = await self._prepare_repeated_send_async( - operation=operation, request=request, validated=validated - ) + preparation = replace(validated) + try: + if preparation.target_binding is not None: + preparation.target = await get_target_service().resolve_binding_async( + preparation.target_binding + ) + preparation.owns_target = True + messages = await self._prepare_repeated_send_async( + operation=operation, request=request, validated=preparation + ) + finally: + if preparation.owns_target and preparation.target: + await release_component_async(preparation.target) if request.request_converter_mode == RequestConverterMode.SHARED: validated = replace( validated, @@ -506,7 +557,7 @@ async def _run_conversation_async( await self._execute_validated_message_async( attack_result_id=operation.status.attack_result_id, request=request, - validated=validated, + validated=replace(validated), progress=progress, prepared_message=message, ) diff --git a/pyrit/backend/services/scorer_service.py b/pyrit/backend/services/scorer_service.py index 98282610bf..ed1c7f450c 100644 --- a/pyrit/backend/services/scorer_service.py +++ b/pyrit/backend/services/scorer_service.py @@ -13,6 +13,7 @@ ScorerListResponse, ScorerTypeEntry, ScorerTypeResponse, + UnregisteredScorer, ) from pyrit.models.catalog.scorer import ScorerInstance from pyrit.models.identifiers.scorer_identifier import ScorerIdentifier @@ -104,7 +105,7 @@ def get_instance() -> ScorerInstance | None: return await asyncio.to_thread(get_instance) - async def create_scorer_async(self, *, request: CreateScorerRequest) -> ScorerInstance: + async def create_scorer_async(self, *, request: CreateScorerRequest) -> ScorerInstance | UnregisteredScorer: """ Build and register a scorer through the shared registry resolver. @@ -112,9 +113,16 @@ async def create_scorer_async(self, *, request: CreateScorerRequest) -> ScorerIn ScorerInstance: The registered scorer and its identifier. """ - def create() -> ScorerInstance: + def create() -> ScorerInstance | UnregisteredScorer: if request.type not in self._registry: raise ValueError(f"Scorer type '{request.type}' not found") + if not request.register: + scorer = self._registry.create_instance(request.type, **request.params) + return UnregisteredScorer( + identifier=ScorerIdentifier.from_component_identifier(scorer.get_identifier()) + ) + if request.name is None: + raise ValueError("name is required when register=true") scorer = self._registry.create_named_instance( name=request.name, type_name=request.type, diff --git a/pyrit/backend/services/target_service.py b/pyrit/backend/services/target_service.py index eb56dbce62..e425f2dccf 100644 --- a/pyrit/backend/services/target_service.py +++ b/pyrit/backend/services/target_service.py @@ -25,10 +25,15 @@ TargetListResponse, TargetTypeEntry, TargetTypeResponse, + UnregisteredTarget, ) +from pyrit.backend.services.component_lifecycle import construct_component_async, release_component_async from pyrit.common import REQUIRED_VALUE +from pyrit.models import TargetIdentifier from pyrit.models.catalog.target import TargetInstance +from pyrit.models.component_spec import SourceInstanceSpec, TargetBinding from pyrit.models.parameter import Parameter +from pyrit.prompt_target import PromptTarget from pyrit.registry import TargetRegistry logger = logging.getLogger(__name__) @@ -219,7 +224,32 @@ async def list_target_types_async(self) -> TargetTypeResponse: ] return TargetTypeResponse(items=items) - async def create_target_async(self, *, request: CreateTargetRequest) -> TargetInstance: + async def build_from_source_async(self, spec: SourceInstanceSpec) -> PromptTarget: + """ + Rebuild a private target and verify its saved effective identity. + + Returns: + PromptTarget: The caller-owned target. + """ + source = self._registry.resolve_source(name=spec.source_name, identifier_hash=spec.source_hash) + target = await construct_component_async( + self._registry.recreate_instance, source=source, params=spec.params, external_input=True + ) + if spec.effective_hash and target.get_identifier().hash != spec.effective_hash: + await release_component_async(target) + raise ValueError("The saved target configuration cannot be reconstructed") + return target + + async def resolve_binding_async(self, binding: TargetBinding) -> PromptTarget: + """ + Rebuild an attack-owned configuration from its current source. + + Returns: + PromptTarget: The caller-owned target. + """ + return await self.build_from_source_async(binding.to_spec()) + + async def create_target_async(self, *, request: CreateTargetRequest) -> TargetInstance | UnregisteredTarget: """ Create a new target instance from API request. @@ -271,11 +301,40 @@ async def create_target_async(self, *, request: CreateTargetRequest) -> TargetIn # LEGACY COMPATIBILITY: The current configuration UI omits the name. # Remove this generated fallback after that UI sends an explicit name. target_registry_name = request.name or f"compat_{uuid.uuid4().hex}" - self._registry.instances.validate_name_available(target_registry_name) - target_obj = self._registry.create_instance_from_external_input(request.type, params=params) - target = self._build_instance_from_object(target_registry_name=target_registry_name, target_obj=target_obj) - self._registry.instances.register(target_obj, name=target_registry_name) - return target + if request.register: + self._registry.instances.validate_name_available(target_registry_name) + if ( + request.source + and type( + self._registry.resolve_source( + name=request.source.source_name, identifier_hash=request.source.source_hash + ) + ) + is not target_cls + ): + raise ValueError("The source target type does not match the requested type") + target_obj = ( + await self.build_from_source_async(request.source) + if request.source + else await construct_component_async( + self._registry.create_instance_from_external_input, request.type, params=params + ) + ) + if not request.register: + try: + return UnregisteredTarget( + identifier=TargetIdentifier.from_component_identifier(target_obj.get_identifier()), + capabilities=target_obj.capabilities, + ) + finally: + await release_component_async(target_obj) + try: + target = self._build_instance_from_object(target_registry_name=target_registry_name, target_obj=target_obj) + self._registry.instances.register(target_obj, name=target_registry_name) + return target + except Exception: + await release_component_async(target_obj) + raise @lru_cache(maxsize=1) diff --git a/pyrit/common/apply_defaults.py b/pyrit/common/apply_defaults.py index 9ce20f87bb..a3afe1ff53 100644 --- a/pyrit/common/apply_defaults.py +++ b/pyrit/common/apply_defaults.py @@ -302,8 +302,26 @@ def wrapper(self: object, *args: object, **kwargs: object) -> T: f"Either pass a valid value or register a default using set_default_value()." ) - # Call the original method with updated arguments - return method(*bound_args.args, **bound_args.kwargs) + resolved: dict[str, object] | None = None + if getattr(method, "__name__", None) == "__init__" and getattr( + type(self).__init__, "__pyrit_capture_constructor__", False + ): + from pyrit.common.constructor_capture import copy_constructor_inputs + + resolved = {} + for name, value in bound_args.arguments.items(): + if name == "self": + continue + if sig.parameters[name].kind is inspect.Parameter.VAR_KEYWORD: + resolved.update({key: copy_constructor_inputs(item) for key, item in value.items()}) + elif sig.parameters[name].kind is not inspect.Parameter.VAR_POSITIONAL: + resolved[name] = copy_constructor_inputs(value) + + result = method(*bound_args.args, **bound_args.kwargs) + if resolved is not None: + retained = vars(self).setdefault("_resolved_constructor_parameters", {}) + retained[inspect.unwrap(method)] = resolved + return result return wrapper diff --git a/pyrit/common/brick_contract.py b/pyrit/common/brick_contract.py index 1f3d503e56..0569c6502b 100644 --- a/pyrit/common/brick_contract.py +++ b/pyrit/common/brick_contract.py @@ -63,6 +63,23 @@ def init_parameters_are_forwarded(init: Callable[..., object]) -> bool: return bool(getattr(init, _FORWARD_INIT_PARAMETERS_ATTRIBUTE, False)) +def get_constructor_owners(cls: type) -> list[type]: + """ + Get the constructor owners that define the class's build contract. + + Returns: + list[type]: Owners in child-to-parent order, including explicitly forwarded parents. + """ + owners: list[type] = [] + for owner in cls.__mro__: + if "__init__" not in owner.__dict__: + continue + owners.append(owner) + if not init_parameters_are_forwarded(owner.__dict__["__init__"]): + break + return owners + + def enforce_keyword_only_init(cls: type, *, base_name: str) -> None: """ Validate that ``cls.__init__`` only accepts keyword-only parameters. diff --git a/pyrit/common/constructor_capture.py b/pyrit/common/constructor_capture.py new file mode 100644 index 0000000000..28ae93a1b7 --- /dev/null +++ b/pyrit/common/constructor_capture.py @@ -0,0 +1,81 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Opt-in capture of constructor inputs for component reconstruction.""" + +import inspect +from collections.abc import Callable +from functools import wraps +from typing import ParamSpec + +from pyrit.common.brick_contract import get_constructor_owners + +_Params = ParamSpec("_Params") + + +def copy_constructor_inputs(value: object, *, memo: dict[int, object] | None = None) -> object: + """ + Copy containers while retaining live component and client references. + + Returns: + object: An independent container tree with unchanged leaf objects. + """ + if memo is None: + memo = {} + if id(value) in memo: + return memo[id(value)] + if isinstance(value, dict): + copied_dict: dict[object, object] = {} + memo[id(value)] = copied_dict + copied_dict.update({key: copy_constructor_inputs(item, memo=memo) for key, item in value.items()}) + return copied_dict + if isinstance(value, list): + copied_list: list[object] = [] + memo[id(value)] = copied_list + copied_list.extend(copy_constructor_inputs(item, memo=memo) for item in value) + return copied_list + if isinstance(value, tuple): + copied_tuple = tuple(copy_constructor_inputs(item, memo=memo) for item in value) + return memo.setdefault(id(value), copied_tuple) + if isinstance(value, set): + copied_set = {copy_constructor_inputs(item, memo=memo) for item in value} + memo[id(value)] = copied_set + return copied_set + return value + + +def capture_constructor_parameters(init: Callable[_Params, None]) -> Callable[_Params, None]: + """ + Retain inputs on the constructed object, not in a global container. + + Returns: + Callable: A signature-preserving constructor wrapper. + """ + signature = inspect.signature(init) + + @wraps(init) + def captured(*args: _Params.args, **kwargs: _Params.kwargs) -> None: + bound = signature.bind(*args, **kwargs) + bound.apply_defaults() + parameters: dict[str, object] = {} + for key, value in bound.arguments.items(): + kind = signature.parameters[key].kind + if kind is inspect.Parameter.VAR_KEYWORD: + parameters.update({name: copy_constructor_inputs(item) for name, item in value.items()}) + elif key != "self" and kind is not inspect.Parameter.VAR_POSITIONAL: + parameters[key] = copy_constructor_inputs(value) + init(*args, **kwargs) + instance = bound.arguments["self"] + if type(instance).__init__ is captured: + resolved = vars(instance).pop("_resolved_constructor_parameters", {}) + resolved_parameters: dict[str, object] = {} + for owner in reversed(get_constructor_owners(type(instance))): + constructor = inspect.unwrap(owner.__dict__["__init__"]) + for name, value in resolved.get(constructor, {}).items(): + if value is not None or name not in resolved_parameters: + resolved_parameters[name] = value + parameters.update(resolved_parameters) + instance._reconstruction_parameters = parameters + + vars(captured)["__pyrit_capture_constructor__"] = True + return captured diff --git a/pyrit/converter/converter.py b/pyrit/converter/converter.py index 75b73a544b..0bffa4a4d0 100644 --- a/pyrit/converter/converter.py +++ b/pyrit/converter/converter.py @@ -86,6 +86,10 @@ def __init_subclass__(cls, **kwargs: object) -> None: from pyrit.common.brick_contract import enforce_keyword_only_init enforce_keyword_only_init(cls, base_name="Converter") + if "__init__" in cls.__dict__: + from pyrit.common.constructor_capture import capture_constructor_parameters + + type.__setattr__(cls, "__init__", capture_constructor_parameters(cls.__init__)) # Only validate concrete (non-abstract) classes if not inspect.isabstract(cls): if not cls.SUPPORTED_INPUT_TYPES: @@ -139,6 +143,23 @@ def __init__(self, *, converter_target: PromptTarget | None = None) -> None: if converter_target is not None: type(self).TARGET_REQUIREMENTS.validate(target=converter_target) + def get_reconstruction_parameters(self) -> dict[str, object]: + """ + Return server-only constructor inputs for a separate converter. + + Returns: + dict[str, object]: Inputs retained when the source was constructed. + + Raises: + ValueError: If the constructor did not retain its inputs. + """ + parameters = getattr(self, "_reconstruction_parameters", None) + if parameters is None: + if type(self).__init__ is Converter.__init__: + return {} + raise ValueError(f"{type(self).__name__} cannot be reconstructed") + return dict(parameters) + @abc.abstractmethod async def convert_async(self, *, prompt: str, input_type: PromptDataType = "text") -> ConverterResult: """ diff --git a/pyrit/models/catalog/target.py b/pyrit/models/catalog/target.py index 6fad8fe073..614a69bf4f 100644 --- a/pyrit/models/catalog/target.py +++ b/pyrit/models/catalog/target.py @@ -44,6 +44,9 @@ class TargetInstance(BaseModel): """ target_registry_name: str = Field(..., description="Target registry key (e.g., 'azure_openai_chat')") + reconstructable: bool = False + reconstruction_error: str | None = None + supports_temperature_override: bool = False identifier: TargetIdentifier = Field( ..., description=( diff --git a/pyrit/models/component_spec.py b/pyrit/models/component_spec.py new file mode 100644 index 0000000000..b6eb57760e --- /dev/null +++ b/pyrit/models/component_spec.py @@ -0,0 +1,57 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""References for constructing private variants of registered components.""" + +from typing import Literal, Self + +from pydantic import BaseModel, ConfigDict, Field + +from pyrit.models.identifiers.component_identifier import JSONValue + + +class SourceInstanceSpec(BaseModel): + """A source identity and constructor overrides, never a live object handle.""" + + model_config = ConfigDict(extra="forbid") + + source_name: str = Field(min_length=1) + source_hash: str = Field(min_length=1) + params: dict[str, JSONValue] = Field(default_factory=dict) + effective_hash: str | None = None + + +class TargetBinding(BaseModel): + """Non-secret reconstruction information saved with a manual attack.""" + + model_config = ConfigDict(extra="forbid") + + version: Literal[1] = 1 + source_name: str = Field(min_length=1) + source_hash: str = Field(min_length=1) + temperature: float = Field(ge=0, le=2, allow_inf_nan=False) + effective_hash: str = Field(min_length=1) + + def to_spec(self) -> SourceInstanceSpec: + """Return the constructor specification.""" + return SourceInstanceSpec( + source_name=self.source_name, + source_hash=self.source_hash, + params={"temperature": self.temperature}, + effective_hash=self.effective_hash, + ) + + def to_metadata(self) -> dict[str, JSONValue]: + """Return the metadata entry owned by this value.""" + return {"target_binding": self.model_dump(mode="json")} + + @classmethod + def from_metadata(cls, metadata: dict[str, object]) -> Self | None: + """ + Read a binding without treating invalid saved data as a default. + + Returns: + Self | None: The validated binding, or None when no binding is stored. + """ + value = metadata.get("target_binding") + return cls.model_validate(value) if value is not None else None diff --git a/pyrit/prompt_target/common/utils.py b/pyrit/prompt_target/common/utils.py index d2acc514fc..77d330a349 100644 --- a/pyrit/prompt_target/common/utils.py +++ b/pyrit/prompt_target/common/utils.py @@ -22,7 +22,7 @@ def _get_rate_limit_lock(target: Any) -> asyncio.Lock: """Return the target's pacing lock, rebuilding it when the event loop changes.""" loop = asyncio.get_running_loop() - target_vars = vars(target) + target_vars = vars(getattr(target, "_rate_limit_source", target)) lock = target_vars.get("_rate_limit_lock") if lock is None or target_vars.get("_rate_limit_lock_loop") is not loop: lock = asyncio.Lock() diff --git a/pyrit/prompt_target/openai/openai_target.py b/pyrit/prompt_target/openai/openai_target.py index fde956649b..eff600dbe6 100644 --- a/pyrit/prompt_target/openai/openai_target.py +++ b/pyrit/prompt_target/openai/openai_target.py @@ -68,6 +68,15 @@ class OpenAITarget(PromptTarget): api_key_environment_variable: str _async_client: AsyncOpenAI | None = None + _reconstruction_parameters: dict[str, object] = {} + + def __init_subclass__(cls, **kwargs: object) -> None: + """Opt in OpenAI implementations to server-only reconstruction.""" + super().__init_subclass__(**kwargs) + if "__init__" in cls.__dict__: + from pyrit.common.constructor_capture import capture_constructor_parameters + + type.__setattr__(cls, "__init__", capture_constructor_parameters(cls.__init__)) @property def _client(self) -> AsyncOpenAI: @@ -161,6 +170,46 @@ def __init__( self._initialize_openai_client() + def get_reconstruction_parameters(self) -> dict[str, object]: + """ + Return server-only inputs, including resolved authentication. + + Returns: + dict[str, object]: Resolved constructor inputs. + + Raises: + ValueError: If the source has an unsupported resource or temperature setting. + """ + if "http_client" in self._httpx_client_kwargs: + raise ValueError("Reconstruction with an externally owned HTTP client is not supported") + if "temperature" in (getattr(self, "_extra_body_parameters", None) or {}): + raise ValueError( + "Separate temperature settings are not supported when extra_body_parameters sets temperature" + ) + identifier = self.get_identifier() + return { + **self._reconstruction_parameters, + **identifier.params, + **identifier.promoted_scalar_values(), + "model_name": self._model_name, + "endpoint": self._endpoint, + "api_key": self._api_key, + "headers": json.dumps(self._headers), + "max_requests_per_minute": self._max_requests_per_minute, + "httpx_client_kwargs": dict(self._httpx_client_kwargs), + "underlying_model": self._underlying_model, + "custom_configuration": self.configuration, + } + + async def cleanup_target_async(self) -> None: + """Close this target's SDK client.""" + if self._async_client is not None: + await self._async_client.close() + + def attach_reconstruction_source(self, source: object) -> None: + """Share request pacing, not clients or behavioral configuration.""" + self._rate_limit_source = getattr(source, "_rate_limit_source", source) + @staticmethod def _parse_request_headers(value: object) -> dict[str, str]: """ diff --git a/pyrit/registry/registry.py b/pyrit/registry/registry.py index 2376a5fb5e..c9be4afd6e 100644 --- a/pyrit/registry/registry.py +++ b/pyrit/registry/registry.py @@ -30,7 +30,7 @@ import logging import threading from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar +from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar, runtime_checkable from pyrit.registry.instance_registry import DefaultInstanceRegistry, InstanceRegistry from pyrit.registry.registry_metadata import RegistryMetadata @@ -50,6 +50,25 @@ logger = logging.getLogger(__name__) + +@runtime_checkable +class Reconstructable(Protocol): + """A component that can supply private constructor inputs.""" + + def get_reconstruction_parameters(self) -> dict[str, object]: + """Return server-only inputs; these must not be serialized as credentials.""" + ... + + +@runtime_checkable +class ReconstructionContext(Protocol): + """A component that shares source-owned operational context.""" + + def attach_reconstruction_source(self, source: object) -> None: + """Attach context without changing the source's behavioral settings.""" + ... + + T = TypeVar("T") InstanceT = TypeVar("InstanceT", bound="Identifiable") MetadataT = TypeVar("MetadataT", bound=RegistryMetadata) @@ -853,6 +872,91 @@ def create_named_instance( self.instances.register(instance, name=name, metadata=registry_metadata) return instance + def resolve_source(self, *, name: str, identifier_hash: str) -> InstanceT: + """ + Resolve an unchanged source, allowing only an unambiguous rename. + + Returns: + InstanceT: The matching registered source. + + Raises: + ValueError: If the source is missing, changed, or ambiguous. + """ + source = self.instances.get(name) + if source is not None: + if source.get_identifier().hash != identifier_hash: + raise ValueError(f"Source instance '{name}' has changed") + return source + matches = [ + entry.instance + for entry in self.instances.get_all_instances() + if entry.instance.get_identifier().hash == identifier_hash + ] + if len(matches) != 1: + raise ValueError(f"Source instance '{name}' is missing or ambiguous") + return matches[0] + + def get_reconstruction_parameters(self, source: InstanceT) -> dict[str, object]: + """ + Select declared constructor inputs from an opt-in source. + + Returns: + dict[str, object]: Independent input containers with shared live dependencies. + + Raises: + ValueError: If the source does not support reconstruction. + """ + if not isinstance(source, Reconstructable): + raise ValueError(f"{type(source).__name__} does not support reconstruction") + inputs = source.get_reconstruction_parameters() + parameters = derive_parameters(cls=type(source), identifier_type=self._identifier_type()) + retained = getattr(source, "_reconstruction_parameters", {}) + unsupported = set(retained) - {parameter.name for parameter in parameters} + if unsupported: + raise ValueError(f"The source uses undeclared constructor inputs: {', '.join(sorted(unsupported))}") + from pyrit.common.constructor_capture import copy_constructor_inputs + + return { + parameter.name: copy_constructor_inputs(inputs[parameter.name]) + for parameter in parameters + if parameter.name in inputs + } + + def recreate_instance( + self, *, source: InstanceT, params: Mapping[str, object], external_input: bool = False + ) -> InstanceT: + """ + Build an independent component; never change or register the source. + + Returns: + InstanceT: The new component, owned by the caller. + + Raises: + ValueError: If overrides are unknown or the source class is unavailable. + """ + inputs = self.get_reconstruction_parameters(source) + unknown = set(params) - { + parameter.name for parameter in derive_parameters(cls=type(source), identifier_type=self._identifier_type()) + } + if unknown: + raise ValueError(f"Unknown constructor parameters: {', '.join(sorted(unknown))}") + if external_input: + params = resolve_constructor_args( + cls=type(source), + raw_args=dict(params), + identifier_type=self._identifier_type(), + external_input=True, + ) + with self._catalog_lock: + self._ensure_discovered() + names = sorted(name for name, cls in self._classes.items() if cls is type(source)) + if not names: + raise ValueError("The source component's class is no longer registered") + instance = self.create_instance(names[0], **{**inputs, **params}) + if isinstance(instance, ReconstructionContext): + instance.attach_reconstruction_source(source) + return instance + class ParamBagRegistry(Registry[ConfigurableT, MetadataT]): """ diff --git a/pyrit/registry/resolution.py b/pyrit/registry/resolution.py index 17e16dd498..05c7b8feb4 100644 --- a/pyrit/registry/resolution.py +++ b/pyrit/registry/resolution.py @@ -50,7 +50,7 @@ from pydantic import TypeAdapter, ValidationError from pyrit.common.apply_defaults import REQUIRED_VALUE, _RequiredValueSentinel -from pyrit.common.brick_contract import init_parameters_are_forwarded +from pyrit.common.brick_contract import get_constructor_owners from pyrit.models import StructuredParameterValue from pyrit.models.parameter import ComponentType, Parameter, RegistryReference @@ -229,9 +229,8 @@ def _constructor_sources(cls: type) -> list[tuple[type, inspect.Signature]]: Raises: ValueError: If a constructor signature cannot be inspected. """ - owners = [owner for owner in cls.__mro__ if "__init__" in owner.__dict__] sources: list[tuple[type, inspect.Signature]] = [] - for index, owner in enumerate(owners): + for owner in get_constructor_owners(cls): init = owner.__dict__["__init__"] try: signature = inspect.signature(init) @@ -253,8 +252,6 @@ def _constructor_sources(cls: type) -> list[tuple[type, inspect.Signature]]: parameters.append(param.replace(annotation=annotation)) sources.append((owner, signature.replace(parameters=parameters))) - if not init_parameters_are_forwarded(init) or index + 1 == len(owners): - break return sources diff --git a/tests/unit/backend/test_temporary_components.py b/tests/unit/backend/test_temporary_components.py new file mode 100644 index 0000000000..e9963c7587 --- /dev/null +++ b/tests/unit/backend/test_temporary_components.py @@ -0,0 +1,586 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT license. + +"""Private construction, reconstruction, lifetime, and provenance regressions.""" + +import asyncio +import threading +import uuid +from collections.abc import Iterator +from typing import Any +from unittest.mock import AsyncMock, patch + +import pytest +from fastapi import APIRouter, FastAPI +from fastapi.testclient import TestClient + +from pyrit.backend.models.attacks import ( + CreateAttackRequest, + MessagePieceRequest, + MessageRequest, + SaveConversationRequest, +) +from pyrit.backend.models.converters import ConverterPreviewRequest, CreateConverterRequest +from pyrit.backend.models.message_sends import MessageSendRequest, MessageSendState +from pyrit.backend.models.scorers import CreateScorerRequest +from pyrit.backend.models.targets import CreateTargetRequest +from pyrit.backend.routes import converters, scorers, targets +from pyrit.backend.services.attack_service import AttackService +from pyrit.backend.services.component_lifecycle import construct_component_async +from pyrit.backend.services.converter_service import ConverterService, get_converter_service +from pyrit.backend.services.manual_send_scheduler import ManualSendScheduler +from pyrit.backend.services.message_send_service import MessageSendService, resolve_applied_converter_identifiers +from pyrit.backend.services.scorer_service import ScorerService, get_scorer_service +from pyrit.backend.services.target_service import TargetService, get_target_service +from pyrit.common.apply_defaults import reset_default_values, set_default_value +from pyrit.common.brick_contract import forward_init_parameters +from pyrit.converter import CaesarConverter, Converter, ConverterResult, LLMGenericTextConverter, TranslationConverter +from pyrit.memory import SQLiteMemory +from pyrit.models import JSONValue, PromptDataType +from pyrit.models.component_spec import SourceInstanceSpec, TargetBinding +from pyrit.prompt_target import OpenAIChatTarget +from pyrit.prompt_target.common.utils import _get_rate_limit_lock +from pyrit.registry import ConverterRegistry, ScorerRegistry, TargetRegistry +from unit.backend.mocks import _settle_send_async + + +@pytest.fixture(autouse=True) +def isolated_registries() -> Iterator[None]: + for registry in (TargetRegistry, ConverterRegistry, ScorerRegistry): + registry.reset_registry_singleton() + for factory in (get_target_service, get_converter_service, get_scorer_service): + factory.cache_clear() + yield + for factory in (get_target_service, get_converter_service, get_scorer_service): + factory.cache_clear() + for registry in (TargetRegistry, ConverterRegistry, ScorerRegistry): + registry.reset_registry_singleton() + + +def source_target() -> OpenAIChatTarget: + return OpenAIChatTarget( + endpoint="https://example.test/v1", + model_name="test", + api_key="test-only-key", + temperature=0.2, + extra_body_parameters={"example": "preserved"}, + max_requests_per_minute=30, + ) + + +class _MutableConverter(Converter): + SUPPORTED_INPUT_TYPES = ("text",) + SUPPORTED_OUTPUT_TYPES = ("text",) + + def __init__(self, *, settings: dict[str, list[str]]) -> None: + super().__init__() + self.settings = settings + self.settings["values"].append("constructed") + + async def convert_async(self, *, prompt: str, input_type: PromptDataType = "text") -> ConverterResult: + return ConverterResult(output_text=prompt, output_type=input_type) + + +class _ForwardedConverter(LLMGenericTextConverter): + @forward_init_parameters + def __init__(self, **kwargs: Any) -> None: + super().__init__(**kwargs) + + +@pytest.mark.usefixtures("patch_central_database") +class TestTemporaryComponents: + @pytest.mark.parametrize( + ("router", "path", "payload", "registry"), + [ + ( + targets.router, + "/api/targets", + { + "type": "OpenAIChatTarget", + "register": False, + "params": {"endpoint": "https://example.test/v1", "model_name": "test", "api_key": "test-only-key"}, + }, + TargetRegistry, + ), + ( + converters.router, + "/api/converters", + { + "type": "CaesarConverter", + "register": False, + "params": {"caesar_offset": 3}, + }, + ConverterRegistry, + ), + ( + scorers.router, + "/api/scorers", + { + "type": "SubStringScorer", + "register": False, + "params": {"substring": "example"}, + }, + ScorerRegistry, + ), + ], + ) + def test_rest_unregistered_build_has_no_handle( + self, + *, + router: APIRouter, + path: str, + payload: dict[str, JSONValue], + registry: type[TargetRegistry] | type[ConverterRegistry] | type[ScorerRegistry], + ) -> None: + app = FastAPI() + app.include_router(router, prefix="/api") + with TestClient(app) as client: + response = client.post(path, json=payload) + assert response.status_code == 201, response.text + assert "identifier" in response.json() + assert not {"converter_id", "target_registry_name", "scorer_registry_name"} & response.json().keys() + assert "test-only-key" not in response.text + assert registry.get_registry_singleton().instances.get_names() == [] + + async def test_temperature_reconstruction_is_private_and_preserves_auth_async(self) -> None: + source = source_target() + registry = TargetRegistry.get_registry_singleton() + registry.instances.register(source, name="source") + original = source.get_identifier() + service = TargetService() + spec = SourceInstanceSpec(source_name="source", source_hash=original.hash, params={"temperature": 1.5}) + built = await service.create_target_async( + request=CreateTargetRequest( + type="OpenAIChatTarget", + register=False, + source=spec, + ) + ) + assert built.identifier.temperature == 1.5 + binding = TargetBinding( + source_name="source", + source_hash=original.hash, + temperature=1.5, + effective_hash=built.identifier.hash, + ) + restored = await TargetService().resolve_binding_async(TargetBinding.from_metadata(binding.to_metadata())) + try: + assert restored is not source + assert restored.get_identifier().hash == built.identifier.hash + assert restored.get_reconstruction_parameters()["api_key"] == "test-only-key" + assert restored._extra_body_parameters == {"example": "preserved"} + assert restored._rate_limit_source is source + assert _get_rate_limit_lock(restored) is _get_rate_limit_lock(source) + assert not source._client.is_closed() + assert source.get_identifier() is original + assert source._temperature == 0.2 + assert registry.instances.get_names() == ["source"] + assert "test-only-key" not in str(binding.to_metadata()) + finally: + await restored.cleanup_target_async() + await source.cleanup_target_async() + + async def test_two_temperature_variants_are_isolated_async(self) -> None: + source = source_target() + TargetRegistry.get_registry_singleton().instances.register(source, name="source") + service = TargetService() + first, second = await asyncio.gather( + *[ + service.build_from_source_async( + SourceInstanceSpec( + source_name="source", + source_hash=source.get_identifier().hash, + params={"temperature": temperature}, + ) + ) + for temperature in (0.8, 1.5) + ] + ) + try: + assert first._temperature == 0.8 + assert second._temperature == 1.5 + assert source._temperature == 0.2 + assert first._client is not second._client + first._extra_body_parameters["example"] = "private" + assert second._extra_body_parameters == {"example": "preserved"} + assert source._extra_body_parameters == {"example": "preserved"} + finally: + await first.cleanup_target_async() + await second.cleanup_target_async() + await source.cleanup_target_async() + + async def test_effective_hash_mismatch_releases_only_the_private_target_async(self) -> None: + source = source_target() + registry = TargetRegistry.get_registry_singleton() + registry.instances.register(source, name="source") + private = registry.recreate_instance(source=source, params={"temperature": 0.8}) + spec = SourceInstanceSpec( + source_name="source", + source_hash=source.get_identifier().hash, + params={"temperature": 0.8}, + effective_hash="changed", + ) + try: + with patch.object(registry, "recreate_instance", return_value=private): + with pytest.raises(ValueError, match="cannot be reconstructed"): + await TargetService().build_from_source_async(spec) + assert private._client.is_closed() + assert not source._client.is_closed() + finally: + await private.cleanup_target_async() + await source.cleanup_target_async() + + def test_constructor_inputs_are_copied_and_custom_alias_is_resolved(self) -> None: + source = _MutableConverter(settings={"values": ["source"]}) + registry = ConverterRegistry.get_registry_singleton() + registry.register_class(_MutableConverter, name="custom") + derived = registry.recreate_instance(source=source, params={}) + assert source.settings == {"values": ["source", "constructed"]} + assert derived.settings == source.settings + derived.settings["values"].append("derived") + assert source.settings == {"values": ["source", "constructed"]} + + async def test_resolved_default_target_is_retained_async(self) -> None: + source = source_target() + replacement = source_target() + try: + set_default_value(class_type=LLMGenericTextConverter, parameter_name="converter_target", value=source) + converter = LLMGenericTextConverter() + set_default_value(class_type=LLMGenericTextConverter, parameter_name="converter_target", value=replacement) + rebuilt = ConverterRegistry.get_registry_singleton().recreate_instance(source=converter, params={}) + assert rebuilt._converter_target is source + finally: + reset_default_values() + await source.cleanup_target_async() + await replacement.cleanup_target_async() + + async def test_translation_reconstruction_uses_only_its_declared_inputs_async(self) -> None: + target = source_target() + replacement = source_target() + try: + set_default_value(class_type=TranslationConverter, parameter_name="converter_target", value=target) + source = TranslationConverter(language="Spanish", max_retries=7) + identifier = source.get_identifier() + set_default_value(class_type=TranslationConverter, parameter_name="converter_target", value=replacement) + registry = ConverterRegistry.get_registry_singleton() + inputs = registry.get_reconstruction_parameters(source) + assert inputs["converter_target"] is target + assert "system_prompt_template" not in inputs + rebuilt = registry.recreate_instance(source=source, params={"language": "French", "max_retries": "4"}) + assert rebuilt.language == "french" + assert rebuilt._prompt_kwargs["language"] == "french" + assert rebuilt.converter_target is target + assert rebuilt._max_retry_attempts == 4 + assert source.language == "spanish" + assert source._max_retry_attempts == 7 + assert source.get_identifier() is identifier + finally: + reset_default_values() + await target.cleanup_target_async() + await replacement.cleanup_target_async() + + async def test_forwarded_constructor_reconstruction_retains_parent_defaults_async(self) -> None: + target = source_target() + replacement = source_target() + try: + set_default_value(class_type=_ForwardedConverter, parameter_name="converter_target", value=target) + source = _ForwardedConverter() + set_default_value(class_type=_ForwardedConverter, parameter_name="converter_target", value=replacement) + registry = ConverterRegistry.get_registry_singleton() + registry.register_class(_ForwardedConverter, name="forwarded") + rebuilt = registry.recreate_instance(source=source, params={"max_retry_attempts": "5"}) + assert rebuilt._converter_target is target + assert rebuilt._max_retry_attempts == 5 + finally: + reset_default_values() + await target.cleanup_target_async() + await replacement.cleanup_target_async() + + async def test_opaque_constructor_inputs_are_not_silently_discarded_async(self) -> None: + target = source_target() + try: + converter = LLMGenericTextConverter(converter_target=target, language="French") + with pytest.raises(ValueError, match="undeclared constructor inputs: language"): + ConverterRegistry.get_registry_singleton().recreate_instance(source=converter, params={}) + finally: + await target.cleanup_target_async() + + async def test_target_settings_capability_rejects_undeclared_inputs_async(self) -> None: + target = source_target() + TargetRegistry.get_registry_singleton().instances.register(target, name="source") + try: + with patch.dict(target._reconstruction_parameters, {"opaque_setting": "preserve"}): + descriptor = await TargetService().get_target_async(target_registry_name="source") + assert descriptor is not None + assert not descriptor.reconstructable + assert not descriptor.supports_temperature_override + assert descriptor.reconstruction_error == ( + "The source uses undeclared constructor inputs: opaque_setting" + ) + finally: + await target.cleanup_target_async() + + async def test_cancelled_thread_build_releases_component_async(self) -> None: + started = threading.Event() + finish = threading.Event() + component = source_target() + cleanup = AsyncMock() + + def build() -> OpenAIChatTarget: + started.set() + assert finish.wait(timeout=5) + return component + + with patch.object(component, "cleanup_target_async", cleanup): + task = asyncio.create_task(construct_component_async(build)) + assert await asyncio.to_thread(started.wait, 5) + task.cancel() + finish.set() + with pytest.raises(asyncio.CancelledError): + await task + cleanup.assert_awaited_once() + await component.cleanup_target_async() + + async def test_attack_binding_persists_and_same_attack_is_immutable_async( + self, sqlite_instance: SQLiteMemory + ) -> None: + source = source_target() + TargetRegistry.get_registry_singleton().instances.register(source, name="source") + built = await TargetService().build_from_source_async( + SourceInstanceSpec( + source_name="source", + source_hash=source.get_identifier().hash, + params={"temperature": 0.8}, + ) + ) + binding = TargetBinding( + source_name="source", + source_hash=source.get_identifier().hash, + temperature=0.8, + effective_hash=built.get_identifier().hash, + ) + await built.cleanup_target_async() + sender = MessageSendService(scheduler=ManualSendScheduler()) + service = AttackService(message_send_service=sender) + try: + created = await service.create_attack_async( + request=CreateAttackRequest( + target_registry_name="source", + target_binding=binding, + ) + ) + [stored] = await sqlite_instance.get_attack_results_async(attack_result_ids=[created.attack_result_id]) + assert TargetBinding.from_metadata(stored.metadata) == binding + detail = await AttackService().get_attack_async(attack_result_id=created.attack_result_id) + assert detail.target.binding == binding + assert detail.target.identifier_hash == binding.effective_hash + request = SaveConversationRequest( + save_id=uuid.uuid4(), + destination="same_attack", + attack_result_id=created.attack_result_id, + target_registry_name="source", + target_binding=binding, + messages=[MessageRequest(role="user", pieces=[MessagePieceRequest(original_value="test")])], + ) + await service.save_conversation_async(request=request) + request.target_binding = binding.model_copy(update={"temperature": 1.2}) + request.save_id = uuid.uuid4() + with pytest.raises(ValueError, match="keep its temperature"): + await service.save_conversation_async(request=request) + finally: + await sender.shutdown_async() + await source.cleanup_target_async() + + async def test_reconstruction_rejects_wrong_type_unknown_override_and_effective_hash_async(self) -> None: + source = CaesarConverter(caesar_offset=1) + ConverterRegistry.get_registry_singleton().instances.register(source, name="source") + spec = SourceInstanceSpec( + source_name="source", + source_hash=source.get_identifier().hash, + params={"unknown": True}, + ) + service = ConverterService() + try: + with pytest.raises(ValueError, match="Unknown constructor parameters"): + await service.create_converter_async( + request=CreateConverterRequest( + type="CaesarConverter", + source=spec, + register=False, + ) + ) + with pytest.raises(ValueError, match="type does not match"): + await service.create_converter_async( + request=CreateConverterRequest( + type="Base64Converter", + source=spec, + register=False, + ) + ) + spec.params = {} + spec.effective_hash = "different" + with pytest.raises(ValueError, match="configuration has changed"): + await service.create_converter_async( + request=CreateConverterRequest( + type="CaesarConverter", + source=spec, + register=False, + ) + ) + finally: + await service.close_async() + + async def test_missing_changed_and_renamed_sources_async(self) -> None: + source = source_target() + registry = TargetRegistry.get_registry_singleton() + registry.instances.register(source, name="source") + spec = SourceInstanceSpec( + source_name="source", + source_hash=source.get_identifier().hash, + params={"temperature": 0.8}, + ) + registry.instances.unregister("source") + with pytest.raises(ValueError, match="missing or ambiguous"): + await TargetService().build_from_source_async(spec) + registry.instances.register(source, name="renamed") + restored = await TargetService().build_from_source_async(spec) + await restored.cleanup_target_async() + changed = OpenAIChatTarget( + endpoint="https://example.test/v1", + model_name="different", + api_key="test-only-key", + ) + registry.instances.register(changed, name="source") + with pytest.raises(ValueError, match="changed"): + await TargetService().build_from_source_async(spec) + await changed.cleanup_target_async() + await source.cleanup_target_async() + + async def test_source_overrides_keep_external_input_validation_async(self) -> None: + source = source_target() + TargetRegistry.get_registry_singleton().instances.register(source, name="source") + try: + with pytest.raises(ValueError, match="cannot be set through the API"): + await TargetService().build_from_source_async( + SourceInstanceSpec( + source_name="source", + source_hash=source.get_identifier().hash, + params={"httpx_client_kwargs": {"timeout": 5}}, + ) + ) + finally: + await source.cleanup_target_async() + + async def test_repeated_send_owns_each_temperature_target_async(self, sqlite_instance: SQLiteMemory) -> None: + source = source_target() + TargetRegistry.get_registry_singleton().instances.register(source, name="source") + built = await TargetService().build_from_source_async( + SourceInstanceSpec( + source_name="source", source_hash=source.get_identifier().hash, params={"temperature": 0.8} + ) + ) + binding = TargetBinding( + source_name="source", + source_hash=source.get_identifier().hash, + temperature=0.8, + effective_hash=built.get_identifier().hash, + ) + await built.cleanup_target_async() + sender = MessageSendService(scheduler=ManualSendScheduler()) + try: + created = await AttackService(message_send_service=sender).create_attack_async( + request=CreateAttackRequest(target_registry_name="source", target_binding=binding) + ) + with patch.object(sender, "_execute_message_async", new_callable=AsyncMock) as execute: + status = await sender.submit_async( + attack_result_id=created.attack_result_id, + request=MessageSendRequest( + target_registry_name="source", + target_conversation_id=created.conversation_id, + submission_id="temporary-repeat", + count=3, + role="user", + pieces=[MessagePieceRequest(original_value="test")], + ), + ) + status = await _settle_send_async(service=sender, status=status) + assert status.state == MessageSendState.COMPLETED + targets = [call.kwargs["target"] for call in execute.await_args_list] + assert len(targets) == 3 + assert len({id(target) for target in targets}) == 3 + assert all(target._temperature == 0.8 and target._client.is_closed() for target in targets) + assert all(target is not source for target in targets) + assert not source._client.is_closed() + finally: + await sender.shutdown_async() + await source.cleanup_target_async() + + async def test_converter_preview_isolated_and_provenance_survives_discard_async(self) -> None: + source = CaesarConverter(caesar_offset=1) + registry = ConverterRegistry.get_registry_singleton() + registry.instances.register(source, name="caesar") + service = get_converter_service() + spec = SourceInstanceSpec( + source_name="caesar", + source_hash=source.get_identifier().hash, + params={"caesar_offset": 3}, + ) + built = await service.create_converter_async( + request=CreateConverterRequest( + type="CaesarConverter", + register=False, + source=spec, + ) + ) + assert built.identifier.hash != source.get_identifier().hash + preview = await service.preview_conversion_async( + request=ConverterPreviewRequest( + original_value="abc", + converter_ids=["caesar", "caesar"], + converter_specs=[spec, None], + ) + ) + assert preview.steps[0].output_value == "def" + assert preview.converted_value == "efg" + assert (await source.convert_async(prompt="abc", input_type="text")).output_text == "bcd" + token = preview.steps[0].provenance + assert token is not None + registry.instances.unregister("caesar") + piece = MessagePieceRequest( + original_value="abc", + converted_value="def", + applied_converter_ids=["caesar"], + applied_converter_provenance=[token], + ) + assert resolve_applied_converter_identifiers([piece])[0][0].hash == preview.steps[0].identifier.hash + with pytest.raises(ValueError, match="invalid"): + service.read_provenance(("0" if token[0] != "0" else "1") + token[1:]) + replacement = ConverterService() + with pytest.raises(ValueError, match="earlier runtime"): + replacement.read_provenance(token) + await replacement.close_async() + await service.close_async() + + async def test_scorer_build_and_registration_default_async(self) -> None: + service = ScorerService() + await service.create_scorer_async( + request=CreateScorerRequest( + type="SubStringScorer", + params={"substring": "test"}, + register=False, + ) + ) + assert ScorerRegistry.get_registry_singleton().instances.get_names() == [] + await service.create_scorer_async( + request=CreateScorerRequest( + name="named", + type="SubStringScorer", + params={"substring": "test"}, + ) + ) + assert ScorerRegistry.get_registry_singleton().instances.get_names() == ["named"] + + @pytest.mark.parametrize("temperature", [-0.1, 2.1, float("nan"), float("inf")]) + def test_binding_rejects_invalid_temperature(self, temperature: float) -> None: + with pytest.raises(ValueError): + TargetBinding(source_name="source", source_hash="hash", effective_hash="hash", temperature=temperature) diff --git a/tests/unit/prompt_target/test_target_utils.py b/tests/unit/prompt_target/test_target_utils.py index 9175098c53..25fa2b1992 100644 --- a/tests/unit/prompt_target/test_target_utils.py +++ b/tests/unit/prompt_target/test_target_utils.py @@ -11,6 +11,7 @@ from pyrit.exceptions import PyritException, RateLimitException, pyrit_target_retry from pyrit.models import MessagePiece +from pyrit.prompt_target import PromptTarget from pyrit.prompt_target.common.utils import ( build_empty_truncated_response, limit_requests_per_minute, @@ -87,7 +88,7 @@ def test_validate_top_p_above_one_raises(): async def test_limit_requests_per_minute_no_rpm(): - mock_self = MagicMock() + mock_self = MagicMock(spec=PromptTarget) mock_self._max_requests_per_minute = None inner_func = AsyncMock(return_value="response") @@ -100,7 +101,7 @@ async def test_limit_requests_per_minute_no_rpm(): async def test_limit_requests_per_minute_with_rpm(): - mock_self = MagicMock() + mock_self = MagicMock(spec=PromptTarget) mock_self._max_requests_per_minute = 30 inner_func = AsyncMock(return_value="response") @@ -113,7 +114,7 @@ async def test_limit_requests_per_minute_with_rpm(): async def test_limit_requests_per_minute_zero_rpm(): - mock_self = MagicMock() + mock_self = MagicMock(spec=PromptTarget) mock_self._max_requests_per_minute = 0 inner_func = AsyncMock(return_value="response") @@ -125,9 +126,16 @@ async def test_limit_requests_per_minute_zero_rpm(): assert result == "response" -async def test_limit_requests_per_minute_serializes_concurrent_starts() -> None: - target = MagicMock() +@pytest.mark.parametrize("shared_source", [False, True]) +async def test_limit_requests_per_minute_serializes_concurrent_starts(shared_source: bool) -> None: + target = MagicMock(spec=PromptTarget) target._max_requests_per_minute = 60 + targets = [target, target, target] + if shared_source: + temporary_target = MagicMock(spec=PromptTarget) + temporary_target._max_requests_per_minute = 60 + temporary_target._rate_limit_source = target + targets[1] = temporary_target sleep_started: asyncio.Queue[int] = asyncio.Queue() release_sleep: asyncio.Queue[None] = asyncio.Queue() provider_calls: asyncio.Queue[int] = asyncio.Queue() @@ -145,7 +153,10 @@ async def send_async(target: MagicMock, *, request_index: int) -> int: decorated = limit_requests_per_minute(send_async) with patch("pyrit.prompt_target.common.utils.asyncio.sleep", side_effect=controlled_sleep_async): - tasks = [asyncio.create_task(decorated(target, request_index=index)) for index in range(3)] + tasks = [ + asyncio.create_task(decorated(current_target, request_index=index)) + for index, current_target in enumerate(targets) + ] assert await sleep_started.get() == 0 assert delays == [1.0] @@ -161,7 +172,7 @@ async def send_async(target: MagicMock, *, request_index: int) -> int: async def test_limit_requests_per_minute_cancellation_releases_lock() -> None: - target = MagicMock() + target = MagicMock(spec=PromptTarget) target._max_requests_per_minute = 60 sleep_started = asyncio.Event() release_sleep = asyncio.Event() @@ -193,7 +204,7 @@ async def send_async(target: MagicMock, *, value: str) -> str: async def test_target_retry_paces_every_attempt() -> None: - target = MagicMock() + target = MagicMock(spec=PromptTarget) target._max_requests_per_minute = 1 attempts = 0 @@ -234,16 +245,23 @@ def test_retrying_rate_limited_targets_pace_every_attempt() -> None: ) -def test_limit_requests_per_minute_rebuilds_lock_for_new_event_loop() -> None: - target = MagicMock() +@pytest.mark.parametrize("shared_source", [False, True]) +def test_limit_requests_per_minute_rebuilds_lock_for_new_event_loop(shared_source: bool) -> None: + target = MagicMock(spec=PromptTarget) target._max_requests_per_minute = 60 + source = target + if shared_source: + source = MagicMock(spec=PromptTarget) + target._rate_limit_source = source decorated = limit_requests_per_minute(AsyncMock(return_value="response")) async def invoke_async() -> asyncio.Lock: with patch("pyrit.prompt_target.common.utils.asyncio.sleep", new_callable=AsyncMock): await decorated(target) - lock = vars(target)["_rate_limit_lock"] + lock = vars(source)["_rate_limit_lock"] assert isinstance(lock, asyncio.Lock) + if shared_source: + assert "_rate_limit_lock" not in vars(target) return lock first_lock = asyncio.run(invoke_async())