diff --git a/ghostwriter/api/tests/test_views.py b/ghostwriter/api/tests/test_views.py index 79f37bf86..71277108a 100644 --- a/ghostwriter/api/tests/test_views.py +++ b/ghostwriter/api/tests/test_views.py @@ -3001,6 +3001,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( @@ -3014,112 +3049,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() @@ -3128,33 +3233,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() @@ -3163,36 +3253,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() @@ -3235,10 +3307,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 ed2b9efdb..b490d762a 100644 --- a/ghostwriter/api/views.py +++ b/ghostwriter/api/views.py @@ -55,7 +55,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 @@ -1859,16 +1863,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..6ece86f77 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: int, count: int) -> int: + """Return a one-based position that fits inside a group of findings.""" + if count < 1: + return 1 + 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: + if target_position is None: + target_position = len(findings) + 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,51 @@ 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. + normalize_finding_positions( + model, + report_id, + new_sev, + moving_instance_id=instance.id, + target_position=None, + ) return None diff --git a/javascript/src/collab_server/base_handler.ts b/javascript/src/collab_server/base_handler.ts index cafcffb84..9c68274a8 100644 --- a/javascript/src/collab_server/base_handler.ts +++ b/javascript/src/collab_server/base_handler.ts @@ -1,5 +1,9 @@ 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 = { @@ -20,13 +24,20 @@ 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( +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 + mkQueryVars: (doc: Y.Doc, id: number, data: T) => SetQueryVars, + onSaveSuccess?: (queryVars: SetQueryVars, data: T) => void ): ModelHandler { return { load: async (client, id) => { @@ -47,17 +58,19 @@ export function simpleModelHandler( return [doc, data!]; }, async save(client, id, doc, data) { - let queryVars; + let queryVars: SetQueryVars | undefined; doc.transact(() => { queryVars = mkQueryVars(doc, id, data); }); + const savedQueryVars = queryVars!; const res = await client.mutate({ mutation: setQuery, - variables: queryVars, + 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 d5f4ed4ea..fa1eeeb99 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,68 @@ 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, + }; + + // 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; + } + 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)), }; + }, + (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; + } } ); export default ReportFindingLinkHandler; diff --git a/javascript/src/collab_server/index.ts b/javascript/src/collab_server/index.ts index 54b0504cf..37317f8bd 100644 --- a/javascript/src/collab_server/index.ts +++ b/javascript/src/collab_server/index.ts @@ -21,6 +21,7 @@ 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][] = [ @@ -35,10 +36,11 @@ 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 }) => { @@ -240,17 +242,18 @@ const server = new Server({ 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, + context.id, + data.document, + docData + ); } 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); - }); + 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) { 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); + }); +});