From dc88eb71d9380a7b993446f6000bbefa70173b71 Mon Sep 17 00:00:00 2001 From: Behnam Ousat Date: Tue, 25 Aug 2026 10:08:58 -0700 Subject: [PATCH 01/24] Support adding custom initializers in CopyRIT --- frontend/package.json | 2 + .../Initializers/CustomInitializers.styles.ts | 51 +++++ .../Initializers/CustomInitializers.tsx | 177 ++++++++++++++++++ .../Initializers/Initializers.test.tsx | 59 ++++++ .../components/Initializers/Initializers.tsx | 135 ++++++++++--- .../Initializers/PythonCode.styles.ts | 95 ++++++++++ .../components/Initializers/PythonCode.tsx | 100 ++++++++++ .../src/components/Sidebar/Navigation.tsx | 1 + frontend/src/services/api.ts | 17 ++ frontend/src/types/index.ts | 10 + pyrit/backend/main.py | 4 +- pyrit/backend/routes/initializers.py | 17 +- pyrit/backend/services/initializer_service.py | 60 +++++- ...c8d9e0f1a_add_custom_initializers_table.py | 39 ++++ pyrit/memory/memory_interface.py | 37 ++++ pyrit/memory/memory_models.py | 41 ++++ pyrit/models/__init__.py | 2 + pyrit/models/custom_initializer.py | 21 +++ .../unit/backend/test_initializer_service.py | 107 ++++++++++- tests/unit/backend/test_main.py | 45 ++++- .../memory/test_custom_initializer_memory.py | 42 +++++ 21 files changed, 1024 insertions(+), 38 deletions(-) create mode 100644 frontend/src/components/Initializers/CustomInitializers.styles.ts create mode 100644 frontend/src/components/Initializers/CustomInitializers.tsx create mode 100644 frontend/src/components/Initializers/PythonCode.styles.ts create mode 100644 frontend/src/components/Initializers/PythonCode.tsx create mode 100644 pyrit/memory/alembic/versions/6b7c8d9e0f1a_add_custom_initializers_table.py create mode 100644 pyrit/models/custom_initializer.py create mode 100644 tests/unit/memory/test_custom_initializer_memory.py diff --git a/frontend/package.json b/frontend/package.json index 2a5db9683b..2d4acaa632 100644 --- a/frontend/package.json +++ b/frontend/package.json @@ -27,6 +27,7 @@ "@fluentui/react-components": "9.74.6", "@fluentui/react-icons": "2.0.335", "axios": "1.19.0", + "prismjs": "^1.30.0", "react": "19.2.8", "react-dom": "19.2.8", "react-error-boundary": "6.1.2", @@ -44,6 +45,7 @@ "@testing-library/user-event": "14.6.4", "@types/jest": "30.0.0", "@types/node": "26.2.0", + "@types/prismjs": "^1.26.6", "@types/react": "19.2.18", "@types/react-dom": "19.2.4", "@typescript-eslint/eslint-plugin": "8.67.0", diff --git a/frontend/src/components/Initializers/CustomInitializers.styles.ts b/frontend/src/components/Initializers/CustomInitializers.styles.ts new file mode 100644 index 0000000000..65b79bb86b --- /dev/null +++ b/frontend/src/components/Initializers/CustomInitializers.styles.ts @@ -0,0 +1,51 @@ +import { makeStyles, tokens } from '@fluentui/react-components' + +export const useCustomInitializersStyles = makeStyles({ + root: { + display: 'flex', + flexDirection: 'column', + gap: tokens.spacingVerticalM, + minWidth: 0, + }, + header: { + display: 'flex', + alignItems: 'center', + justifyContent: 'space-between', + flexWrap: 'wrap', + gap: tokens.spacingVerticalM, + }, + tableWrap: { + overflowX: 'auto', + border: `1px solid ${tokens.colorNeutralStroke2}`, + backgroundColor: tokens.colorNeutralBackground1, + }, + nameCell: { + verticalAlign: 'top', + }, + clickableRow: { + cursor: 'pointer', + '&:hover': { + backgroundColor: tokens.colorNeutralBackground1Hover, + }, + '&:focus-within': { + backgroundColor: tokens.colorNeutralBackground1Hover, + }, + }, + actionCell: { + width: '7rem', + verticalAlign: 'top', + }, + dialogBody: { + display: 'flex', + flexDirection: 'column', + gap: tokens.spacingVerticalM, + }, + sourceDialog: { + width: 'min(60rem, 90vw)', + maxWidth: 'none', + }, + emptyState: { + padding: tokens.spacingVerticalXL, + color: tokens.colorNeutralForeground3, + }, +}) \ No newline at end of file diff --git a/frontend/src/components/Initializers/CustomInitializers.tsx b/frontend/src/components/Initializers/CustomInitializers.tsx new file mode 100644 index 0000000000..02ae091d3f --- /dev/null +++ b/frontend/src/components/Initializers/CustomInitializers.tsx @@ -0,0 +1,177 @@ +import { useState } from 'react' + +import { + Button, + Dialog, + DialogActions, + DialogBody, + DialogContent, + DialogSurface, + DialogTitle, + DialogTrigger, + Field, + Input, + Table, + TableBody, + TableCell, + TableHeader, + TableHeaderCell, + TableRow, + Text, +} from '@fluentui/react-components' +import { AddRegular, DeleteRegular, EyeRegular } from '@fluentui/react-icons' + +import type { CustomInitializer } from '@/types' + +import { useCustomInitializersStyles } from './CustomInitializers.styles' +import { PythonCodeBlock, PythonCodeEditor } from './PythonCode' + +interface CustomInitializersProps { + items: CustomInitializer[] + registering: boolean + deletingName: string | null + onRegister: (name: string, scriptContent: string) => Promise + onDelete: (name: string) => Promise +} + +export default function CustomInitializers({ + items, + registering, + deletingName, + onRegister, + onDelete, +}: CustomInitializersProps) { + const styles = useCustomInitializersStyles() + const [dialogOpen, setDialogOpen] = useState(false) + const [name, setName] = useState('') + const [scriptContent, setScriptContent] = useState('') + const [viewingInitializer, setViewingInitializer] = useState(null) + + const handleSubmit = async (event: React.FormEvent): Promise => { + event.preventDefault() + if (await onRegister(name.trim(), scriptContent)) { + setDialogOpen(false) + setName('') + setScriptContent('') + } + } + + return ( +
+
+
+ + Custom initializers + + Persisted Python definitions available to startup initializer configuration. +
+ setDialogOpen(data.open)}> + + + + +
+ + Register custom initializer + + + setName(data.value)} + disabled={registering} + autoComplete="off" + /> + + + + + + + + + + + + +
+
+
+
+ + {items.length === 0 ? ( + No custom initializers registered. + ) : ( +
+ + + + Name + Actions + + + + {items.map((item) => ( + setViewingInitializer(item)} + > + + + + + + + + ))} + +
+
+ )} + + { + if (!data.open) { + setViewingInitializer(null) + } + }} + > + + + {viewingInitializer?.initializer_name} + + + + + + + + + +
+ ) +} \ No newline at end of file diff --git a/frontend/src/components/Initializers/Initializers.test.tsx b/frontend/src/components/Initializers/Initializers.test.tsx index 3ddc8b6954..4cc7726cae 100644 --- a/frontend/src/components/Initializers/Initializers.test.tsx +++ b/frontend/src/components/Initializers/Initializers.test.tsx @@ -8,6 +8,7 @@ import type { BaselineInitializerSetting, InitializerSettingsResponse, RegisteredInitializer, + CustomInitializer, } from '@/types' import Initializers from './Initializers' @@ -16,6 +17,9 @@ jest.mock('@/services/api', () => ({ initializersApi: { getSettings: jest.fn(), listRegistered: jest.fn(), + listCustom: jest.fn(), + register: jest.fn(), + unregister: jest.fn(), createAdditional: jest.fn(), updateAdditional: jest.fn(), deleteAdditional: jest.fn(), @@ -96,6 +100,11 @@ const sampleSettings: InitializerSettingsResponse = { additional: [additionalItem], } +const customInitializer: CustomInitializer = { + initializer_name: 'custom_target', + script_content: 'class CustomTargetInitializer: pass', +} + function renderInitializers(): void { render( @@ -112,6 +121,15 @@ describe('Initializers', () => { items: [targetInitializer, scorerInitializer], pagination: { limit: 200, has_more: false }, }) + mockedInitializersApi.listCustom.mockResolvedValue([customInitializer]) + mockedInitializersApi.register.mockResolvedValue({ + initializer_name: 'new_custom', + initializer_type: 'NewCustomInitializer', + description: 'New custom initializer.', + required_env_vars: [], + supported_parameters: [], + }) + mockedInitializersApi.unregister.mockResolvedValue() mockedInitializersApi.createAdditional.mockResolvedValue({ id: 'additional-2', initializer_name: 'target', @@ -150,6 +168,47 @@ describe('Initializers', () => { expect(screen.getByTestId('initializer-row-additional-1')).toHaveTextContent('scorer') }) + it('should list and register persisted custom initializers in the custom tab', async () => { + const user = userEvent.setup() + const writeText = jest.fn().mockResolvedValue(undefined) + Object.defineProperty(navigator, 'clipboard', { configurable: true, value: { writeText } }) + renderInitializers() + + await user.click(await screen.findByRole('tab', { name: 'Custom' })) + const customRow = await screen.findByTestId('custom-initializer-custom_target') + expect(customRow).not.toHaveTextContent('class CustomTargetInitializer: pass') + fireEvent.click(within(customRow).getByRole('button', { name: /custom_target/ })) + const sourceDialog = await screen.findByRole('dialog') + const sourceBlock = within(sourceDialog).getByLabelText('Python source') + expect(sourceBlock).toHaveTextContent( + 'class CustomTargetInitializer: pass', + ) + expect(sourceBlock.querySelector('.token.keyword')).toHaveTextContent('class') + fireEvent.click(within(sourceDialog).getByRole('button', { name: 'Copy Python source', hidden: true })) + await waitFor(() => expect(writeText).toHaveBeenCalledWith('class CustomTargetInitializer: pass')) + await user.click(within(sourceDialog).getByRole('button', { name: 'Close', hidden: true })) + + fireEvent.click(screen.getByRole('button', { name: 'Register initializer' })) + const dialog = await screen.findByRole('dialog') + fireEvent.change(within(dialog).getByLabelText(/Initializer name/), { target: { value: 'new_custom' } }) + fireEvent.change(within(dialog).getByLabelText('Python source', { selector: 'textarea' }), { + target: { value: 'class NewCustom: pass' }, + }) + expect(within(dialog).getAllByRole('textbox')).toHaveLength(2) + expect(dialog.querySelector('pre[aria-hidden="true"] .token.keyword')).toHaveTextContent('class') + const form = dialog.querySelector('form') + expect(form).not.toBeNull() + fireEvent.submit(form as HTMLFormElement) + + await waitFor(() => { + expect(mockedInitializersApi.register).toHaveBeenCalledWith({ + name: 'new_custom', + script_content: 'class NewCustom: pass', + }) + expect(screen.getByText('Registered new_custom.')).toBeInTheDocument() + }) + }) + it('should refresh settings when the refresh button is clicked', async () => { const user = userEvent.setup() renderInitializers() diff --git a/frontend/src/components/Initializers/Initializers.tsx b/frontend/src/components/Initializers/Initializers.tsx index 55c7547d7f..7d69f15a6f 100644 --- a/frontend/src/components/Initializers/Initializers.tsx +++ b/frontend/src/components/Initializers/Initializers.tsx @@ -1,15 +1,22 @@ import { useEffect, useState } from 'react' -import { Button, MessageBar, MessageBarBody, Spinner, Text } from '@fluentui/react-components' +import { Button, MessageBar, MessageBarBody, Spinner, Tab, TabList, Text } from '@fluentui/react-components' +import type { SelectTabData, SelectTabEvent } from '@fluentui/react-components' import { ArrowSyncRegular } from '@fluentui/react-icons' import { initializersApi } from '@/services/api' import { toApiError } from '@/services/errors' -import type { InitializerSettingsResponse, RegisteredInitializer, UpdateAdditionalInitializerRequest } from '@/types' +import type { + CustomInitializer, + InitializerSettingsResponse, + RegisteredInitializer, + UpdateAdditionalInitializerRequest, +} from '@/types' import AdditionalInitializers from './AdditionalInitializers' import AvailableInitializersDialog from './AvailableInitializersDialog' import BaselineInitializers from './BaselineInitializers' +import CustomInitializers from './CustomInitializers' import { useInitializersStyles } from './Initializers.styles' interface StatusMessage { @@ -17,6 +24,8 @@ interface StatusMessage { text: string } +type InitializerTab = 'startup' | 'custom' + const EMPTY_SETTINGS: InitializerSettingsResponse = { baseline: [], additional: [], @@ -26,6 +35,8 @@ export default function Initializers() { const styles = useInitializersStyles() const [settings, setSettings] = useState(EMPTY_SETTINGS) const [registeredInitializers, setRegisteredInitializers] = useState([]) + const [customInitializers, setCustomInitializers] = useState([]) + const [selectedTab, setSelectedTab] = useState('startup') const [loading, setLoading] = useState(true) const [statusMessage, setStatusMessage] = useState(null) const [refetchCount, setRefetchCount] = useState(0) @@ -34,14 +45,17 @@ export default function Initializers() { const [saveErrors, setSaveErrors] = useState>({}) const [applyingInitializerId, setApplyingInitializerId] = useState(null) const [deletingInitializerId, setDeletingInitializerId] = useState(null) + const [registeringCustom, setRegisteringCustom] = useState(false) + const [deletingCustomName, setDeletingCustomName] = useState(null) useEffect(() => { let cancelled = false const loadInitializersAsync = async (): Promise => { - const [settingsResult, registeredResult] = await Promise.allSettled([ + const [settingsResult, registeredResult, customResult] = await Promise.allSettled([ initializersApi.getSettings(), initializersApi.listRegistered(), + initializersApi.listCustom(), ]) if (cancelled) { return @@ -64,6 +78,12 @@ export default function Initializers() { ) } + if (customResult.status === 'fulfilled') { + setCustomInitializers(customResult.value) + } else { + setStatusMessage({ intent: 'error', text: toApiError(customResult.reason).detail }) + } + setLoading(false) } @@ -84,6 +104,49 @@ export default function Initializers() { setSettings(response) } + const refetchCustomCatalog = async (): Promise => { + const [custom, registered] = await Promise.all([ + initializersApi.listCustom(), + initializersApi.listRegistered(), + ]) + setCustomInitializers(custom) + setRegisteredInitializers(registered.items) + } + + const handleTabSelect = (_: SelectTabEvent, data: SelectTabData): void => { + if (data.value === 'startup' || data.value === 'custom') { + setSelectedTab(data.value) + } + } + + const handleRegisterCustom = async (name: string, scriptContent: string): Promise => { + setRegisteringCustom(true) + try { + await initializersApi.register({ name, script_content: scriptContent }) + await refetchCustomCatalog() + setStatusMessage({ intent: 'success', text: `Registered ${name}.` }) + return true + } catch (error) { + setStatusMessage({ intent: 'error', text: toApiError(error).detail }) + return false + } finally { + setRegisteringCustom(false) + } + } + + const handleDeleteCustom = async (name: string): Promise => { + setDeletingCustomName(name) + try { + await initializersApi.unregister(name) + await refetchCustomCatalog() + setStatusMessage({ intent: 'success', text: `Removed ${name}.` }) + } catch (error) { + setStatusMessage({ intent: 'error', text: toApiError(error).detail }) + } finally { + setDeletingCustomName(null) + } + } + const handleAdd = async ( initializerName: string, parameters: Record | null, @@ -171,15 +234,16 @@ export default function Initializers() {
Initializers - Browse every registered initializer, review the read-only baseline that ran at startup, and manage - additional initializer invocations that run after it. + Manage startup initializer invocations and persisted custom initializer definitions.
- + {selectedTab === 'startup' && ( + + )}
+ ) +} + +export function PythonCodeBlock({ source, ariaLabel }: PythonCodeBlockProps) { + const styles = usePythonCodeStyles() + + return ( +
+ +
+        
+      
+
+ ) +} + +export function PythonCodeEditor({ source, disabled, onChange }: PythonCodeEditorProps) { + const styles = usePythonCodeStyles() + const highlightRef = useRef(null) + + const handleScroll = (event: React.UIEvent): void => { + if (highlightRef.current) { + highlightRef.current.scrollTop = event.currentTarget.scrollTop + highlightRef.current.scrollLeft = event.currentTarget.scrollLeft + } + } + + return ( +
+ +
+ +