Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
336 changes: 204 additions & 132 deletions ghostwriter/api/tests/test_views.py

Large diffs are not rendered by default.

21 changes: 10 additions & 11 deletions ghostwriter/api/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
152 changes: 104 additions & 48 deletions ghostwriter/modules/model_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -38,15 +40,71 @@ 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:
"""
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**

Expand All @@ -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
23 changes: 18 additions & 5 deletions javascript/src/collab_server/base_handler.ts
Original file line number Diff line number Diff line change
@@ -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<T> = {
Expand All @@ -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<GetRes, SetRes, SetQueryVars, T>(
export function simpleModelHandler<
GetRes,
SetRes,
SetQueryVars extends OperationVariables,
T,
>(
getQuery: TypedDocumentNode<GetRes, IdVars>,
setQuery: TypedDocumentNode<SetRes, SetQueryVars>,
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<T> {
return {
load: async (client, id) => {
Expand All @@ -47,17 +58,19 @@ export function simpleModelHandler<GetRes, SetRes, SetQueryVars, T>(
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);
},
};
}
92 changes: 61 additions & 31 deletions javascript/src/collab_server/handlers/report_finding_link.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down Expand Up @@ -86,43 +91,68 @@ const ReportFindingLinkHandler = simpleModelHandler(
);
tagsToYjs(res.tags.tags, doc.get("tags", Y.Map<boolean>));
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<any>);
const extraFields = extraFieldsFromYdoc(extraFieldSpec, doc);
const extraFields = extraFieldsFromYdoc(data.extraFieldSpec, doc);
const severityId = plainFields.get("severityId") ?? null;
const set: Record<string, unknown> = {
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<boolean>)),
};
},
(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;
Loading