From 3506e691a7816be6ffc47634affcf058ae27c545 Mon Sep 17 00:00:00 2001 From: YouGottaHackThat Date: Thu, 2 Jul 2026 19:19:08 +0100 Subject: [PATCH 1/4] Prevent duplicate collab saves --- javascript/src/collab_server/base_handler.ts | 35 +++--- .../handlers/report_finding_link.ts | 85 +++++++++------ javascript/src/collab_server/index.ts | 103 +++++++++++++++--- 3 files changed, 166 insertions(+), 57 deletions(-) diff --git a/javascript/src/collab_server/base_handler.ts b/javascript/src/collab_server/base_handler.ts index cafcffb84..8415be069 100644 --- a/javascript/src/collab_server/base_handler.ts +++ b/javascript/src/collab_server/base_handler.ts @@ -1,15 +1,16 @@ import * as Y from "yjs"; -import { ApolloClient, TypedDocumentNode } from "@apollo/client"; +import { + ApolloClient, + OperationVariables, + TypedDocumentNode, +} from "@apollo/client"; /** Functions for loading, saving, and converting a type from/to YJS. */ -export type ModelHandler = { +export type ModelHandler = { load: (client: ApolloClient, id: number) => Promise<[Y.Doc, T]>; - save: ( - client: ApolloClient, - id: number, - doc: Y.Doc, - data: T - ) => Promise; + getSavePayload: (id: number, doc: Y.Doc, data: T) => P; + save: (client: ApolloClient, payload: P, data: T) => Promise; + markSaved?: (payload: P, data: T) => void; }; export type IdVars = { id: number }; @@ -22,12 +23,17 @@ export type IdVars = { id: number }; * @param mkQueryVars Function to get the parameters for the `setQuery` to save the model. Called in a YJS transaction. * @returns The model handler. */ -export function simpleModelHandler( +export function simpleModelHandler< + GetRes, + SetRes, + SetQueryVars extends OperationVariables, + T, +>( getQuery: TypedDocumentNode, setQuery: TypedDocumentNode, fillFields: (doc: Y.Doc, res: GetRes) => T, mkQueryVars: (doc: Y.Doc, id: number, data: T) => SetQueryVars -): ModelHandler { +): ModelHandler { return { load: async (client, id) => { const res = await client.query({ @@ -46,14 +52,17 @@ export function simpleModelHandler( }); return [doc, data!]; }, - async save(client, id, doc, data) { - let queryVars; + getSavePayload(id, doc, data) { + let queryVars: SetQueryVars | undefined; doc.transact(() => { queryVars = mkQueryVars(doc, id, data); }); + return queryVars!; + }, + async save(client, payload) { const res = await client.mutate({ mutation: setQuery, - variables: queryVars, + variables: payload, }); if (res.errors) { throw res.errors; diff --git a/javascript/src/collab_server/handlers/report_finding_link.ts b/javascript/src/collab_server/handlers/report_finding_link.ts index d5f4ed4ea..f677c652b 100644 --- a/javascript/src/collab_server/handlers/report_finding_link.ts +++ b/javascript/src/collab_server/handlers/report_finding_link.ts @@ -4,6 +4,11 @@ import * as Y from "yjs"; import { htmlToYjs, tagsToYjs, yjsToHtml, yjsToTags } from "../yjs_converters"; import { extraFieldsFromYdoc, extraFieldsToYdoc } from "../extra_fields"; +type ReportFindingLinkData = { + extraFieldSpec: { internalName: string; type: string }[]; + savedSeverityId: number | null; +}; + const GET = gql(` query GET_REPORT_FINDING_LINK($id: bigint!) { reportedFinding_by_pk(id: $id) { @@ -86,43 +91,61 @@ const ReportFindingLinkHandler = simpleModelHandler( ); tagsToYjs(res.tags.tags, doc.get("tags", Y.Map)); extraFieldsToYdoc(res.extraFieldSpec, doc, obj.extraFields); - return res.extraFieldSpec; + return { + extraFieldSpec: res.extraFieldSpec, + savedSeverityId: obj.severity.id, + } satisfies ReportFindingLinkData; }, - (doc, id, extraFieldSpec) => { + (doc, id, data) => { const plainFields = doc.get("plain_fields", Y.Map); - const extraFields = extraFieldsFromYdoc(extraFieldSpec, doc); + const extraFields = extraFieldsFromYdoc(data.extraFieldSpec, doc); + const severityId = plainFields.get("severityId") ?? null; + const set: Record = { + title: plainFields.get("title") ?? "", + cvssScore: plainFields.get("cvssScore") ?? null, + cvssVector: plainFields.get("cvssVector") ?? "", + findingTypeId: plainFields.get("findingTypeId"), + + description: yjsToHtml(doc.get("description", Y.XmlFragment)), + impact: yjsToHtml(doc.get("impact", Y.XmlFragment)), + mitigation: yjsToHtml(doc.get("mitigation", Y.XmlFragment)), + replication_steps: yjsToHtml( + doc.get("replicationSteps", Y.XmlFragment) + ), + hostDetectionTechniques: yjsToHtml( + doc.get("hostDetectionTechniques", Y.XmlFragment) + ), + networkDetectionTechniques: yjsToHtml( + doc.get("networkDetectionTechniques", Y.XmlFragment) + ), + references: yjsToHtml(doc.get("references", Y.XmlFragment)), + findingGuidance: yjsToHtml( + doc.get("findingGuidance", Y.XmlFragment) + ), + affectedEntities: yjsToHtml( + doc.get("affectedEntities", Y.XmlFragment) + ), + extraFields, + }; + + if (severityId !== data.savedSeverityId) { + set.severityId = severityId; + } + return { id, - set: { - title: plainFields.get("title") ?? "", - cvssScore: plainFields.get("cvssScore") ?? null, - cvssVector: plainFields.get("cvssVector") ?? "", - findingTypeId: plainFields.get("findingTypeId"), - severityId: plainFields.get("severityId"), - - description: yjsToHtml(doc.get("description", Y.XmlFragment)), - impact: yjsToHtml(doc.get("impact", Y.XmlFragment)), - mitigation: yjsToHtml(doc.get("mitigation", Y.XmlFragment)), - replication_steps: yjsToHtml( - doc.get("replicationSteps", Y.XmlFragment) - ), - hostDetectionTechniques: yjsToHtml( - doc.get("hostDetectionTechniques", Y.XmlFragment) - ), - networkDetectionTechniques: yjsToHtml( - doc.get("networkDetectionTechniques", Y.XmlFragment) - ), - references: yjsToHtml(doc.get("references", Y.XmlFragment)), - findingGuidance: yjsToHtml( - doc.get("findingGuidance", Y.XmlFragment) - ), - affectedEntities: yjsToHtml( - doc.get("affectedEntities", Y.XmlFragment) - ), - extraFields, - }, + set, tags: yjsToTags(doc.get("tags", Y.Map)), }; } ); +ReportFindingLinkHandler.markSaved = (payload, data) => { + const variables = payload as { set?: { severityId?: number | null } }; + if ( + variables.set && + Object.prototype.hasOwnProperty.call(variables.set, "severityId") + ) { + data.savedSeverityId = variables.set.severityId ?? null; + } +}; export default ReportFindingLinkHandler; diff --git a/javascript/src/collab_server/index.ts b/javascript/src/collab_server/index.ts index 54b0504cf..716294432 100644 --- a/javascript/src/collab_server/index.ts +++ b/javascript/src/collab_server/index.ts @@ -23,7 +23,7 @@ import ReportHandler from "./handlers/report"; import ProjectHandler from "./handlers/project"; // Extend this with your model handlers. See how-to-collab.md. -const HANDLERS_ARR: [string, ModelHandler][] = [ +const HANDLERS_ARR: [string, ModelHandler][] = [ ["observation", ObservationHandler], ["report_observation_link", ReportObservationLinkHandler], ["finding", FindingHandler], @@ -31,14 +31,15 @@ const HANDLERS_ARR: [string, ModelHandler][] = [ ["report", ReportHandler], ["project", ProjectHandler], ]; -const HANDLERS: Map> = new Map(HANDLERS_ARR); +const HANDLERS: Map> = new Map(HANDLERS_ARR); // Graphql Client -const graphql_engine_hostname: string = env["HASURA_GRAPHQL_SERVER_HOSTNAME"] || "graphql_engine"; +const graphql_engine_hostname: string = + env["HASURA_GRAPHQL_SERVER_HOSTNAME"] || "graphql_engine"; const httpLink = createHttpLink({ - uri: "http://" + graphql_engine_hostname + ":8080/v1/graphql" + uri: "http://" + graphql_engine_hostname + ":8080/v1/graphql", }); const authLink = setContext((_, { headers }) => { @@ -86,6 +87,45 @@ class AuthError extends Error { const BASE_LOGGER = pino({}); const documentData = new Map(); +const lastSavedSignatures = new Map(); +const lastFailedSignatures = new Map(); + +function stableStringify(value: unknown): string { + return JSON.stringify(normalizeForSignature(value)); +} + +function normalizeForSignature(value: unknown): unknown { + if (Array.isArray(value)) { + return value.map(normalizeForSignature); + } + + if (value !== null && typeof value === "object") { + return Object.fromEntries( + Object.entries(value as Record) + .filter(([, entryValue]) => entryValue !== undefined) + .sort(([leftKey], [rightKey]) => + leftKey.localeCompare(rightKey) + ) + .map(([entryKey, entryValue]) => [ + entryKey, + normalizeForSignature(entryValue), + ]) + ); + } + + return value; +} + +function setSaveError(doc: Y.Doc, value: boolean) { + const serverInfo = doc.get("serverInfo", Y.Map); + if (serverInfo.get("saveError") === value) { + return; + } + + doc.transact(() => { + serverInfo.set("saveError", value); + }); +} const server = new Server({ port: 8000, @@ -227,6 +267,15 @@ const server = new Server({ serverInfo.set("saveError", false); }); documentData.set(data.documentName, docData); + const initialPayload = handler.getSavePayload( + context.id, + doc, + docData + ); + lastSavedSignatures.set( + data.documentName, + stableStringify(initialPayload) + ); return doc; } catch (e) { context.log.error({ msg: "Could not load document", err: e }); @@ -236,21 +285,47 @@ const server = new Server({ async onStoreDocument(data) { const context = data.context as Context; + const docData = documentData.get(data.documentName); + const handler = HANDLERS.get(context.model)!; + const payload = handler.getSavePayload( + context.id, + data.document, + docData + ); + const signature = stableStringify(payload); + + if (lastSavedSignatures.get(data.documentName) === signature) { + context.log.info("Skipping unchanged document save"); + setSaveError(data.document, false); + return; + } + + if (lastFailedSignatures.get(data.documentName) === signature) { + context.log.warn( + "Skipping retry for unchanged failed document save" + ); + setSaveError(data.document, true); + return; + } + try { - const docData = documentData.get(data.documentName); context.log.info("Saving document"); - const handler = HANDLERS.get(context.model)!; - await handler.save(gqlClient, context.id, data.document, docData); + await handler.save(gqlClient, payload, docData); + handler.markSaved?.(payload, docData); + lastSavedSignatures.set( + data.documentName, + stableStringify( + handler.getSavePayload(context.id, data.document, docData) + ) + ); + lastFailedSignatures.delete(data.documentName); } catch (e) { context.log.error({ msg: "Could not save document", err: e }); - data.document.transact((tx) => { - tx.doc.get("serverInfo", Y.Map).set("saveError", true); - }); + lastFailedSignatures.set(data.documentName, signature); + setSaveError(data.document, true); return; } - data.document.transact((tx) => { - tx.doc.get("serverInfo", Y.Map).set("saveError", false); - }); + setSaveError(data.document, false); }, async onDisconnect(data) { @@ -259,6 +334,8 @@ const server = new Server({ async afterUnloadDocument(data) { documentData.delete(data.documentName); + lastSavedSignatures.delete(data.documentName); + lastFailedSignatures.delete(data.documentName); }, }); From 76798c71a30eabb3a993fe5a8d0da808c49998e5 Mon Sep 17 00:00:00 2001 From: YouGottaHackThat Date: Thu, 2 Jul 2026 22:01:23 +0100 Subject: [PATCH 2/4] Stabilise report finding position events --- ghostwriter/api/tests/test_views.py | 336 +++++++++++++++++----------- ghostwriter/api/views.py | 21 +- ghostwriter/modules/model_utils.py | 155 +++++++++---- 3 files changed, 321 insertions(+), 191 deletions(-) diff --git a/ghostwriter/api/tests/test_views.py b/ghostwriter/api/tests/test_views.py index c724b43a3..4106eaff9 100644 --- a/ghostwriter/api/tests/test_views.py +++ b/ghostwriter/api/tests/test_views.py @@ -2977,6 +2977,41 @@ def setUpTestData(cls): def setUp(self): self.client = Client() + def assert_position_sequence(self, severity, expected_findings): + """Assert a severity group has deterministic contiguous positions.""" + for expected_position, finding in enumerate(expected_findings, start=1): + finding.refresh_from_db() + self.assertEqual(finding.severity_id, severity.id) + self.assertEqual(finding.position, expected_position) + + actual_positions = list( + self.ReportFindingLink.objects.filter( + report=self.report, + severity=severity, + ) + .order_by("position", "id") + .values_list("position", flat=True) + ) + self.assertEqual(actual_positions, list(range(1, len(actual_positions) + 1))) + + def post_change_event(self, op, old, new): + return self.client.post( + self.change_uri, + content_type="application/json", + data={ + "event": { + "op": op, + "data": { + "old": old, + "new": new, + }, + }, + }, + **{ + "HTTP_HASURA_ACTION_SECRET": f"{ACTION_SECRET}", + }, + ) + def test_model_cleaning_position(self): self.ReportFindingLink.objects.all().delete() first_finding = ReportFindingLinkFactory( @@ -2990,112 +3025,182 @@ def test_model_cleaning_position(self): ) # Simulate an event changing the position of the first finding to `3` + old_position = first_finding.position first_finding.position = 3 first_finding.save() - sample_data = { - "event": { - "op": "UPDATE", - "data": { - "old": { - "id": first_finding.id, - "position": 1, - "severity_id": first_finding.severity.id, - }, - "new": { - "id": first_finding.id, - "position": 3, - "severity_id": first_finding.severity.id, - }, - }, + response = self.post_change_event( + "UPDATE", + { + "id": first_finding.id, + "position": old_position, + "severity_id": first_finding.severity.id, }, - } - - # Submit the POST request to the event webhook - response = self.client.post( - self.change_uri, - content_type="application/json", - data=sample_data, - **{ - "HTTP_HASURA_ACTION_SECRET": f"{ACTION_SECRET}", + { + "id": first_finding.id, + "position": first_finding.position, + "severity_id": first_finding.severity.id, }, ) self.assertEqual(response.status_code, 200) - first_finding.refresh_from_db() - self.assertEqual(first_finding.position, 3) - second_finding.refresh_from_db() - self.assertEqual(second_finding.position, 1) - third_finding.refresh_from_db() - self.assertEqual(third_finding.position, 2) + self.assert_position_sequence( + self.critical_severity, + [second_finding, third_finding, first_finding], + ) # Repeat for an `UPDATE` event with a severity change + old_position = second_finding.position second_finding.severity = self.high_severity second_finding.save() - sample_data = { - "event": { - "op": "UPDATE", - "data": { - "old": { - "id": second_finding.id, - "position": second_finding.position, - "severity_id": self.critical_severity.id, - }, - "new": { - "id": second_finding.id, - "position": second_finding.position, - "severity_id": self.high_severity.id, - }, - }, + response = self.post_change_event( + "UPDATE", + { + "id": second_finding.id, + "position": old_position, + "severity_id": self.critical_severity.id, }, - } - - response = self.client.post( - self.change_uri, - content_type="application/json", - data=sample_data, - **{ - "HTTP_HASURA_ACTION_SECRET": f"{ACTION_SECRET}", + { + "id": second_finding.id, + "position": old_position, + "severity_id": self.high_severity.id, }, ) self.assertEqual(response.status_code, 200) - first_finding.refresh_from_db() - second_finding.refresh_from_db() - third_finding.refresh_from_db() - self.assertEqual(second_finding.position, 1) - self.assertEqual(first_finding.position, 2) - self.assertEqual(third_finding.position, 1) + self.assert_position_sequence( + self.critical_severity, + [third_finding, first_finding], + ) + self.assert_position_sequence(self.high_severity, [second_finding]) # Repeat for an `INSERT` event new_finding = ReportFindingLinkFactory( report=self.report, severity=self.critical_severity ) - sample_data = { - "event": { - "op": "INSERT", - "data": { - "old": None, - "new": { - "id": new_finding.id, - "position": 1, - "severity_id": new_finding.severity.id, - }, - }, + response = self.post_change_event( + "INSERT", + None, + { + "id": new_finding.id, + "position": new_finding.position, + "severity_id": new_finding.severity.id, }, - } + ) - response = self.client.post( - self.change_uri, - content_type="application/json", - data=sample_data, - **{ - "HTTP_HASURA_ACTION_SECRET": f"{ACTION_SECRET}", + self.assertEqual(response.status_code, 200) + self.assert_position_sequence( + self.critical_severity, + [third_finding, first_finding, new_finding], + ) + + def test_duplicate_insert_positions_converge_and_follow_up_events_noop(self): + self.ReportFindingLink.objects.all().delete() + first_finding = ReportFindingLinkFactory( + report=self.report, severity=self.critical_severity, position=1 + ) + second_finding = ReportFindingLinkFactory( + report=self.report, severity=self.critical_severity, position=1 + ) + third_finding = ReportFindingLinkFactory( + report=self.report, severity=self.critical_severity, position=1 + ) + + response = self.post_change_event( + "INSERT", + None, + { + "id": third_finding.id, + "position": third_finding.position, + "severity_id": third_finding.severity.id, + }, + ) + + self.assertEqual(response.status_code, 200) + self.assert_position_sequence( + self.critical_severity, + [first_finding, second_finding, third_finding], + ) + + positions_before_follow_up = list( + self.ReportFindingLink.objects.filter(report=self.report) + .order_by("id") + .values_list("id", "position") + ) + response = self.post_change_event( + "UPDATE", + { + "id": second_finding.id, + "position": 1, + "severity_id": second_finding.severity.id, + }, + { + "id": second_finding.id, + "position": 2, + "severity_id": second_finding.severity.id, + }, + ) + + self.assertEqual(response.status_code, 200) + positions_after_follow_up = list( + self.ReportFindingLink.objects.filter(report=self.report) + .order_by("id") + .values_list("id", "position") + ) + self.assertEqual(positions_after_follow_up, positions_before_follow_up) + + response = self.post_change_event( + "UPDATE", + { + "id": first_finding.id, + "position": 1, + "severity_id": first_finding.severity.id, + }, + { + "id": first_finding.id, + "position": 3, + "severity_id": first_finding.severity.id, + }, + ) + + self.assertEqual(response.status_code, 200) + positions_after_stale_follow_up = list( + self.ReportFindingLink.objects.filter(report=self.report) + .order_by("id") + .values_list("id", "position") + ) + self.assertEqual(positions_after_stale_follow_up, positions_before_follow_up) + + def test_duplicate_position_update_honors_requested_target(self): + self.ReportFindingLink.objects.all().delete() + first_finding = ReportFindingLinkFactory( + report=self.report, severity=self.critical_severity, position=1 + ) + second_finding = ReportFindingLinkFactory( + report=self.report, severity=self.critical_severity, position=2 + ) + third_finding = ReportFindingLinkFactory( + report=self.report, severity=self.critical_severity, position=2 + ) + + response = self.post_change_event( + "UPDATE", + { + "id": third_finding.id, + "position": 1, + "severity_id": third_finding.severity.id, + }, + { + "id": third_finding.id, + "position": 2, + "severity_id": third_finding.severity.id, }, ) self.assertEqual(response.status_code, 200) - new_finding.refresh_from_db() - self.assertEqual(new_finding.position, 3) + self.assert_position_sequence( + self.critical_severity, + [first_finding, third_finding, second_finding], + ) def test_position_set_to_zero(self): self.ReportFindingLink.objects.all().delete() @@ -3104,33 +3209,18 @@ def test_position_set_to_zero(self): ) # Simulate an event changing the position of the first finding to `0` - sample_data = { - "event": { - "op": "INSERT", - "data": { - "old": None, - "new": { - "id": finding.id, - "position": 0, - "severity_id": finding.severity.id, - }, - }, - }, - } - - # Submit the POST request to the event webhook - response = self.client.post( - self.change_uri, - content_type="application/json", - data=sample_data, - **{ - "HTTP_HASURA_ACTION_SECRET": f"{ACTION_SECRET}", + response = self.post_change_event( + "INSERT", + None, + { + "id": finding.id, + "position": 0, + "severity_id": finding.severity.id, }, ) self.assertEqual(response.status_code, 200) - finding.refresh_from_db() - self.assertEqual(finding.position, 1) + self.assert_position_sequence(self.critical_severity, [finding]) def test_position_set_higher_than_count(self): self.ReportFindingLink.objects.all().delete() @@ -3139,36 +3229,18 @@ def test_position_set_higher_than_count(self): ) # Simulate an event changing the position of the first finding to `100` - sample_data = { - "event": { - "op": "INSERT", - "data": { - "old": None, - "new": { - "id": finding.id, - "position": 100, - "severity_id": finding.severity.id, - }, - }, - }, - } - - # Submit the POST request to the event webhook - response = self.client.post( - self.change_uri, - content_type="application/json", - data=sample_data, - **{ - "HTTP_HASURA_ACTION_SECRET": f"{ACTION_SECRET}", + response = self.post_change_event( + "INSERT", + None, + { + "id": finding.id, + "position": 100, + "severity_id": finding.severity.id, }, ) self.assertEqual(response.status_code, 200) - finding.refresh_from_db() - total_findings = self.ReportFindingLink.objects.filter( - report=self.report - ).count() - self.assertEqual(finding.position, total_findings) + self.assert_position_sequence(self.critical_severity, [finding]) def test_position_change_on_delete(self): self.ReportFindingLink.objects.all().delete() @@ -3211,10 +3283,10 @@ def test_position_change_on_delete(self): ) self.assertEqual(response.status_code, 200) - first_finding.refresh_from_db() - third_finding.refresh_from_db() - self.assertEqual(first_finding.position, 1) - self.assertEqual(third_finding.position, 2) + self.assert_position_sequence( + self.critical_severity, + [first_finding, third_finding], + ) class GraphqlProjectContactUpdateEventTests(TestCase): diff --git a/ghostwriter/api/views.py b/ghostwriter/api/views.py index d3718d19c..8a908b6aa 100644 --- a/ghostwriter/api/views.py +++ b/ghostwriter/api/views.py @@ -54,7 +54,11 @@ ) from ghostwriter.commandcenter.models import ExtraFieldModel, GeneralConfiguration from ghostwriter.modules import codenames -from ghostwriter.modules.model_utils import set_finding_positions, to_dict +from ghostwriter.modules.model_utils import ( + normalize_finding_positions, + set_finding_positions, + to_dict, +) from ghostwriter.modules.passive_voice.detector import get_detector from ghostwriter.modules.reportwriter.report.json import ExportReportJson from ghostwriter.oplog.models import OplogEntry, OplogEntryEvidence, OplogEntryRecording @@ -1843,16 +1847,11 @@ class GraphqlReportFindingDeleteEvent(HasuraEventView): def post(self, request, *args, **kwargs): try: - findings_queryset = ReportFindingLink.objects.filter( - Q(report=self.old_data["report_id"]) - & Q(severity=self.old_data["severity_id"]) - ) - if findings_queryset: - counter = 1 - for finding in findings_queryset: - # Adjust position to close gap created by the removed finding - findings_queryset.filter(id=finding.id).update(position=counter) - counter += 1 + normalize_finding_positions( + ReportFindingLink, + self.old_data["report_id"], + self.old_data["severity_id"], + ) except Report.DoesNotExist: # pragma: no cover # Report was deleted, so no need to adjust positions pass diff --git a/ghostwriter/modules/model_utils.py b/ghostwriter/modules/model_utils.py index cd4ef87f7..fc43c5125 100644 --- a/ghostwriter/modules/model_utils.py +++ b/ghostwriter/modules/model_utils.py @@ -2,9 +2,11 @@ # Standard Libraries from itertools import chain +from typing import Optional, Type # Django Imports import django +from django.db import transaction from django.db.models import ForeignKey, Q @@ -38,6 +40,61 @@ def to_dict(instance: django.db.models.Model, include_id: bool = False, resolve_ return data +def _clamp_position(position: Optional[int], count: int) -> int: + """Return a one-based position that fits inside a group of findings.""" + if count < 1: + return 1 + if position is None: + return count + return min(max(position, 1), count) + + +def normalize_finding_positions( + model: Type[django.db.models.Model], + report_id: int, + severity_id: int, + moving_instance_id: Optional[int] = None, + target_position: Optional[int] = None, +) -> None: + """ + Normalise report finding positions for a single report/severity group. + + Ordering uses ``position, id`` so duplicate positions converge + deterministically. Updates are emitted only for rows that need to change, + allowing Hasura events generated by this function to become no-ops. + """ + report_model = model._meta.get_field("report").related_model + with transaction.atomic(): + # Locking the report row serialises reorder work for every severity + # group in that report, including callers that only normalise one group. + report_model.objects.select_for_update().get(id=report_id) + + findings = list( + model.objects.select_for_update() + .filter(Q(report_id=report_id) & Q(severity_id=severity_id)) + .order_by("position", "id") + ) + if not findings: + return + + if moving_instance_id is not None: + moving_finding = next( + (finding for finding in findings if finding.id == moving_instance_id), + None, + ) + if moving_finding is not None: + ordered_findings = [ + finding for finding in findings if finding.id != moving_instance_id + ] + insert_at = _clamp_position(target_position, len(findings)) - 1 + ordered_findings.insert(insert_at, moving_finding) + findings = ordered_findings + + for position, finding in enumerate(findings, start=1): + if finding.position != position: + model.objects.filter(id=finding.id).update(position=position) + + def set_finding_positions( instance: django.db.models.Model, old_pos: [int, None], old_sev: [int, None], new_pos: int, new_sev: int ) -> None: @@ -45,8 +102,9 @@ def set_finding_positions( Updates the ``position`` value for a finding in a report. This is used when a finding is moved to a new position or changes severity. - The following adjustments use the queryset ``update()`` method (direct SQL statement) instead of calling ``save()`` - on the individual model instance. This avoids forever looping through position changes. + Reorder processing is serialised on the parent report row. The position + updates use ``update()`` and are only issued when a row's position must + change, so Hasura events generated by normalisation converge to no-ops. **Parameters** @@ -63,53 +121,54 @@ def set_finding_positions( """ # We don't import the model at the top of the file because it causes a circular import model = instance._meta.model + report_model = instance.report._meta.model + report_id = instance.report_id + + with transaction.atomic(): + # Serialise all reorder work for a report so concurrent Hasura events + # cannot interleave read/renumber/write cycles for the same findings. + report_model.objects.select_for_update().get(id=report_id) + + if old_pos and old_sev: + if old_pos == new_pos and old_sev == new_sev: + return None + + # Concurrent Hasura events can arrive after a newer reorder has + # already moved this finding. Treat those stale events as clean-up + # passes instead of replaying an old requested position. + if instance.position != new_pos or instance.severity_id != new_sev: + severity_ids = {old_sev, new_sev, instance.severity_id} + for severity_id in severity_ids: + normalize_finding_positions(model, report_id, severity_id) + return None - if old_pos and old_sev: - # Only run db queries if ``position`` or ``severity`` changed - if old_pos != new_pos or old_sev != new_sev: - # Get all findings in report that share the instance's severity rating - finding_queryset = model.objects.filter( - Q(report__pk=instance.report.pk) & Q(severity=instance.severity) - ).order_by("position") - - # If severity rating changed, adjust positioning in the previous severity group if old_sev != new_sev: - # Get a list of findings for the old severity rating - old_sev_queryset = model.objects.filter( - Q(report__pk=instance.report.pk) & Q(severity=old_sev) - ).order_by("position") - if old_sev_queryset: - for finding in old_sev_queryset: - # Adjust position to close gap created by moved finding - if finding.position > old_pos: - new_pos = finding.position - 1 - old_sev_queryset.filter(id=finding.id).order_by("position").update(position=new_pos) - - # The ``modelUpdateForm`` sets minimum number to 0, but check again for funny business - instance.position = max(instance.position, 1) - - # The ``position`` value should not be larger than total findings - if instance.position > finding_queryset.count(): - finding_queryset.filter(id=instance.id).update(position=finding_queryset.count()) - - counter = 1 - if finding_queryset: - # Loop from top position down and look for a match - for finding in finding_queryset: - # Check if finding in loop is the finding being updated - if not instance.pk == finding.pk: - # Increment position counter when counter equals new value - if instance.position == counter: - counter += 1 - finding_queryset.filter(id=finding.id).update(position=counter) - counter += 1 - else: - pass - # No other findings with the chosen severity, so set ``position`` to ``1`` + normalize_finding_positions(model, report_id, old_sev) + normalize_finding_positions( + model, + report_id, + new_sev, + moving_instance_id=instance.id, + target_position=new_pos, + ) else: - instance.position = 1 - # Place newly created findings at the end of the current list - else: - finding_queryset = model.objects.filter(Q(report__pk=instance.report.pk) & Q(severity=instance.severity)) - finding_queryset.filter(id=instance.id).update(position=finding_queryset.count()) + normalize_finding_positions( + model, + report_id, + new_sev, + moving_instance_id=instance.id, + target_position=new_pos, + ) + else: + # Insert events append new findings to the end of the severity group. + group_count = model.objects.filter( + Q(report_id=report_id) & Q(severity_id=new_sev) + ).count() + normalize_finding_positions( + model, + report_id, + new_sev, + moving_instance_id=instance.id, + target_position=group_count, + ) return None From cf77303bf6e5c1475849cda9e0aacf94293b3522 Mon Sep 17 00:00:00 2001 From: yg-ht Date: Thu, 16 Jul 2026 13:04:32 +0100 Subject: [PATCH 3/4] Remove redundant finding position count query --- ghostwriter/modules/model_utils.py | 11 ++++------- 1 file changed, 4 insertions(+), 7 deletions(-) diff --git a/ghostwriter/modules/model_utils.py b/ghostwriter/modules/model_utils.py index fc43c5125..6ece86f77 100644 --- a/ghostwriter/modules/model_utils.py +++ b/ghostwriter/modules/model_utils.py @@ -40,12 +40,10 @@ def to_dict(instance: django.db.models.Model, include_id: bool = False, resolve_ return data -def _clamp_position(position: Optional[int], count: int) -> int: +def _clamp_position(position: int, count: int) -> int: """Return a one-based position that fits inside a group of findings.""" if count < 1: return 1 - if position is None: - return count return min(max(position, 1), count) @@ -83,6 +81,8 @@ def normalize_finding_positions( None, ) if moving_finding is not None: + if target_position is None: + target_position = len(findings) ordered_findings = [ finding for finding in findings if finding.id != moving_instance_id ] @@ -161,14 +161,11 @@ def set_finding_positions( ) else: # Insert events append new findings to the end of the severity group. - group_count = model.objects.filter( - Q(report_id=report_id) & Q(severity_id=new_sev) - ).count() normalize_finding_positions( model, report_id, new_sev, moving_instance_id=instance.id, - target_position=group_count, + target_position=None, ) return None From 615c8152775d12edd3dc319b1591813705ccd14c Mon Sep 17 00:00:00 2001 From: yg-ht Date: Thu, 16 Jul 2026 15:37:55 +0100 Subject: [PATCH 4/4] Simplify collaboration save-error handling --- javascript/src/collab_server/base_handler.ts | 26 ++-- .../handlers/report_finding_link.ts | 25 +-- javascript/src/collab_server/index.ts | 94 ++--------- javascript/src/collab_server/save_error.ts | 17 ++ .../tests/e2e/collab_save_error.spec.ts | 32 ++++ .../e2e/report_finding_link_save.spec.ts | 146 ++++++++++++++++++ 6 files changed, 236 insertions(+), 104 deletions(-) create mode 100644 javascript/src/collab_server/save_error.ts create mode 100644 javascript/tests/e2e/collab_save_error.spec.ts create mode 100644 javascript/tests/e2e/report_finding_link_save.spec.ts diff --git a/javascript/src/collab_server/base_handler.ts b/javascript/src/collab_server/base_handler.ts index 8415be069..9c68274a8 100644 --- a/javascript/src/collab_server/base_handler.ts +++ b/javascript/src/collab_server/base_handler.ts @@ -6,11 +6,14 @@ import { } from "@apollo/client"; /** Functions for loading, saving, and converting a type from/to YJS. */ -export type ModelHandler = { +export type ModelHandler = { load: (client: ApolloClient, id: number) => Promise<[Y.Doc, T]>; - getSavePayload: (id: number, doc: Y.Doc, data: T) => P; - save: (client: ApolloClient, payload: P, data: T) => Promise; - markSaved?: (payload: P, data: T) => void; + save: ( + client: ApolloClient, + id: number, + doc: Y.Doc, + data: T + ) => Promise; }; export type IdVars = { id: number }; @@ -21,6 +24,7 @@ export type IdVars = { id: number }; * @param setQuery The GraphQL query to save the model. * @param fillFields Function to set fields on a `Y.Doc` based on the results returned from the `getQuery`. Called in a YJS transaction. * @param mkQueryVars Function to get the parameters for the `setQuery` to save the model. Called in a YJS transaction. + * @param onSaveSuccess Optional callback to update handler state after the generated variables have been saved successfully. * @returns The model handler. */ export function simpleModelHandler< @@ -32,8 +36,9 @@ export function simpleModelHandler< getQuery: TypedDocumentNode, setQuery: TypedDocumentNode, fillFields: (doc: Y.Doc, res: GetRes) => T, - mkQueryVars: (doc: Y.Doc, id: number, data: T) => SetQueryVars -): ModelHandler { + mkQueryVars: (doc: Y.Doc, id: number, data: T) => SetQueryVars, + onSaveSuccess?: (queryVars: SetQueryVars, data: T) => void +): ModelHandler { return { load: async (client, id) => { const res = await client.query({ @@ -52,21 +57,20 @@ export function simpleModelHandler< }); return [doc, data!]; }, - getSavePayload(id, doc, data) { + async save(client, id, doc, data) { let queryVars: SetQueryVars | undefined; doc.transact(() => { queryVars = mkQueryVars(doc, id, data); }); - return queryVars!; - }, - async save(client, payload) { + const savedQueryVars = queryVars!; const res = await client.mutate({ mutation: setQuery, - variables: payload, + variables: savedQueryVars, }); if (res.errors) { throw res.errors; } + onSaveSuccess?.(savedQueryVars, data); }, }; } diff --git a/javascript/src/collab_server/handlers/report_finding_link.ts b/javascript/src/collab_server/handlers/report_finding_link.ts index f677c652b..fa1eeeb99 100644 --- a/javascript/src/collab_server/handlers/report_finding_link.ts +++ b/javascript/src/collab_server/handlers/report_finding_link.ts @@ -128,6 +128,8 @@ const ReportFindingLinkHandler = simpleModelHandler( extraFields, }; + // Omit an unchanged severity so an unrelated collaborative save does + // not overwrite a newer value set through another application view. if (severityId !== data.savedSeverityId) { set.severityId = severityId; } @@ -137,15 +139,20 @@ const ReportFindingLinkHandler = simpleModelHandler( set, tags: yjsToTags(doc.get("tags", Y.Map)), }; + }, + (queryVars, data) => { + const variables = queryVars as { + set?: { severityId?: number | null }; + }; + + // Advance the baseline only for a severity value that the backend + // accepted. Failed saves must continue to include the local change. + if ( + variables.set && + Object.prototype.hasOwnProperty.call(variables.set, "severityId") + ) { + data.savedSeverityId = variables.set.severityId ?? null; + } } ); -ReportFindingLinkHandler.markSaved = (payload, data) => { - const variables = payload as { set?: { severityId?: number | null } }; - if ( - variables.set && - Object.prototype.hasOwnProperty.call(variables.set, "severityId") - ) { - data.savedSeverityId = variables.set.severityId ?? null; - } -}; export default ReportFindingLinkHandler; diff --git a/javascript/src/collab_server/index.ts b/javascript/src/collab_server/index.ts index 716294432..37317f8bd 100644 --- a/javascript/src/collab_server/index.ts +++ b/javascript/src/collab_server/index.ts @@ -21,9 +21,10 @@ import FindingHandler from "./handlers/finding"; import ReportFindingLinkHandler from "./handlers/report_finding_link"; import ReportHandler from "./handlers/report"; import ProjectHandler from "./handlers/project"; +import { setSaveError } from "./save_error"; // Extend this with your model handlers. See how-to-collab.md. -const HANDLERS_ARR: [string, ModelHandler][] = [ +const HANDLERS_ARR: [string, ModelHandler][] = [ ["observation", ObservationHandler], ["report_observation_link", ReportObservationLinkHandler], ["finding", FindingHandler], @@ -31,7 +32,7 @@ const HANDLERS_ARR: [string, ModelHandler][] = [ ["report", ReportHandler], ["project", ProjectHandler], ]; -const HANDLERS: Map> = new Map(HANDLERS_ARR); +const HANDLERS: Map> = new Map(HANDLERS_ARR); // Graphql Client @@ -87,45 +88,6 @@ class AuthError extends Error { const BASE_LOGGER = pino({}); const documentData = new Map(); -const lastSavedSignatures = new Map(); -const lastFailedSignatures = new Map(); - -function stableStringify(value: unknown): string { - return JSON.stringify(normalizeForSignature(value)); -} - -function normalizeForSignature(value: unknown): unknown { - if (Array.isArray(value)) { - return value.map(normalizeForSignature); - } - - if (value !== null && typeof value === "object") { - return Object.fromEntries( - Object.entries(value as Record) - .filter(([, entryValue]) => entryValue !== undefined) - .sort(([leftKey], [rightKey]) => - leftKey.localeCompare(rightKey) - ) - .map(([entryKey, entryValue]) => [ - entryKey, - normalizeForSignature(entryValue), - ]) - ); - } - - return value; -} - -function setSaveError(doc: Y.Doc, value: boolean) { - const serverInfo = doc.get("serverInfo", Y.Map); - if (serverInfo.get("saveError") === value) { - return; - } - - doc.transact(() => { - serverInfo.set("saveError", value); - }); -} const server = new Server({ port: 8000, @@ -267,15 +229,6 @@ const server = new Server({ serverInfo.set("saveError", false); }); documentData.set(data.documentName, docData); - const initialPayload = handler.getSavePayload( - context.id, - doc, - docData - ); - lastSavedSignatures.set( - data.documentName, - stableStringify(initialPayload) - ); return doc; } catch (e) { context.log.error({ msg: "Could not load document", err: e }); @@ -285,43 +238,18 @@ const server = new Server({ async onStoreDocument(data) { const context = data.context as Context; - const docData = documentData.get(data.documentName); - const handler = HANDLERS.get(context.model)!; - const payload = handler.getSavePayload( - context.id, - data.document, - docData - ); - const signature = stableStringify(payload); - - if (lastSavedSignatures.get(data.documentName) === signature) { - context.log.info("Skipping unchanged document save"); - setSaveError(data.document, false); - return; - } - - if (lastFailedSignatures.get(data.documentName) === signature) { - context.log.warn( - "Skipping retry for unchanged failed document save" - ); - setSaveError(data.document, true); - return; - } - try { + const docData = documentData.get(data.documentName); context.log.info("Saving document"); - await handler.save(gqlClient, payload, docData); - handler.markSaved?.(payload, docData); - lastSavedSignatures.set( - data.documentName, - stableStringify( - handler.getSavePayload(context.id, data.document, docData) - ) + const handler = HANDLERS.get(context.model)!; + await handler.save( + gqlClient, + context.id, + data.document, + docData ); - lastFailedSignatures.delete(data.documentName); } catch (e) { context.log.error({ msg: "Could not save document", err: e }); - lastFailedSignatures.set(data.documentName, signature); setSaveError(data.document, true); return; } @@ -334,8 +262,6 @@ const server = new Server({ async afterUnloadDocument(data) { documentData.delete(data.documentName); - lastSavedSignatures.delete(data.documentName); - lastFailedSignatures.delete(data.documentName); }, }); diff --git a/javascript/src/collab_server/save_error.ts b/javascript/src/collab_server/save_error.ts new file mode 100644 index 000000000..7d19ddbeb --- /dev/null +++ b/javascript/src/collab_server/save_error.ts @@ -0,0 +1,17 @@ +import * as Y from "yjs"; + +/** Update the shared save-error flag only when its value has changed. */ +export function setSaveError(doc: Y.Doc, value: boolean): void { + const serverInfo = doc.get("serverInfo", Y.Map); + + // Avoid generating another Yjs update when the error state is unchanged. + // Hocuspocus stores document updates, so redundant writes can create a + // feedback loop when this flag is set from its own storage callback. + if (serverInfo.get("saveError") === value) { + return; + } + + doc.transact(() => { + serverInfo.set("saveError", value); + }); +} diff --git a/javascript/tests/e2e/collab_save_error.spec.ts b/javascript/tests/e2e/collab_save_error.spec.ts new file mode 100644 index 000000000..8ec8b1451 --- /dev/null +++ b/javascript/tests/e2e/collab_save_error.spec.ts @@ -0,0 +1,32 @@ +import { expect, test } from "@playwright/test"; +import * as Y from "yjs"; + +import { setSaveError } from "../../src/collab_server/save_error"; + +test.describe("collaboration save-error state", () => { + test("emits Yjs updates only when the error state changes", () => { + const doc = new Y.Doc(); + const serverInfo = doc.get("serverInfo", Y.Map); + serverInfo.set("saveError", false); + + let updateCount = 0; + doc.on("update", () => { + updateCount += 1; + }); + + setSaveError(doc, false); + expect(updateCount).toBe(0); + + setSaveError(doc, true); + expect(updateCount).toBe(1); + + setSaveError(doc, true); + expect(updateCount).toBe(1); + + setSaveError(doc, false); + expect(updateCount).toBe(2); + + setSaveError(doc, false); + expect(updateCount).toBe(2); + }); +}); diff --git a/javascript/tests/e2e/report_finding_link_save.spec.ts b/javascript/tests/e2e/report_finding_link_save.spec.ts new file mode 100644 index 000000000..9f776d317 --- /dev/null +++ b/javascript/tests/e2e/report_finding_link_save.spec.ts @@ -0,0 +1,146 @@ +import { type ApolloClient } from "@apollo/client"; +import { expect, test } from "@playwright/test"; +import * as Y from "yjs"; + +import ReportFindingLinkHandler from "../../src/collab_server/handlers/report_finding_link"; + +type ReportFindingLinkData = { + extraFieldSpec: { internalName: string; type: string }[]; + savedSeverityId: number | null; +}; + +type SaveVariables = { + set: Record; +}; + +function createDocument(severityId: number): Y.Doc { + const doc = new Y.Doc(); + doc.get("plain_fields", Y.Map).set("severityId", severityId); + return doc; +} + +function createData(savedSeverityId: number): ReportFindingLinkData { + return { + extraFieldSpec: [], + savedSeverityId, + }; +} + +function createClient( + mutate: (variables: SaveVariables) => Promise +): ApolloClient { + return { + mutate: ({ variables }: { variables: SaveVariables }) => + mutate(variables), + } as unknown as ApolloClient; +} + +test.describe("report finding collaboration saves", () => { + test("omits an unchanged severity from unrelated saves", async () => { + const doc = createDocument(1); + const data = createData(1); + let databaseSeverityId = 2; + let savedVariables: SaveVariables | undefined; + const client = createClient(async (variables) => { + savedVariables = variables; + if ( + Object.prototype.hasOwnProperty.call( + variables.set, + "severityId" + ) + ) { + databaseSeverityId = variables.set.severityId as number; + } + return {}; + }); + + await ReportFindingLinkHandler.save(client, 7, doc, data); + + expect(savedVariables).toBeDefined(); + expect(savedVariables!.set).not.toHaveProperty("severityId"); + expect(databaseSeverityId).toBe(2); + expect(data.savedSeverityId).toBe(1); + }); + + test("records a changed severity only after a successful save", async () => { + const doc = createDocument(2); + const data = createData(1); + const saves: SaveVariables[] = []; + const client = createClient(async (variables) => { + saves.push(variables); + return {}; + }); + + await ReportFindingLinkHandler.save(client, 7, doc, data); + await ReportFindingLinkHandler.save(client, 7, doc, data); + + expect(saves[0].set).toHaveProperty("severityId", 2); + expect(saves[1].set).not.toHaveProperty("severityId"); + expect(data.savedSeverityId).toBe(2); + }); + + test("keeps a failed severity change dirty for retry", async () => { + const doc = createDocument(2); + const data = createData(1); + const failure = new Error("save failed"); + const failingClient = createClient(async () => ({ + errors: [failure], + })); + + await expect( + ReportFindingLinkHandler.save(failingClient, 7, doc, data) + ).rejects.toEqual([failure]); + expect(data.savedSeverityId).toBe(1); + + let retryVariables: SaveVariables | undefined; + const successfulClient = createClient(async (variables) => { + retryVariables = variables; + return {}; + }); + await ReportFindingLinkHandler.save(successfulClient, 7, doc, data); + + expect(retryVariables).toBeDefined(); + expect(retryVariables!.set).toHaveProperty("severityId", 2); + expect(data.savedSeverityId).toBe(2); + }); + + test("keeps a newer in-flight severity change dirty", async () => { + const doc = createDocument(2); + const data = createData(1); + let firstVariables: SaveVariables | undefined; + let resolveFirstSave: ((value: unknown) => void) | undefined; + const firstClient = createClient( + (variables) => + new Promise((resolve) => { + firstVariables = variables; + resolveFirstSave = resolve; + }) + ); + + const firstSave = ReportFindingLinkHandler.save( + firstClient, + 7, + doc, + data + ); + expect(firstVariables).toBeDefined(); + expect(firstVariables!.set).toHaveProperty("severityId", 2); + + doc.get("plain_fields", Y.Map).set("severityId", 3); + resolveFirstSave!({}); + await firstSave; + + expect(data.savedSeverityId).toBe(2); + + let secondVariables: SaveVariables | undefined; + const secondClient = createClient(async (variables) => { + secondVariables = variables; + return {}; + }); + await ReportFindingLinkHandler.save(secondClient, 7, doc, data); + + expect(secondVariables).toBeDefined(); + expect(secondVariables!.set).toHaveProperty("severityId", 3); + expect(data.savedSeverityId).toBe(3); + }); +});