Compare commits

..
Author SHA1 Message Date
stumpylog 74abdf8ff8 Docs: remove stale drf-writable-nested references, DRY up audit-actor update()
Final-review cleanup on refactor/remove-drf-writable-nested: several
comments/docstrings still described drf-writable-nested's removed
NestedUpdateMixin behavior in the present tense. Rewrote them to explain
the shared context-cache and delete-query-count guarantees in terms of
the current code (bulk-edit's per-field CustomFieldInstanceSerializer
construction, and _sync_custom_fields' single hard-delete query).

Also collapsed the duplicated `if custom_fields_data is not None:
self._sync_custom_fields(...)` line in DocumentSerializer.update()'s
audit-log branches into a single code path using
contextlib.nullcontext(), removing the drift risk that caused the
earlier audit-actor bug.

No behavior change.
2026-08-25 11:53:42 -07:00
stumpylogandClaude Sonnet 5 ce5d5c33b0 chore: remove drf-writable-nested dependency
Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-08-25 11:07:11 -07:00
stumpylog d9d44520bd test: add query-count regression coverage for custom_fields sync on update 2026-08-25 10:59:47 -07:00
stumpylog 500ff563cf refactor: replace drf-writable-nested's NestedUpdateMixin with explicit custom_fields sync 2026-08-25 10:38:51 -07:00
stumpylog 7f314b5149 test: characterize custom field instance delete-on-omit as a hard delete 2026-08-25 10:28:24 -07:00
stumpylog 132918821d Merge remote-tracking branch 'origin/perf/batch-custom-field-lookup' into tmp/perf-integration 2026-08-25 10:15:52 -07:00
stumpylog 512a3fe196 Merge remote-tracking branch 'origin/perf/batch-modify-custom-fields' into tmp/perf-integration 2026-08-25 10:15:48 -07:00
stumpylog b01e0368b7 Merge remote-tracking branch 'origin/perf/batch-tags-field-lookup' into tmp/perf-integration 2026-08-25 10:15:44 -07:00
stumpylog 1c82bd15c5 Merge branch 'perf/batch-set-permissions' into tmp/perf-integration 2026-08-25 10:15:39 -07:00
stumpylog cbac71c165 Perf: avoid unnecessary full-row fetches in batch permission assignment
set_permissions_for_objects now takes a model + pks instead of instances,
and identity filtering resolves straight to ids, so bulk-editing
permissions no longer materializes full Document/User/Group rows just to
read their pk/id. Row construction for bulk_create is also chunked to
bound peak memory for very large "apply to all" operations.
2026-08-25 09:27:24 -07:00
stumpylog 8780bcd5c7 Perf: batch guardian permission assignment in bulk-edit
bulk_edit.set_permissions and BulkEditObjectPermissionsView both
looped documents/objects and called set_permissions_for_object per
object, which itself calls guardian's assign_perm/remove_perm once
per (object, user) pair -- ~10-20+ queries per object, scaling with
selection size.

Added set_permissions_for_objects, a bulk equivalent that resolves
existing permission holders once across the whole batch (not once per
object) and applies changes with a small, batch-size-independent
number of queries per action instead of one per (object, user) pair.

Deliberately does not use guardian's queryset-aware assign_perm:
passing a list as the target routes to bulk_assign_perm, which skips
creating a direct permission row for anyone who already has the
permission via ANY group membership (checked via
ObjectPermissionChecker.has_perm, which is group-inheritance-aware) --
unlike the single-object assign_perm this replaces, which always
ensures a direct row via get_or_create. Losing that guarantee would
mean a later revocation of the group's grant silently strips access
an admin explicitly asked to be direct. Bulk-creates rows straight
against UserObjectPermission/GroupObjectPermission instead
(ignore_conflicts=True, relying on the existing (identity, permission,
object_pk) unique constraint), which preserves the original semantics
exactly while still batching every object and identity into one query
per action. Also raises Permission.DoesNotExist for an unrecognized
action name instead of silently no-op-ing, matching the original
per-object path -- BulkEditObjectsSerializer never actually validates
action keys against the raw client-supplied permissions dict, so this
is reachable from client input, not just internal callers.

Verified via CaptureQueriesContext: query count is now identical at 5
vs. 50 documents/objects (was 1,123 queries for 20 documents on the
Document path, 2,806 for 50 tags on the BulkEditObjectPermissionsView
path, both now flat). Full documents test suite green (2,148 passed,
1 skipped).
2026-08-24 20:04:58 -07:00
stumpylog 9d2416c435 Perf: batch id resolution for TagsField and friends
TagsField/CorrespondentField/DocumentTypeField/StoragePathField were
plain PrimaryKeyRelatedField subclasses with no batching. When used
with many=True (only tags today: DocumentSerializer.tags,
WorkflowActionSerializer.assign_tags), DRF's ManyRelatedField resolves
each submitted id with its own query -- one query per tag on every
PATCH/PUT that sets tags.

Added BatchResolvingPrimaryKeyRelatedField as the shared base for all
four field classes and overrode many_init so the many=True form
(_BatchingManyRelatedField) resolves the whole id list with one
pk__in query, falling back to the child relation's normal per-item
validation for anything not found in that batch. Only TagsField uses
many=True today, but the fix isn't tag-specific -- if a future PR puts
many=True on one of the others, it inherits the same batching instead
of reintroducing this as a new bug.

Independent review caught a real regression: Django's IntegerFieldOverflow
guard (out-of-range int -> EmptyResultSet) only covers exact/gt/gte/lt/lte
lookups, not `in`, so an absurdly large tag id reached the batched
pk__in= query as-is and raised an unhandled OverflowError (SQLite) /
DataError (Postgres) instead of the normal 400 the original per-item
`exact` lookup produced. Guarded the batch query and fall through to
per-item resolution (which goes through the protected `exact` lookup)
on failure.

Verified via CaptureQueriesContext against a real API PATCH: 20 tags
dropped from 54 to 35 queries per request (exactly the 19 saved by
collapsing 20 individual lookups into one batched query). Full
documents/workflows/bulk-edit/retagger/custom-fields suites green
(443 passed).
2026-08-24 15:15:00 -07:00
stumpylog a98d0669e4 Perf: batch CustomField/Document lookups in modify_custom_fields
modify_custom_fields looped documents x fields, re-.get()-ing the
CustomField queryset per iteration and Document.objects.get() per doc
for DOCUMENTLINK fields -- same shape as the earlier custom_fields
serializer N+1 (#13779), just nested one level deeper. Resolve both
into dicts once up front instead. Also pass the resolved objects
(not bare ids) to update_or_create so newly-created CustomFieldInstance
rows cache their field/document FK, avoiding a re-fetch when auditlog's
post_save receiver calls str(instance) (which touches .field.name).

docs_by_id defers `content` (the one field guaranteed both large and
unused by this function or its receivers) rather than using .only(),
since .only() would just turn the filename-generation signal's other
field access into a deferred-reload N+1.

Verified via CaptureQueriesContext: 6 docs x 4 fields dropped from 48
CustomField queries to 1; DOCUMENTLINK per-doc Document lookups dropped
from N to 0 (single batched query instead).
2026-08-24 14:21:59 -07:00
stumpylog bda506968b Handles a bad client sending malformed JSON or non-int primary keys 2026-08-24 12:35:06 -07:00
Trenton Holmes 4ccb34a70b Perf: avoid per-instance CustomField reload in DocumentMetadataOverrides
send_websocket_document_updated calls document.refresh_from_db()
before building overrides, which drops the custom_fields prefetch
(and its select_related("field")) set up by the view's queryset.
DocumentMetadataOverrides.from_document() then lazily reloads field
once per custom field instance. Since from_document() can't rely on
the caller having a prefetched document, select_related explicitly at
the point of use instead.
2026-08-23 17:33:05 -07:00
Trenton Holmes 0a466c9fcf Perf: reuse resolved CustomField objects across drf-writable-nested's per-item revalidation
drf-writable-nested's update_or_create_reverse_relations rebuilds a
fresh serializer -- and fresh field instances -- per custom_fields item
while matching existing vs. new instances during save(), so the
per-instance lookup cache alone only helped the first validation pass.
It passes the same context dict (by reference) to every one of those
serializers, so stash resolved CustomField objects there instead:
later passes reuse them for free rather than re-querying.
2026-08-23 17:12:25 -07:00
Trenton Holmes c9cc4f427d Perf: batch CustomField lookups when validating a document's custom_fields
DocumentSerializer.custom_fields validates each item's field id via a
plain PrimaryKeyRelatedField, which issues one SELECT per custom field
per validation pass (discussion #13690). Batch-resolve all field ids in
one query and cache them on the field instance so per-item validation
is free instead of re-querying.
2026-08-23 14:53:46 -07:00
29 changed files with 1210 additions and 1543 deletions
-29
View File
@@ -667,35 +667,6 @@ The action takes no options, its presence is what enables remote OCR for a match
If the remote engine is not configured, or does not support the document's file type, the document is
processed locally instead and a warning is written to the log.
##### Apply AI Suggestions {#workflow-action-apply-ai-suggestions}
"Apply AI Suggestions" actions ask the configured AI service for title and metadata suggestions,
the same as the AI suggestions shown on the document detail page, except applied automatically and in bulk.
It requires [AI features](configuration.md#ai) to be enabled. You can specify:
- Which suggestions to apply: title, tags, correspondent, document type, storage path and / or created
date. Suggestions for fields you did not select are discarded.
- Whether to create missing items. By default only tags, correspondents and document types that
already exist are assigned and any other suggestion is dropped. With this enabled, suggested items
that do not exist are created. Storage paths are never created.
- Whether to overwrite existing values. By default a field is only filled in if it is currently empty.
Note that documents almost always already have a title and created date, so if you select those you
will usually want to enable this too. Tags are an exception: suggested tags are always added and
never replace the document's existing tags.
The action works with every trigger **except Consumption Started**, because suggestions are made from
the document's text, which does not exist until after the document has been processed.
Because the query to the AI service is slow, the action is queued and runs in the background rather
than as part of the workflow run itself. The document is updated once the suggestions come back.
!!! warning
Every matching document results in a query to the AI service, which may incur costs and have privacy
implications. Queries can be slow, so a workflow matching a large number of documents can occupy the
task queue, and delay consumption of new documents, etc. Consider narrowing the trigger filters,
running in small batches and / or increasing workers.
#### Workflow placeholders
Titles and webhook payloads can be generated by workflows using [Jinja templates](https://jinja.palletsprojects.com/en/3.1.x/templates/).
-1
View File
@@ -40,7 +40,6 @@ dependencies = [
"djangorestframework~=3.16",
"drf-spectacular~=0.30",
"drf-spectacular-sidecar~=2026.7.1",
"drf-writable-nested~=0.7.1",
"filelock~=3.32.0",
"flower~=2.0.1",
"gotenberg-client~=0.14.0",
@@ -462,45 +462,6 @@
</div>
</div>
}
@case (WorkflowActionType.ApplyAiSuggestions) {
<div class="row">
<div class="col">
<p class="text-muted small" i18n>The document will be sent to the configured AI service for suggestions. Consider costs and privacy.</p>
<pngx-input-select
i18n-title
title="Apply suggestions for"
[items]="aiSuggestionFieldOptions"
[multiple]="true"
formControlName="ai_suggestion_fields"
[error]="error?.actions?.[i]?.ai_suggestion_fields"
hint="Suggestions for fields that are not selected are discarded."
i18n-hint
></pngx-input-select>
</div>
</div>
<div class="row">
<div class="col-md-6">
<pngx-input-switch
[horizontal]="true"
i18n-title
title="Create missing items"
formControlName="ai_create_missing"
hint="Create suggested tags, correspondents and document types that do not exist yet."
i18n-hint
></pngx-input-switch>
</div>
<div class="col-md-6">
<pngx-input-switch
[horizontal]="true"
i18n-title
title="Overwrite existing values"
formControlName="ai_overwrite_existing"
hint="Apply suggestions even if the document already has a value. Tags are always added, never replaced."
i18n-hint
></pngx-input-switch>
</div>
</div>
}
}
</div>
</ng-template>
@@ -22,7 +22,6 @@ import {
} from 'src/app/data/matching-model'
import { Workflow } from 'src/app/data/workflow'
import {
AISuggestionField,
WorkflowAction,
WorkflowActionType,
} from 'src/app/data/workflow-action'
@@ -50,7 +49,6 @@ import { TagsComponent } from '../../input/tags/tags.component'
import { TextComponent } from '../../input/text/text.component'
import { EditDialogMode } from '../edit-dialog.component'
import {
AI_SUGGESTION_FIELD_OPTIONS,
DOCUMENT_SOURCE_OPTIONS,
SCHEDULE_DATE_FIELD_OPTIONS,
TriggerFilterType,
@@ -241,15 +239,14 @@ describe('WorkflowEditDialogComponent', () => {
SCHEDULE_DATE_FIELD_OPTIONS
)
// Email, remote OCR and AI all disabled
// Email disabled
jest.spyOn(settingsService, 'get').mockReturnValue(false)
component.ngOnInit()
expect(component.actionTypeOptions).toEqual(
WORKFLOW_ACTION_OPTIONS.filter(
(a) =>
a.id !== WorkflowActionType.Email &&
a.id !== WorkflowActionType.RemoteOcr &&
a.id !== WorkflowActionType.ApplyAiSuggestions
a.id !== WorkflowActionType.RemoteOcr
)
)
})
@@ -347,125 +344,6 @@ describe('WorkflowEditDialogComponent', () => {
)
})
it('should offer apply AI suggestions unless every trigger is consumption', () => {
jest.spyOn(settingsService, 'get').mockReturnValue(true)
// Consumption runs before the document has been parsed, so there would be
// no content to make suggestions from
component.object = {
name: 'Workflow 1',
order: 0,
enabled: true,
triggers: [{ type: WorkflowTriggerType.Consumption }],
actions: [],
} as Workflow
component.ngOnInit()
expect(component.actionTypeOptions.map((a) => a.id)).not.toContain(
WorkflowActionType.ApplyAiSuggestions
)
// A second, usable trigger is enough
component.object = {
name: 'Workflow 2',
order: 0,
enabled: true,
triggers: [
{ type: WorkflowTriggerType.Consumption },
{ type: WorkflowTriggerType.DocumentAdded },
],
actions: [],
} as Workflow
component.ngOnInit()
expect(component.actionTypeOptions.map((a) => a.id)).toContain(
WorkflowActionType.ApplyAiSuggestions
)
})
it('should keep apply AI suggestions listed when an action already uses it', () => {
jest.spyOn(settingsService, 'get').mockReturnValue(true)
// Otherwise changing the trigger would silently blank the selection
component.object = {
name: 'Workflow 1',
order: 0,
enabled: true,
triggers: [{ type: WorkflowTriggerType.Consumption }],
actions: [{ type: WorkflowActionType.ApplyAiSuggestions }],
} as Workflow
component.ngOnInit()
expect(component.actionTypeOptions.map((a) => a.id)).toContain(
WorkflowActionType.ApplyAiSuggestions
)
})
it('should not offer apply AI suggestions when AI is disabled', () => {
jest
.spyOn(settingsService, 'get')
.mockImplementation((key) => key !== SETTINGS_KEYS.AI_ENABLED)
component.object = {
name: 'Workflow 1',
order: 0,
enabled: true,
triggers: [{ type: WorkflowTriggerType.DocumentAdded }],
actions: [],
} as Workflow
component.ngOnInit()
expect(component.actionTypeOptions.map((a) => a.id)).not.toContain(
WorkflowActionType.ApplyAiSuggestions
)
})
it('should create form fields for apply AI suggestions options', () => {
component.object = {
name: 'Workflow 1',
order: 0,
enabled: true,
triggers: [{ type: WorkflowTriggerType.DocumentAdded }],
actions: [
{
type: WorkflowActionType.ApplyAiSuggestions,
ai_suggestion_fields: [
AISuggestionField.Title,
AISuggestionField.Tags,
],
ai_create_missing: true,
ai_overwrite_existing: true,
},
],
} as Workflow
component.ngOnInit()
const action = component.actionFields.at(0)
expect(action.get('ai_suggestion_fields').value).toEqual([
AISuggestionField.Title,
AISuggestionField.Tags,
])
expect(action.get('ai_create_missing').value).toBeTruthy()
expect(action.get('ai_overwrite_existing').value).toBeTruthy()
expect(component.aiSuggestionFieldOptions).toEqual(
AI_SUGGESTION_FIELD_OPTIONS
)
})
it('should default apply AI suggestions options on a new action', () => {
component.object = {
name: 'Workflow 1',
order: 0,
enabled: true,
triggers: [{ type: WorkflowTriggerType.DocumentAdded }],
actions: [],
} as Workflow
component.addAction()
const action = component.actionFields.at(component.actionFields.length - 1)
expect(action.get('ai_suggestion_fields').value).toEqual([])
expect(action.get('ai_create_missing').value).toBeFalsy()
expect(action.get('ai_overwrite_existing').value).toBeFalsy()
})
it('should support add and remove triggers and actions', () => {
component.object = workflow
component.addTrigger()
@@ -30,7 +30,6 @@ import { StoragePath } from 'src/app/data/storage-path'
import { SETTINGS_KEYS } from 'src/app/data/ui-settings'
import { Workflow } from 'src/app/data/workflow'
import {
AISuggestionField,
WorkflowAction,
WorkflowActionType,
} from 'src/app/data/workflow-action'
@@ -153,37 +152,6 @@ export const WORKFLOW_ACTION_OPTIONS = [
id: WorkflowActionType.RemoteOcr,
name: $localize`Remote OCR`,
},
{
id: WorkflowActionType.ApplyAiSuggestions,
name: $localize`Apply AI suggestions`,
},
]
export const AI_SUGGESTION_FIELD_OPTIONS = [
{
id: AISuggestionField.Title,
name: $localize`Title`,
},
{
id: AISuggestionField.Tags,
name: $localize`Tags`,
},
{
id: AISuggestionField.Correspondent,
name: $localize`Correspondent`,
},
{
id: AISuggestionField.DocumentType,
name: $localize`Document type`,
},
{
id: AISuggestionField.StoragePath,
name: $localize`Storage path`,
},
{
id: AISuggestionField.Created,
name: $localize`Created date`,
},
]
export enum TriggerFilterType {
@@ -608,24 +576,6 @@ export class WorkflowEditDialogComponent
allowed = allowed.filter((a) => a.id !== WorkflowActionType.RemoteOcr)
}
// Only available after consumption. Unlike remote OCR this is hidden only
// once every trigger is consumption, so it stays offered on a workflow
// that has no triggers yet.
const aiSuggestionsUsable =
this.settingsService.get(SETTINGS_KEYS.AI_ENABLED) &&
(!formWorkflow?.triggers?.length ||
formWorkflow.triggers.some(
(trigger) => trigger.type !== WorkflowTriggerType.Consumption
) ||
formWorkflow.actions?.some(
(action) => action.type === WorkflowActionType.ApplyAiSuggestions
))
if (!aiSuggestionsUsable) {
allowed = allowed.filter(
(a) => a.id !== WorkflowActionType.ApplyAiSuggestions
)
}
if (
this.allowedActionTypes?.length === allowed.length &&
this.allowedActionTypes.every((a, i) => a.id === allowed[i].id)
@@ -1277,11 +1227,6 @@ export class WorkflowEditDialogComponent
passwords: new FormControl(
this.formatPasswords(action.passwords ?? [])
),
ai_suggestion_fields: new FormControl(
action.ai_suggestion_fields ?? []
),
ai_create_missing: new FormControl(!!action.ai_create_missing),
ai_overwrite_existing: new FormControl(!!action.ai_overwrite_existing),
}),
{ emitEvent }
)
@@ -1371,10 +1316,6 @@ export class WorkflowEditDialogComponent
return this.actionTypeOptions.find((t) => t.id === type)?.name ?? ''
}
get aiSuggestionFieldOptions() {
return AI_SUGGESTION_FIELD_OPTIONS
}
addAction() {
if (!this.object) {
this.object = Object.assign({}, this.objectForm.value)
@@ -1428,9 +1369,6 @@ export class WorkflowEditDialogComponent
include_document: false,
},
passwords: [],
ai_suggestion_fields: [],
ai_create_missing: false,
ai_overwrite_existing: false,
}
this.object.actions.push(action)
this.createActionField(action)
-17
View File
@@ -8,17 +8,6 @@ export enum WorkflowActionType {
PasswordRemoval = 5,
MoveToTrash = 6,
RemoteOcr = 7,
ApplyAiSuggestions = 8,
}
// see src/documents/models.py AISuggestionField
export enum AISuggestionField {
Title = 'title',
Tags = 'tags',
Correspondent = 'correspondent',
DocumentType = 'document_type',
StoragePath = 'storage_path',
Created = 'created',
}
export interface WorkflowActionEmail extends ObjectWithId {
@@ -113,10 +102,4 @@ export interface WorkflowAction extends ObjectWithId {
webhook?: WorkflowActionWebhook
passwords?: string[]
ai_suggestion_fields?: AISuggestionField[]
ai_create_missing?: boolean
ai_overwrite_existing?: boolean
}
+46 -29
View File
@@ -27,7 +27,7 @@ from documents.models import DocumentType
from documents.models import PaperlessTask
from documents.models import StoragePath
from documents.models import Tag
from documents.permissions import set_permissions_for_object
from documents.permissions import set_permissions_for_objects
from documents.plugins.helpers import DocumentsStatusManager
from documents.tasks import bulk_update_documents
from documents.tasks import consume_file
@@ -305,33 +305,49 @@ def modify_custom_fields(
else [(field, None) for field in add_custom_fields]
)
custom_fields = CustomField.objects.filter(
id__in=[int(field) for field, _ in add_custom_fields],
).distinct()
custom_fields_by_id: dict[int, CustomField] = {
cf.id: cf
for cf in CustomField.objects.filter(
id__in=[int(field) for field, _ in add_custom_fields],
)
}
# Deferred, not `.only()`: these objects get cached onto the FK
# descriptor of newly-created CustomFieldInstance rows below, and
# downstream post_save receivers (e.g. the filename-generation signal)
# touch other Document fields -- `.only("pk")` would just turn that into
# a deferred-field reload per document, trading one N+1 for another.
# `content` is the one field guaranteed to be both large (full OCR text)
# and unused by anything this function or its receivers touch.
docs_by_id: dict[int, Document] = {
doc.id: doc
for doc in Document.objects.filter(id__in=affected_docs).defer("content")
}
for field_id, value in add_custom_fields:
custom_field = custom_fields_by_id[field_id]
value_field = CustomFieldInstance.TYPE_TO_DATA_STORE_NAME_MAP[
custom_field.data_type
]
for doc_id in affected_docs:
defaults = {}
custom_field = custom_fields.get(id=field_id)
if custom_field:
value_field = CustomFieldInstance.TYPE_TO_DATA_STORE_NAME_MAP[
custom_field.data_type
]
defaults[value_field] = value
if (
custom_field.data_type == CustomField.FieldDataType.DOCUMENTLINK
and value
and doc_id in value
):
# Prevent self-linking
continue
defaults = {value_field: value}
if (
custom_field.data_type == CustomField.FieldDataType.DOCUMENTLINK
and value
and doc_id in value
):
# Prevent self-linking
continue
# Pass the already-resolved objects, not bare ids: this caches
# them on the FK descriptor of any newly-created instance, so a
# later `.field`/`.document` access (e.g. auditlog's post_save
# receiver calling `str(instance)`, which touches `.field.name`)
# doesn't trigger its own per-instance re-fetch.
CustomFieldInstance.objects.update_or_create(
document_id=doc_id,
field_id=field_id,
document=docs_by_id[doc_id],
field=custom_field,
defaults=defaults,
)
if custom_field.data_type == CustomField.FieldDataType.DOCUMENTLINK:
doc = Document.objects.get(id=doc_id)
reflect_doclinks(doc, custom_field, value)
reflect_doclinks(docs_by_id[doc_id], custom_field, value)
# For doc link fields that are being removed, remove symmetrical links
for doclink_being_removed_instance in CustomFieldInstance.objects.filter(
@@ -339,12 +355,10 @@ def modify_custom_fields(
field__id__in=remove_custom_fields,
field__data_type=CustomField.FieldDataType.DOCUMENTLINK,
value_document_ids__isnull=False,
):
).select_related("field"):
for target_doc_id in doclink_being_removed_instance.value:
remove_doclink(
document=Document.objects.get(
id=doclink_being_removed_instance.document.id,
),
document=docs_by_id[doclink_being_removed_instance.document_id],
field=doclink_being_removed_instance.field,
target_doc_id=target_doc_id,
)
@@ -430,10 +444,13 @@ def set_permissions(
else:
qs.update(owner=owner)
for doc in qs:
set_permissions_for_object(permissions=set_permissions, object=doc, merge=merge)
affected_docs = list(qs.values_list("pk", flat=True))
set_permissions_for_objects(
permissions=set_permissions,
model=Document,
pks=affected_docs,
merge=merge,
)
bulk_update_documents.apply_async(
kwargs={"document_ids": affected_docs},
+1 -1
View File
@@ -129,7 +129,7 @@ class DocumentMetadataOverrides:
)
overrides.custom_fields = {
custom_field.field.id: custom_field.value
for custom_field in doc.custom_fields.all()
for custom_field in doc.custom_fields.select_related("field").all()
}
groups_with_perms = get_groups_with_perms(
@@ -1,84 +0,0 @@
# Generated by Django 5.2.16 on 2026-08-10 18:26
from django.db import migrations
from django.db import models
class Migration(migrations.Migration):
dependencies = [
("documents", "0024_alter_workflowaction_type"),
]
operations = [
migrations.AddField(
model_name="workflowaction",
name="ai_create_missing",
field=models.BooleanField(
default=False,
help_text="Create suggested tags, correspondents, document types and storage paths that do not already exist instead of skipping them.",
verbose_name="create missing objects",
),
),
migrations.AddField(
model_name="workflowaction",
name="ai_overwrite_existing",
field=models.BooleanField(
default=False,
help_text="Apply suggestions even if the document already has a value for that field. Tags are always added to, never replaced.",
verbose_name="overwrite existing values",
),
),
migrations.AddField(
model_name="workflowaction",
name="ai_suggestion_fields",
field=models.JSONField(
blank=True,
help_text="Which of the AI-suggested fields to apply to the document.",
null=True,
verbose_name="AI suggestion fields",
),
),
migrations.AlterField(
model_name="workflowaction",
name="type",
field=models.PositiveSmallIntegerField(
choices=[
(1, "Assignment"),
(2, "Removal"),
(3, "Email"),
(4, "Webhook"),
(5, "Password removal"),
(6, "Move to trash"),
(7, "Remote OCR"),
(8, "Apply AI suggestions"),
],
default=1,
verbose_name="Workflow Action Type",
),
),
migrations.AlterField(
model_name="paperlesstask",
name="task_type",
field=models.CharField(
choices=[
("consume_file", "Consume File"),
("train_classifier", "Train Classifier"),
("sanity_check", "Sanity Check"),
("index_optimize", "Index Optimize"),
("mail_fetch", "Mail Fetch"),
("llm_index", "LLM Index"),
("empty_trash", "Empty Trash"),
("check_workflows", "Check Workflows"),
("bulk_update", "Bulk Update"),
("reprocess_document", "Reprocess Document"),
("build_share_link", "Build Share Link"),
("bulk_delete", "Bulk Delete"),
("apply_ai_suggestions", "Apply AI Suggestions"),
],
db_index=True,
help_text="The kind of work being performed",
max_length=50,
verbose_name="Task Type",
),
),
]
-40
View File
@@ -766,7 +766,6 @@ class PaperlessTask(ModelWithOwner):
REPROCESS_DOCUMENT = "reprocess_document", _("Reprocess Document")
BUILD_SHARE_LINK = "build_share_link", _("Build Share Link")
BULK_DELETE = "bulk_delete", _("Bulk Delete")
APPLY_AI_SUGGESTIONS = "apply_ai_suggestions", _("Apply AI Suggestions")
COMPLETE_STATUSES = (
Status.SUCCESS,
@@ -1675,18 +1674,6 @@ class WorkflowAction(models.Model):
7,
_("Remote OCR"),
)
APPLY_AI_SUGGESTIONS = (
8,
_("Apply AI suggestions"),
)
class AISuggestionField(models.TextChoices):
TITLE = ("title", _("Title"))
TAGS = ("tags", _("Tags"))
CORRESPONDENT = ("correspondent", _("Correspondent"))
DOCUMENT_TYPE = ("document_type", _("Document type"))
STORAGE_PATH = ("storage_path", _("Storage path"))
CREATED = ("created", _("Created date"))
type = models.PositiveSmallIntegerField(
_("Workflow Action Type"),
@@ -1925,33 +1912,6 @@ class WorkflowAction(models.Model):
),
)
ai_suggestion_fields = models.JSONField(
_("AI suggestion fields"),
null=True,
blank=True,
help_text=_(
"Which of the AI-suggested fields to apply to the document.",
),
)
ai_create_missing = models.BooleanField(
_("create missing objects"),
default=False,
help_text=_(
"Create suggested tags, correspondents, document types and storage "
"paths that do not already exist instead of skipping them.",
),
)
ai_overwrite_existing = models.BooleanField(
_("overwrite existing values"),
default=False,
help_text=_(
"Apply suggestions even if the document already has a value for that "
"field. Tags are always added to, never replaced.",
),
)
class Meta:
verbose_name = _("workflow action")
verbose_name_plural = _("workflow actions")
+176
View File
@@ -173,6 +173,182 @@ def set_permissions_for_object(
)
def _resolve_permissions(codenames: set[str], ctype: ContentType) -> list[Permission]:
"""
Resolves `codenames` to Permission rows, raising like the single-object
assign_perm() this bulk path replaces does (via a `.get()` internally)
if any codename doesn't exist -- e.g. a client-supplied action name that
was never validated (BulkEditObjectsSerializer._validate_permissions
calls validate_set_permissions() only for its side-effecting id checks
and discards the filtered dict it returns, so an unrecognized action key
reaches this function as-is). A plain `.filter()` with no existence
check would otherwise silently build zero rows and no-op instead of
reporting the bad input.
"""
permission_objs = list(
Permission.objects.filter(content_type=ctype, codename__in=codenames),
)
missing = codenames - {p.codename for p in permission_objs}
if missing:
raise Permission.DoesNotExist(
f"Permission matching query does not exist for codename(s): "
f"{', '.join(sorted(missing))}",
)
return permission_objs
# Target number of permission rows to build in Python before handing them to
# bulk_create -- keeps peak memory bounded for a large "apply to all" call,
# independent of bulk_create's own batch_size (which only caps the size of
# each INSERT statement, not how many row objects exist in memory at once).
_PERMISSION_ROW_CHUNK_SIZE = 5000
def _apply_bulk_permission_entry(
*,
perm_model: type[UserObjectPermission] | type[GroupObjectPermission],
identity_model: type[User] | type[Group],
identity_field: str,
ids: list[int],
codename: str,
permission_objs: list[Permission],
ctype: ContentType,
object_pks: list[str],
merge: bool,
) -> None:
# Only the ids are needed to build permission rows (via `<field>_id=`),
# so avoid fetching full User/Group rows for identities that may not
# even end up being granted anything new.
add_ids = set(
identity_model.objects.filter(id__in=ids).values_list("id", flat=True),
)
if not merge:
existing_ids = set(
perm_model.objects.filter(
content_type=ctype,
object_pk__in=object_pks,
permission__codename=codename,
).values_list(f"{identity_field}_id", flat=True),
)
remove_ids = existing_ids - add_ids
if remove_ids:
perm_model.objects.filter(
content_type=ctype,
object_pk__in=object_pks,
permission__codename=codename,
**{f"{identity_field}_id__in": remove_ids},
).delete()
if not add_ids:
return
rows_per_pk = len(permission_objs) * len(add_ids)
pks_per_chunk = max(1, _PERMISSION_ROW_CHUNK_SIZE // rows_per_pk)
for start in range(0, len(object_pks), pks_per_chunk):
pk_chunk = object_pks[start : start + pks_per_chunk]
rows = [
perm_model(
content_type=ctype,
object_pk=pk,
permission=permission_obj,
**{f"{identity_field}_id": identity_id},
)
for permission_obj in permission_objs
for pk in pk_chunk
for identity_id in add_ids
]
# ignore_conflicts skips only rows that already exist as an exact
# (identity, permission, object) match -- the same de-dup the
# underlying (user|group, permission, object_pk) unique constraint
# already enforces for the single-object assign_perm() this
# replaces, so it doesn't change what counts as "already granted".
# batch_size caps how many rows go into a single INSERT so a huge
# chunk doesn't build one enormous statement.
perm_model.objects.bulk_create(rows, ignore_conflicts=True, batch_size=1000)
def set_permissions_for_objects(
permissions: dict,
model: type[Model],
pks: QuerySet | list,
*,
merge: bool = False,
) -> None:
"""
Bulk equivalent of set_permissions_for_object: applies the same
permission changes to every object identified by `pks` at once.
Takes a model + pks (rather than model instances) deliberately -- the
permission rows built below only ever need `pk`, `content_type`, and
identity ids, so callers shouldn't have to fetch full rows (with every
other field) just to hand them to this function.
Deliberately does not use guardian's queryset/list-aware assign_perm:
passing a list as the object routes to bulk_assign_perm, which skips
creating a direct permission row for anyone who already has the
permission via ANY group membership (it checks
ObjectPermissionChecker.has_perm, which is group-inheritance-aware) --
unlike the single-object assign_perm this replaces, which always
ensures a direct row via get_or_create regardless of group-derived
access. Losing that guarantee would mean a later revocation of the
group's grant silently strips access an admin explicitly asked to be
direct. Bulk-creating rows straight against the permission models
instead (see _apply_bulk_permission_entry) preserves the original
always-create-a-direct-row semantics while still batching every object
and every identity into one query per action, rather than one query per
(object, user) pair.
"""
object_pks = [str(pk) for pk in pks]
if not object_pks:
return
model_name = model.__name__.lower()
ctype = ContentType.objects.get_for_model(model)
for action, entry in permissions.items():
codename = f"{action}_{model_name}"
implied_codenames = {codename}
if action == "change":
# change gives view too
implied_codenames.add(f"view_{model_name}")
# Resolved once per action (not once per users/groups branch) and
# shared between both below -- also where an unrecognized action
# name (see _resolve_permissions) is caught.
permission_objs = (
_resolve_permissions(implied_codenames, ctype)
if "users" in entry or "groups" in entry
else []
)
if "users" in entry:
_apply_bulk_permission_entry(
perm_model=UserObjectPermission,
identity_model=User,
identity_field="user",
ids=entry["users"],
codename=codename,
permission_objs=permission_objs,
ctype=ctype,
object_pks=object_pks,
merge=merge,
)
if "groups" in entry:
_apply_bulk_permission_entry(
perm_model=GroupObjectPermission,
identity_model=Group,
identity_field="group",
ids=entry["groups"],
codename=codename,
permission_objs=permission_objs,
ctype=ctype,
object_pks=object_pks,
merge=merge,
)
def permitted_object_ids(
user: User | None,
model: type[Model],
+236 -60
View File
@@ -1,8 +1,10 @@
from __future__ import annotations
import contextlib
import logging
import math
import re
from collections.abc import Iterable
from datetime import datetime
from datetime import timedelta
from decimal import Decimal
@@ -24,6 +26,7 @@ from django.core.validators import MaxValueValidator
from django.core.validators import MinValueValidator
from django.core.validators import RegexValidator
from django.core.validators import integer_validator
from django.db import DataError
from django.db.models import Count
from django.db.models import Q
from django.db.models.functions import Lower
@@ -37,12 +40,12 @@ from django.utils.timezone import make_aware
from django.utils.translation import gettext as _
from drf_spectacular.utils import extend_schema_field
from drf_spectacular.utils import extend_schema_serializer
from drf_writable_nested.serializers import NestedUpdateMixin
from guardian.core import ObjectPermissionChecker
from guardian.shortcuts import get_users_with_perms
from guardian.utils import get_group_obj_perms_model
from guardian.utils import get_user_obj_perms_model
from rest_framework import fields
from rest_framework import relations
from rest_framework import serializers
from rest_framework.exceptions import PermissionDenied
from rest_framework.fields import SerializerMethodField
@@ -742,22 +745,100 @@ class TagSerializer(MatchingModelSerializer, OwnedObjectSerializer):
return super().validate(attrs)
class CorrespondentField(serializers.PrimaryKeyRelatedField[Correspondent]):
class _BatchingManyRelatedField(serializers.ManyRelatedField):
"""
`ManyRelatedField.to_internal_value` resolves each id in the submitted
list with its own `child_relation.to_internal_value()` call -- one query
per item on every PATCH/PUT that sets a `many=True` relation field.
Batch-resolve them instead, falling back to the child relation's normal
(query-per-item) validation for anything that isn't a plausible int pk,
so bad input still gets the usual DRF validation error rather than being
silently dropped.
"""
@staticmethod
def _normalize_pk(item) -> int | None:
# Excludes bool: DRF's own PrimaryKeyRelatedField rejects it too
# (True == 1 would otherwise silently match pk 1).
if isinstance(item, bool):
return None
try:
return int(item)
except (TypeError, ValueError):
return None
def to_internal_value(self, data):
if isinstance(data, str) or not hasattr(data, "__iter__"):
self.fail("not_a_list", input_type=type(data).__name__)
if not self.allow_empty and len(data) == 0:
self.fail("empty")
item_pks = [(item, self._normalize_pk(item)) for item in data]
candidate_pks = {pk for _, pk in item_pks if pk is not None}
# Django's IntegerFieldOverflow guard (-> EmptyResultSet, i.e. no
# match) only covers exact/gt/gte/lt/lte lookups, not `in` -- an
# out-of-range int in `pk__in=` reaches the DB driver as-is and
# raises OverflowError (SQLite) / DataError (Postgres) instead of
# cleanly matching nothing. The per-item `exact`-lookup fallback
# below IS covered, so on that failure just skip the batch and let
# every item resolve individually -- each still costs one query,
# but reports the normal validation error instead of a raw 500.
try:
resolved_by_pk = {
obj.pk: obj
for obj in self.child_relation.get_queryset().filter(
pk__in=candidate_pks,
)
}
except (OverflowError, DataError):
resolved_by_pk = {}
result = []
for item, pk in item_pks:
obj = resolved_by_pk.get(pk) if pk is not None else None
result.append(
obj if obj is not None else self.child_relation.to_internal_value(item),
)
return result
class BatchResolvingPrimaryKeyRelatedField(serializers.PrimaryKeyRelatedField):
"""
A PrimaryKeyRelatedField whose `many=True` form (a DRF ManyRelatedField)
resolves all submitted ids with one batched query instead of one query
per id. Subclasses only need to implement `get_queryset()` as usual --
only `TagsField` is used with `many=True` today, but this is the base
for all four so the fix isn't tag-specific: if a future PR puts
`many=True` on correspondent/document_type/storage_path, it inherits the
same batching instead of reintroducing this as a new bug to rediscover.
"""
@classmethod
def many_init(cls, *args, **kwargs):
list_kwargs = {"child_relation": cls(*args, **kwargs)}
for key, value in kwargs.items():
if key in relations.MANY_RELATION_KWARGS:
list_kwargs[key] = value
return _BatchingManyRelatedField(**list_kwargs)
class CorrespondentField(BatchResolvingPrimaryKeyRelatedField[Correspondent]):
def get_queryset(self):
return Correspondent.objects.all()
class TagsField(serializers.PrimaryKeyRelatedField[Tag]):
class TagsField(BatchResolvingPrimaryKeyRelatedField[Tag]):
def get_queryset(self):
return Tag.objects.all()
class DocumentTypeField(serializers.PrimaryKeyRelatedField[DocumentType]):
class DocumentTypeField(BatchResolvingPrimaryKeyRelatedField[DocumentType]):
def get_queryset(self):
return DocumentType.objects.all()
class StoragePathField(serializers.PrimaryKeyRelatedField[StoragePath]):
class StoragePathField(BatchResolvingPrimaryKeyRelatedField[StoragePath]):
def get_queryset(self):
return StoragePath.objects.all()
@@ -876,8 +957,106 @@ def validate_documentlink_targets(user, doc_ids):
)
# A CustomField lookup cache scoped to a single field/serializer instance
# only helps within that one instance's own validation pass. Several call
# sites, though, build more than one CustomFieldInstanceSerializer (or its
# CustomFieldInstanceListSerializer/field) for the same request and pass
# each of them `context=self.context` -- the *same* dict object, not a
# copy -- e.g. bulk-edit's _validate_custom_field_values() constructing a
# fresh CustomFieldInstanceSerializer per submitted field. That context
# dict is already request-scoped (DRF builds it fresh per request via
# get_serializer_context()), so stashing the resolved CustomField objects
# there -- rather than in some new global/thread-local cache -- lets every
# one of those separately-instantiated serializers reuse them for free
# while staying entirely within DRF's existing, already-request-scoped
# machinery.
_CUSTOM_FIELD_CONTEXT_CACHE_KEY = "_custom_field_lookup_cache"
class _CachingCustomFieldPrimaryKeyField(serializers.PrimaryKeyRelatedField):
"""
Resolves CustomField ids with as few queries as possible: a per-instance
cache for repeat lookups on this exact field instance, backed by a
shared cache on the serializer context (see _CUSTOM_FIELD_CONTEXT_CACHE_KEY
above) so later, separately-instantiated fields for the same request
reuse what was already resolved instead of re-querying.
"""
def __init__(self, **kwargs: Any) -> None:
super().__init__(**kwargs)
self._cache: dict[int, CustomField] = {}
def _shared_cache(self) -> dict[int, CustomField]:
return self.context.setdefault(_CUSTOM_FIELD_CONTEXT_CACHE_KEY, {})
@staticmethod
def _normalize_pk(data: Any) -> int | None:
"""
Returns `data` coerced to the int a valid CustomField pk would be,
or None if `data` isn't a plausible pk (wrong type, unhashable,
non-numeric, or a bool -- DRF itself rejects bools as pks since
`True == 1` would otherwise silently match). None tells callers to
leave `data` alone and let `super().to_internal_value()` report the
normal validation error instead of touching the cache/queryset with
it directly.
"""
if isinstance(data, bool):
return None
try:
return int(data)
except (TypeError, ValueError):
return None
def prefetch(self, ids: Iterable[Any]) -> None:
shared_cache = self._shared_cache()
candidates = {pk for i in ids if (pk := self._normalize_pk(i)) is not None}
missing = {
i for i in candidates if i not in self._cache and i not in shared_cache
}
if missing:
for obj in self.get_queryset().filter(pk__in=missing):
shared_cache[obj.pk] = obj
for i in candidates:
obj = shared_cache.get(i)
if obj is not None:
self._cache[i] = obj
def to_internal_value(self, data: Any) -> CustomField:
pk = self._normalize_pk(data)
if pk is None:
return super().to_internal_value(data)
if pk in self._cache:
return self._cache[pk]
shared_cache = self._shared_cache()
if pk in shared_cache:
obj = shared_cache[pk]
self._cache[pk] = obj
return obj
obj: CustomField = super().to_internal_value(data)
self._cache[obj.pk] = obj
shared_cache[obj.pk] = obj
return obj
class CustomFieldInstanceListSerializer(serializers.ListSerializer):
def to_internal_value(self, data: Any) -> list[Any]:
if isinstance(data, list):
field_ids = []
for item in data:
if not isinstance(item, dict) or "field" not in item:
continue
try:
hash(item["field"])
except TypeError:
continue
field_ids.append(item["field"])
if field_ids:
self.child.fields["field"].prefetch(field_ids)
return super().to_internal_value(data)
class CustomFieldInstanceSerializer(serializers.ModelSerializer[CustomFieldInstance]):
field = serializers.PrimaryKeyRelatedField(queryset=CustomField.objects.all())
field = _CachingCustomFieldPrimaryKeyField(queryset=CustomField.objects.all())
value = ReadWriteSerializerMethodField(allow_null=True)
def create(self, validated_data):
@@ -978,6 +1157,7 @@ class CustomFieldInstanceSerializer(serializers.ModelSerializer[CustomFieldInsta
class Meta:
model = CustomFieldInstance
list_serializer_class = CustomFieldInstanceListSerializer
fields = [
"value",
"field",
@@ -1043,7 +1223,6 @@ class DocumentVersionInfoSerializer(serializers.Serializer[_DocumentVersionInfo]
)
class DocumentSerializer(
OwnedObjectSerializer,
NestedUpdateMixin,
DocumentUpdateFieldsModelSerializer,
):
correspondent = CorrespondentField(allow_null=True)
@@ -1258,16 +1437,60 @@ class DocumentSerializer(
if tag not in inbox_tags_not_being_added
]
if settings.AUDIT_LOG_ENABLED:
with set_actor(self.user):
super().update(instance, validated_data)
else:
custom_fields_data = validated_data.pop("custom_fields", None)
actor_context = (
set_actor(self.user)
if settings.AUDIT_LOG_ENABLED
else contextlib.nullcontext()
)
with actor_context:
super().update(instance, validated_data)
if custom_fields_data is not None:
self._sync_custom_fields(instance, custom_fields_data)
# hard delete custom field instances that were soft deleted
CustomFieldInstance.deleted_objects.filter(document=instance).delete()
return instance
def _sync_custom_fields(
self,
instance: Document,
custom_fields_data: list[dict],
) -> None:
"""
Create/update a CustomFieldInstance for every (field, value) pair in
custom_fields_data, then hard-delete any of the document's existing
instances whose field wasn't included.
Replaces drf-writable-nested's generic
update_or_create_reverse_relations()/delete_reverse_relations_if_need():
that machinery always matched submitted items by an instance "id"
this client payload never sends, so its own pk-matching never did
anything for this field -- the real upsert semantics were always
CustomFieldInstanceSerializer.create()'s update_or_create() below.
On a partial (PATCH) update, DRF skips the "value" field's required
check on the first validation pass because it's absent entirely from
the payload item, not merely null. drf-writable-nested happened to
re-enforce that check itself, by re-validating each item against a
freshly built, non-partial child serializer before saving. Reproduce
that specific guarantee explicitly here, since CustomFieldInstance's
"value" is not optional.
"""
for item in custom_fields_data:
if "value" not in item:
raise serializers.ValidationError(
{"custom_fields": [{"value": ["This field is required."]}]},
)
kept_field_ids: set[int] = set()
serializer = CustomFieldInstanceSerializer()
for item in custom_fields_data:
kept_field_ids.add(item["field"].pk)
serializer.create({**item, "document": instance})
CustomFieldInstance.objects.filter(document=instance).exclude(
field_id__in=kept_field_ids,
).hard_delete()
def __init__(self, *args, **kwargs) -> None:
self.truncate_content = kwargs.pop("truncate_content", False)
@@ -3235,9 +3458,6 @@ class WorkflowActionSerializer(serializers.ModelSerializer[WorkflowAction]):
"email",
"webhook",
"passwords",
"ai_suggestion_fields",
"ai_create_missing",
"ai_overwrite_existing",
]
def validate(self, attrs):
@@ -3295,23 +3515,6 @@ class WorkflowActionSerializer(serializers.ModelSerializer[WorkflowAction]):
"Passwords are required for password removal actions",
)
if (
"type" in attrs
and attrs["type"] == WorkflowAction.WorkflowActionType.APPLY_AI_SUGGESTIONS
):
fields = attrs.get("ai_suggestion_fields")
valid_fields = set(WorkflowAction.AISuggestionField.values)
if (
fields is None
or not isinstance(fields, list)
or len(fields) == 0
or any(field not in valid_fields for field in fields)
):
raise serializers.ValidationError(
"At least one valid field is required for apply AI "
f"suggestions actions, options are: {sorted(valid_fields)}",
)
return attrs
@@ -3340,43 +3543,24 @@ class WorkflowSerializer(serializers.ModelSerializer[Workflow]):
action.get("type") == WorkflowAction.WorkflowActionType.REMOTE_OCR
for action in attrs["actions"]
)
has_ai_suggestions_action = any(
action.get("type")
== WorkflowAction.WorkflowActionType.APPLY_AI_SUGGESTIONS
for action in attrs["actions"]
)
else:
has_remote_ocr_action = self.instance is not None and (
self.instance.actions.filter(
type=WorkflowAction.WorkflowActionType.REMOTE_OCR,
).exists()
)
has_ai_suggestions_action = self.instance is not None and (
self.instance.actions.filter(
type=WorkflowAction.WorkflowActionType.APPLY_AI_SUGGESTIONS,
).exists()
)
if "triggers" in attrs:
has_consumption_trigger = any(
trigger.get("type") == WorkflowTrigger.WorkflowTriggerType.CONSUMPTION
for trigger in attrs["triggers"]
)
has_non_consumption_trigger = any(
trigger.get("type") != WorkflowTrigger.WorkflowTriggerType.CONSUMPTION
for trigger in attrs["triggers"]
)
else:
has_consumption_trigger = self.instance is not None and (
self.instance.triggers.filter(
type=WorkflowTrigger.WorkflowTriggerType.CONSUMPTION,
).exists()
)
has_non_consumption_trigger = self.instance is not None and (
self.instance.triggers.exclude(
type=WorkflowTrigger.WorkflowTriggerType.CONSUMPTION,
).exists()
)
# Remote OCR can only work with consumption triggers
if has_remote_ocr_action and not has_consumption_trigger:
@@ -3384,14 +3568,6 @@ class WorkflowSerializer(serializers.ModelSerializer[Workflow]):
"Remote OCR actions require a consumption started trigger",
)
# Suggestions are made from the document content, which does not exist
# until after consumption has finished
if has_ai_suggestions_action and not has_non_consumption_trigger:
raise serializers.ValidationError(
"Apply AI suggestions actions require a trigger other than "
"consumption started",
)
return attrs
def update_triggers_and_actions(
-29
View File
@@ -984,28 +984,6 @@ def run_workflows(
"triggers, ignoring",
extra={"group": logging_group},
)
elif (
action.type
== WorkflowAction.WorkflowActionType.APPLY_AI_SUGGESTIONS
):
if use_overrides:
# The document has not been parsed yet, so there is no
# content for the LLM to make suggestions from
logger.debug(
"Apply AI suggestions action does not apply to "
"consumption triggers, ignoring",
extra={"group": logging_group},
)
else:
# Queued rather than run sync
from documents.tasks import apply_ai_suggestions
# kwargs so the PaperlessTask record can note the
# document, see _extract_input_data
apply_ai_suggestions.delay(
action_id=action.pk,
document_id=document.pk,
)
if not use_overrides:
# limit title to 128 characters
@@ -1061,7 +1039,6 @@ TRACKED_TASKS: dict[str, PaperlessTask.TaskType] = {
"documents.tasks.update_document_content_maybe_archive_file": PaperlessTask.TaskType.REPROCESS_DOCUMENT,
"documents.tasks.build_share_link_bundle": PaperlessTask.TaskType.BUILD_SHARE_LINK,
"documents.bulk_edit.delete": PaperlessTask.TaskType.BULK_DELETE,
"documents.tasks.apply_ai_suggestions": PaperlessTask.TaskType.APPLY_AI_SUGGESTIONS,
}
_CELERY_STATE_TO_STATUS: dict[str, PaperlessTask.Status] = {
@@ -1115,12 +1092,6 @@ def _extract_input_data(
return {"account_ids": account_ids}
return {}
if task_type == PaperlessTask.TaskType.APPLY_AI_SUGGESTIONS:
document_id = task_kwargs.get("document_id")
if document_id is not None:
return {"document_id": document_id}
return {}
return {}
-39
View File
@@ -714,45 +714,6 @@ def llmindex_index(
)
@shared_task(
bind=True,
autoretry_for=(Exception,),
max_retries=3,
retry_backoff=60,
retry_backoff_max=600,
retry_jitter=True,
)
def apply_ai_suggestions(self, action_id: int, document_id: int) -> None:
"""
Deferred "apply AI suggestions" workflow action.
"""
from documents.models import WorkflowAction
from documents.workflows.ai import apply_ai_suggestions_to_document
try:
action = WorkflowAction.objects.get(pk=action_id)
document = Document.objects.select_related("owner").get(pk=document_id)
except (WorkflowAction.DoesNotExist, Document.DoesNotExist):
logger.warning(
"Workflow action %s or document %s no longer exists, "
"not applying AI suggestions",
action_id,
document_id,
)
return
if not apply_ai_suggestions_to_document(action, document):
return
# No document_updated signal to avoid loop
clear_document_caches(document.pk)
index_document.delay(document.pk)
ai_config = AIConfig()
if ai_config.llm_index_enabled:
update_document_in_llm_index.apply_async(kwargs={"document": document})
@shared_task
def update_document_in_llm_index(document) -> None:
llm_index_add_or_update_document(document)
@@ -5,7 +5,9 @@ from unittest.mock import ANY
from django.contrib.auth.models import Permission
from django.contrib.auth.models import User
from django.db import connection
from django.test import override_settings
from django.test.utils import CaptureQueriesContext
from guardian.shortcuts import assign_perm
from rest_framework import status
from rest_framework.test import APITestCase
@@ -13,6 +15,9 @@ from rest_framework.test import APITestCase
from documents.models import CustomField
from documents.models import CustomFieldInstance
from documents.models import Document
from documents.serialisers import CustomFieldInstanceSerializer
from documents.serialisers import DocumentSerializer
from documents.tests.factories import DocumentFactory
from documents.tests.utils import DirectoriesMixin
@@ -530,6 +535,137 @@ class TestCustomFieldsAPI(DirectoriesMixin, APITestCase):
doc.refresh_from_db()
self.assertEqual(len(doc.custom_fields.all()), 10)
def test_document_serializer_custom_fields_validation_batches_field_lookup(
self,
) -> None:
"""
GIVEN:
- A document is being validated with several custom field values
at once (as happens on every PATCH/PUT/POST)
WHEN:
- The serializer is validated
THEN:
- The referenced CustomField objects are resolved with a single
query, not one query per custom field
"""
doc = DocumentFactory(mime_type="application/pdf")
custom_fields = [
CustomField.objects.create(
name=f"Test Custom Field {i}",
data_type=CustomField.FieldDataType.STRING,
)
for i in range(5)
]
serializer = DocumentSerializer(
doc,
data={
"custom_fields": [
{"field": custom_field.id, "value": "test value"}
for custom_field in custom_fields
],
},
partial=True,
)
with CaptureQueriesContext(connection) as ctx:
self.assertTrue(serializer.is_valid(), serializer.errors)
custom_field_lookups = [
query
for query in ctx.captured_queries
if 'FROM "documents_customfield" WHERE "documents_customfield"."id"'
in query["sql"]
]
self.assertEqual(
len(custom_field_lookups),
1,
"Expected a single batched query to resolve the custom fields, "
f"got {len(custom_field_lookups)}: {custom_field_lookups}",
)
def test_custom_field_lookup_reuses_shared_context_cache(self) -> None:
"""
GIVEN:
- A CustomField has already been resolved once, by a serializer
sharing a given `context` dict
WHEN:
- A second, separately-instantiated CustomFieldInstanceSerializer
validates the same field id, sharing that same context
(this mirrors what real call sites do, e.g. bulk-edit's
validate_custom_fields()/_validate_custom_field_values()
constructing a fresh CustomFieldInstanceSerializer per
submitted field, all sharing the outer serializer's context)
THEN:
- No additional query is issued to resolve the CustomField
"""
custom_field = CustomField.objects.create(
name="Test Custom Field",
data_type=CustomField.FieldDataType.STRING,
)
context: dict = {}
first_pass = CustomFieldInstanceSerializer(
data={"field": custom_field.id, "value": "a"},
context=context,
)
self.assertTrue(first_pass.is_valid(), first_pass.errors)
second_pass = CustomFieldInstanceSerializer(
data={"field": custom_field.id, "value": "b"},
context=context,
)
with CaptureQueriesContext(connection) as ctx:
self.assertTrue(second_pass.is_valid(), second_pass.errors)
custom_field_lookups = [
query
for query in ctx.captured_queries
if 'FROM "documents_customfield" WHERE "documents_customfield"."id"'
in query["sql"]
]
self.assertEqual(
len(custom_field_lookups),
0,
"Expected the second, separately-instantiated serializer to reuse "
f"the already-resolved CustomField, got: {custom_field_lookups}",
)
def test_custom_field_validation_rejects_malformed_field_value(self) -> None:
"""
GIVEN:
- A document is being validated with a malformed custom_fields
entry whose "field" value is neither a valid CustomField id
nor a type DRF's own PrimaryKeyRelatedField can safely reject
on its own (unhashable, or a non-numeric scalar)
WHEN:
- The serializer is validated
THEN:
- A normal validation error is raised, not an unhandled
TypeError/ValueError escaping past DRF's validation layer
"""
doc = DocumentFactory(mime_type="application/pdf")
bad_field_values = {
"unhashable-list": [],
"unhashable-dict": {},
"non-numeric-scalar": "abc",
}
for case_id, bad_field_value in bad_field_values.items():
with self.subTest(case_id):
serializer = DocumentSerializer(
doc,
data={
"custom_fields": [
{"field": bad_field_value, "value": "test value"},
],
},
partial=True,
)
self.assertFalse(serializer.is_valid())
self.assertIn("custom_fields", serializer.errors)
def test_change_custom_field_instance_value(self) -> None:
"""
GIVEN:
@@ -660,6 +796,72 @@ class TestCustomFieldsAPI(DirectoriesMixin, APITestCase):
assert _cf_4 is not None
self.assertEqual(_cf_4.value, date_value)
def test_delete_custom_field_instance_is_hard_deleted_not_soft_deleted(
self,
) -> None:
"""
GIVEN:
- A document has two custom field instances
WHEN:
- A PATCH request updates custom_fields to omit one of them
THEN:
- The omitted instance is immediately hard-deleted: it is gone
from both the default manager and the soft-deleted manager,
not left in a soft-deleted, not-yet-purged state
"""
doc = Document.objects.create(
title="WOW",
content="the content",
checksum="123-hard-delete",
mime_type="application/pdf",
)
kept_field = CustomField.objects.create(
name="Kept Field",
data_type=CustomField.FieldDataType.STRING,
)
removed_field = CustomField.objects.create(
name="Removed Field",
data_type=CustomField.FieldDataType.STRING,
)
resp = self.client.patch(
f"/api/documents/{doc.id}/",
data={
"custom_fields": [
{"field": kept_field.id, "value": "keep me"},
{"field": removed_field.id, "value": "remove me"},
],
},
format="json",
)
self.assertEqual(resp.status_code, status.HTTP_200_OK)
self.assertEqual(CustomFieldInstance.objects.count(), 2)
resp = self.client.patch(
f"/api/documents/{doc.id}/",
data={
"custom_fields": [
{"field": kept_field.id, "value": "keep me"},
],
},
format="json",
)
self.assertEqual(resp.status_code, status.HTTP_200_OK)
self.assertEqual(CustomFieldInstance.objects.count(), 1)
self.assertEqual(
CustomFieldInstance.deleted_objects.filter(field=removed_field).count(),
0,
"Removed custom field instance should be hard-deleted, not left "
"soft-deleted",
)
self.assertEqual(
CustomFieldInstance.global_objects.filter(field=removed_field).count(),
0,
"Removed custom field instance should not exist at all, even in "
"the all-rows manager",
)
def test_custom_field_validation(self) -> None:
"""
GIVEN:
@@ -1351,6 +1553,81 @@ class TestCustomFieldsAPI(DirectoriesMixin, APITestCase):
results = response.data["results"]
self.assertEqual(results[0]["document_count"], 0)
def test_document_update_custom_fields_sync_query_count(self) -> None:
"""
GIVEN:
- A document already has 3 custom field instances attached
WHEN:
- A PATCH request updates 2 of them, adds 1 new one, and omits
the 3rd (which should be deleted)
THEN:
- The omitted custom field is removed via exactly one
delete-shaped query, not drf-writable-nested's old two-step
soft-delete-then-hard-delete-after pattern (two delete
passes)
"""
doc = Document.objects.create(
title="WOW",
content="the content",
checksum="123-sync-count",
mime_type="application/pdf",
)
fields = [
CustomField.objects.create(
name=f"Sync Field {i}",
data_type=CustomField.FieldDataType.STRING,
)
for i in range(4)
]
# Attach the first 3 up front; the 4th is added in the PATCH below,
# and the 3rd is omitted (so it should be deleted).
for field in fields[:3]:
CustomFieldInstance.objects.create(
document=doc,
field=field,
value_text="initial",
)
self.assertEqual(CustomFieldInstance.objects.count(), 3)
with CaptureQueriesContext(connection) as ctx:
resp = self.client.patch(
f"/api/documents/{doc.id}/",
data={
"custom_fields": [
{"field": fields[0].id, "value": "updated 0"},
{"field": fields[1].id, "value": "updated 1"},
{"field": fields[3].id, "value": "new 3"},
],
},
format="json",
)
self.assertEqual(resp.status_code, status.HTTP_200_OK)
delete_queries = [
q
for q in ctx.captured_queries
if 'DELETE FROM "documents_customfieldinstance"' in q["sql"]
or (
'UPDATE "documents_customfieldinstance"' in q["sql"]
and '"deleted_at"' in q["sql"]
and '"deleted_at" = NULL' not in q["sql"]
)
]
self.assertEqual(
len(delete_queries),
1,
"Expected exactly one delete/soft-delete-marking query for the "
f"omitted custom field, got {len(delete_queries)}: {delete_queries}",
)
self.assertEqual(CustomFieldInstance.objects.count(), 3)
doc.refresh_from_db()
values = {cfi.field_id: cfi.value for cfi in doc.custom_fields.all()}
self.assertEqual(values[fields[0].id], "updated 0")
self.assertEqual(values[fields[1].id], "updated 1")
self.assertEqual(values[fields[3].id], "new 3")
self.assertNotIn(fields[2].id, values)
def test_patch_document_invalid_date_custom_field_returns_validation_error(
self,
) -> None:
+79
View File
@@ -1,5 +1,6 @@
import datetime
import json
import re
import shutil
import tempfile
import uuid
@@ -21,7 +22,9 @@ from django.core import mail
from django.core.cache import cache
from django.core.files.uploadedfile import SimpleUploadedFile
from django.db import DataError
from django.db import connection
from django.test import override_settings
from django.test.utils import CaptureQueriesContext
from django.utils import timezone
from guardian.shortcuts import assign_perm
from rest_framework import status
@@ -252,6 +255,82 @@ class TestDocumentApi(DirectoriesMixin, ConsumeTaskMixin, APITestCase):
doc.refresh_from_db()
self.assertEqual(doc.created, date(2023, 6, 28))
def test_document_update_tags_batches_tag_lookup(self) -> None:
"""
GIVEN:
- A document is being updated with several tags at once
WHEN:
- API PATCH request is made setting the document's tags
THEN:
- The referenced Tag objects are resolved with a single batched
query, not one query per tag
"""
doc = Document.objects.create(
title="none",
checksum="123",
mime_type="application/pdf",
)
tags = [TagFactory() for _ in range(8)]
with CaptureQueriesContext(connection) as ctx:
response = self.client.patch(
f"/api/documents/{doc.pk}/",
{"tags": [t.id for t in tags]},
format="json",
)
self.assertEqual(response.status_code, status.HTTP_200_OK)
# Match `"documents_tag"."id" = <literal>` (a single-row WHERE lookup)
# but not the same substring appearing as a JOIN's ON condition
# (`"documents_tag"."id" = "documents_document_tags"."tag_id"`),
# which is a legitimate, unrelated response-serialization query.
single_tag_lookup_re = re.compile(r'"documents_tag"\."id" = \d')
single_tag_lookups = [
q for q in ctx.captured_queries if single_tag_lookup_re.search(q["sql"])
]
self.assertEqual(
len(single_tag_lookups),
0,
"Expected tags to be resolved with a batched query, not "
f"per-tag lookups, got: {single_tag_lookups}",
)
doc.refresh_from_db()
self.assertCountEqual(
doc.tags.values_list("id", flat=True),
[t.id for t in tags],
)
def test_document_update_tags_rejects_out_of_range_id(self) -> None:
"""
GIVEN:
- A document is being updated with a tag id too large for the
database's integer column
WHEN:
- API PATCH request is made setting the document's tags
THEN:
- A normal 400 validation error is returned, not an unhandled
OverflowError/DataError escaping as a 500
Django's IntegerFieldOverflow guard converts an out-of-range int
into a clean "no match" for exact/gt/gte/lt/lte lookups, but not for
`in` -- the batched tag resolution uses `pk__in=`, so this has to be
guarded explicitly rather than relying on Django to do it.
"""
doc = Document.objects.create(
title="none",
checksum="123",
mime_type="application/pdf",
)
response = self.client.patch(
f"/api/documents/{doc.pk}/",
{"tags": [99999999999999999999999999999]},
format="json",
)
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
def test_document_update_legacy_created_format(self) -> None:
"""
GIVEN:
+58
View File
@@ -2,10 +2,15 @@ import datetime
import json
from unittest import mock
from django.contrib.auth.models import Group
from django.contrib.auth.models import Permission
from django.contrib.auth.models import User
from django.db import connection
from django.test import override_settings
from django.test.utils import CaptureQueriesContext
from guardian.shortcuts import assign_perm
from guardian.shortcuts import get_groups_with_perms
from guardian.shortcuts import get_users_with_perms
from rest_framework import status
from rest_framework.test import APITestCase
@@ -815,6 +820,59 @@ class TestBulkEditObjects(APITestCase):
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(StoragePath.objects.count(), 0)
def test_bulk_objects_set_permissions_query_count_independent_of_object_count(
self,
) -> None:
"""
GIVEN:
- Many tags are being bulk-edited to set permissions at once
WHEN:
- bulk_edit_objects API endpoint is called with set_permissions
operation over a small batch vs. a much larger one
THEN:
- The number of queries issued is the same either way -- each
user/group is applied across all tags with one batched call,
not one call per (tag, user) pair
"""
group1 = Group.objects.create(name="perm-group")
permissions = {
"view": {"users": [self.user1.id, self.user2.id], "groups": [group1.id]},
"change": {"users": [self.user1.id], "groups": [group1.id]},
}
def run_with_n_tags(n: int) -> int:
tags = [Tag.objects.create(name=f"perm-tag-{n}-{i}") for i in range(n)]
with CaptureQueriesContext(connection) as ctx:
response = self.client.post(
"/api/bulk_edit_objects/",
json.dumps(
{
"objects": [t.id for t in tags],
"object_type": "tags",
"operation": "set_permissions",
"permissions": permissions,
"merge": False,
},
),
content_type="application/json",
)
self.assertEqual(response.status_code, status.HTTP_200_OK)
for tag in tags:
self.assertEqual(get_users_with_perms(tag).count(), 2)
self.assertEqual(get_groups_with_perms(tag).count(), 1)
return len(ctx.captured_queries)
small_batch_queries = run_with_n_tags(5)
large_batch_queries = run_with_n_tags(50)
self.assertEqual(
small_batch_queries,
large_batch_queries,
"Expected the same query count regardless of tag count, got "
f"{small_batch_queries} queries for 5 tags vs. "
f"{large_batch_queries} for 50",
)
def test_bulk_objects_delete_all_filtered(self) -> None:
"""
GIVEN:
-139
View File
@@ -641,145 +641,6 @@ class TestApiWorkflows(DirectoriesMixin, APITestCase):
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertEqual(self.workflow.triggers.get(), self.trigger)
def _post_ai_suggestions_workflow(self, *, trigger_types, action: dict):
def trigger(trigger_type):
# consumption triggers require a filter of their own
if trigger_type == WorkflowTrigger.WorkflowTriggerType.CONSUMPTION:
return {"type": trigger_type, "filter_filename": "*.pdf"}
return {"type": trigger_type}
return self.client.post(
self.ENDPOINT,
json.dumps(
{
"name": "Apply AI suggestions",
"order": 1,
"triggers": [trigger(t) for t in trigger_types],
"actions": [
{
"type": WorkflowAction.WorkflowActionType.APPLY_AI_SUGGESTIONS,
**action,
},
],
},
),
content_type="application/json",
)
def test_api_create_apply_ai_suggestions_action(self) -> None:
"""
GIVEN:
- API request to create a workflow with an apply AI suggestions
action and a valid set of fields
WHEN:
- API is called
THEN:
- The workflow is created with the chosen options
"""
response = self._post_ai_suggestions_workflow(
trigger_types=[WorkflowTrigger.WorkflowTriggerType.DOCUMENT_ADDED],
action={
"ai_suggestion_fields": ["title", "tags", "correspondent"],
"ai_create_missing": True,
"ai_overwrite_existing": True,
},
)
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
action = Workflow.objects.get(name="Apply AI suggestions").actions.first()
self.assertEqual(
action.ai_suggestion_fields,
["title", "tags", "correspondent"],
)
self.assertTrue(action.ai_create_missing)
self.assertTrue(action.ai_overwrite_existing)
def test_api_create_apply_ai_suggestions_action_requires_fields(self) -> None:
"""
GIVEN:
- API request to create an apply AI suggestions action with no
fields selected, which could never do anything
WHEN:
- API is called
THEN:
- Correct HTTP 400 response
- No objects are created
"""
existing_count = Workflow.objects.count()
response = self._post_ai_suggestions_workflow(
trigger_types=[WorkflowTrigger.WorkflowTriggerType.DOCUMENT_ADDED],
action={"ai_suggestion_fields": []},
)
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertEqual(Workflow.objects.count(), existing_count)
def test_api_create_apply_ai_suggestions_action_rejects_unknown_field(
self,
) -> None:
"""
GIVEN:
- API request to create an apply AI suggestions action naming a
field that does not exist
WHEN:
- API is called
THEN:
- Correct HTTP 400 response
"""
response = self._post_ai_suggestions_workflow(
trigger_types=[WorkflowTrigger.WorkflowTriggerType.DOCUMENT_ADDED],
action={"ai_suggestion_fields": ["title", "not_a_field"]},
)
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
def test_api_create_apply_ai_suggestions_action_rejects_consumption_only(
self,
) -> None:
"""
GIVEN:
- API request to create an apply AI suggestions action whose only
trigger is consumption started, so there is no document content
to make suggestions from yet
WHEN:
- API is called
THEN:
- Correct HTTP 400 response
- No objects are created
"""
existing_count = Workflow.objects.count()
response = self._post_ai_suggestions_workflow(
trigger_types=[WorkflowTrigger.WorkflowTriggerType.CONSUMPTION],
action={"ai_suggestion_fields": ["title"]},
)
self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST)
self.assertEqual(Workflow.objects.count(), existing_count)
def test_api_create_apply_ai_suggestions_action_allows_extra_consumption_trigger(
self,
) -> None:
"""
GIVEN:
- API request to create an apply AI suggestions action with a
consumption trigger alongside a usable one
WHEN:
- API is called
THEN:
- The workflow is created, the action applies to the other trigger
"""
response = self._post_ai_suggestions_workflow(
trigger_types=[
WorkflowTrigger.WorkflowTriggerType.CONSUMPTION,
WorkflowTrigger.WorkflowTriggerType.DOCUMENT_ADDED,
],
action={"ai_suggestion_fields": ["title"]},
)
self.assertEqual(response.status_code, status.HTTP_201_CREATED)
def test_api_create_workflow_trigger_action_empty_fields(self) -> None:
"""
GIVEN:
+224
View File
@@ -5,8 +5,11 @@ from unittest import mock
import pikepdf
from django.contrib.auth.models import Group
from django.contrib.auth.models import Permission
from django.contrib.auth.models import User
from django.db import connection
from django.test import TestCase
from django.test.utils import CaptureQueriesContext
from guardian.shortcuts import assign_perm
from guardian.shortcuts import get_groups_with_perms
from guardian.shortcuts import get_users_with_perms
@@ -19,6 +22,7 @@ from documents.models import Document
from documents.models import DocumentType
from documents.models import StoragePath
from documents.models import Tag
from documents.permissions import set_permissions_for_objects
from documents.tests.utils import DirectoriesMixin
@@ -344,6 +348,100 @@ class TestBulkEdit(DirectoriesMixin, TestCase):
assert _cf_3 is not None
self.assertNotIn(self.doc3.id, _cf_3.value)
def test_modify_custom_fields_batches_field_lookup(self) -> None:
"""
GIVEN:
- Several documents are being bulk-edited to add several custom
fields at once
WHEN:
- modify_custom_fields runs
THEN:
- Each CustomField is resolved with one batched query total, not
once per (field, document) pair
"""
docs = [
Document.objects.create(checksum=f"batch-{i}", title=f"batch-{i}")
for i in range(6)
]
fields = [
CustomField.objects.create(
name=f"Batch Field {i}",
data_type=CustomField.FieldDataType.STRING,
)
for i in range(4)
]
with CaptureQueriesContext(connection) as ctx:
bulk_edit.modify_custom_fields(
[doc.id for doc in docs],
add_custom_fields=[field.id for field in fields],
remove_custom_fields=[],
)
field_lookups = [
q
for q in ctx.captured_queries
if 'FROM "documents_customfield"' in q["sql"]
]
self.assertEqual(
len(field_lookups),
1,
"Expected a single batched query to resolve the custom fields, "
f"got {len(field_lookups)}: {field_lookups}",
)
for doc in docs:
self.assertEqual(doc.custom_fields.count(), len(fields))
def test_modify_custom_fields_batches_document_lookup_for_documentlink(
self,
) -> None:
"""
GIVEN:
- Several documents are being bulk-edited to add a DOCUMENTLINK
custom field at once
WHEN:
- modify_custom_fields runs
THEN:
- The Document rows needed to reflect the symmetrical links are
resolved with one batched query total, not once per document
"""
docs = [
Document.objects.create(checksum=f"link-{i}", title=f"link-{i}")
for i in range(6)
]
target = Document.objects.create(checksum="link-target", title="link-target")
doclink_field = CustomField.objects.create(
name="Related",
data_type=CustomField.FieldDataType.DOCUMENTLINK,
)
with CaptureQueriesContext(connection) as ctx:
bulk_edit.modify_custom_fields(
[doc.id for doc in docs],
add_custom_fields={doclink_field.id: [target.id]},
remove_custom_fields=[],
)
single_document_lookups = [
q
for q in ctx.captured_queries
if 'FROM "documents_document"' in q["sql"]
and '"documents_document"."id" = ' in q["sql"]
]
self.assertEqual(
len(single_document_lookups),
0,
"Expected document rows to come from a batched query, not "
f"per-document lookups, got: {single_document_lookups}",
)
for doc in docs:
self.assertEqual(
doc.custom_fields.get(field=doclink_field).value,
[target.id],
)
def test_modify_custom_fields_doclink_self_link(self) -> None:
"""
GIVEN:
@@ -510,6 +608,132 @@ class TestBulkEdit(DirectoriesMixin, TestCase):
)
self.assertEqual(groups_with_perms.count(), 2)
@mock.patch("documents.tasks.bulk_update_documents.apply_async")
def test_set_permissions_query_count_independent_of_document_count(
self,
m,
) -> None:
"""
GIVEN:
- Many documents are being bulk-edited to set permissions at once
WHEN:
- set_permissions runs over a small batch vs. a much larger one
THEN:
- The number of queries issued is the same either way -- each
user/group is applied across all documents with one batched
call, not one call per (document, user) pair
"""
permissions = {
"view": {
"users": [self.user1.id, self.user2.id],
"groups": [self.group2.id],
},
"change": {
"users": [self.user1.id],
"groups": [self.group2.id],
},
}
def run_with_n_documents(n: int) -> int:
docs = [
Document.objects.create(checksum=f"perm-{n}-{i}", title=f"perm-{n}-{i}")
for i in range(n)
]
with CaptureQueriesContext(connection) as ctx:
bulk_edit.set_permissions(
[doc.id for doc in docs],
set_permissions=permissions,
owner=self.owner,
merge=False,
)
for doc in docs:
self.assertEqual(get_users_with_perms(doc).count(), 2)
self.assertEqual(get_groups_with_perms(doc).count(), 1)
return len(ctx.captured_queries)
small_batch_queries = run_with_n_documents(5)
large_batch_queries = run_with_n_documents(50)
self.assertEqual(
small_batch_queries,
large_batch_queries,
"Expected the same query count regardless of document count, got "
f"{small_batch_queries} queries for 5 documents vs. "
f"{large_batch_queries} for 50",
)
@mock.patch("documents.tasks.bulk_update_documents.apply_async")
def test_set_permissions_grants_direct_perm_even_if_already_granted_via_group(
self,
m,
) -> None:
"""
GIVEN:
- A user already has view access to a document via group
membership, with no direct grant of their own
WHEN:
- set_permissions explicitly grants that same user direct view
access via bulk_edit
THEN:
- A direct permission grant is created for the user, not skipped
because they already have equivalent access via the group
Regression test: guardian's queryset-aware assign_perm() (routed to
when the target is a list/queryset) skips creating a direct row for
anyone whose ObjectPermissionChecker.has_perm() already returns True
-- which includes group-derived access. The single-object assign_perm
this bulk path replaces has no such check; it always ensures a
direct row via get_or_create. Losing that guarantee would mean
revoking the group's grant later silently strips access that was
supposed to be explicit.
"""
self.doc1.owner = self.user1
self.doc1.save()
self.user1.groups.add(self.group1)
assign_perm("view_document", self.group1, self.doc1)
bulk_edit.set_permissions(
[self.doc1.id],
set_permissions={
"view": {"users": [self.user1.id], "groups": []},
},
merge=True,
)
direct_users = get_users_with_perms(
self.doc1,
only_with_perms_in=["view_document"],
with_group_users=False,
)
self.assertIn(self.user1, direct_users)
def test_set_permissions_for_objects_raises_for_unknown_action(self) -> None:
"""
GIVEN:
- An unrecognized permission action name with users to grant it
to
WHEN:
- set_permissions_for_objects is called
THEN:
- Permission.DoesNotExist is raised, not a silent no-op
Regression test: the endpoint that calls this
(BulkEditObjectPermissionsView) never actually validates action
names against the raw client-supplied permissions dict --
BulkEditObjectsSerializer._validate_permissions calls
validate_set_permissions() only for its side-effecting user/group id
checks and discards the filtered dict it returns -- so a bogus
action key reaches this function as-is. Resolving the Permission via
a bare `.filter()` (which returns empty instead of raising) would
silently drop the grant and report success.
"""
with self.assertRaises(Permission.DoesNotExist):
set_permissions_for_objects(
{"not_a_real_action": {"users": [self.user1.id], "groups": []}},
Document,
[self.doc1.pk],
)
@mock.patch("documents.models.Document.delete")
def test_delete_documents_old_uuid_field(self, m) -> None:
m.side_effect = Exception("Data too long for column 'transaction_id' at row 1")
+58
View File
@@ -0,0 +1,58 @@
from django.db import connection
from django.test import TestCase
from django.test.utils import CaptureQueriesContext
from documents.data_models import DocumentMetadataOverrides
from documents.models import CustomField
from documents.models import CustomFieldInstance
from documents.tests.factories import DocumentFactory
from documents.tests.utils import DirectoriesMixin
class TestDocumentMetadataOverridesFromDocument(DirectoriesMixin, TestCase):
def test_from_document_batches_custom_field_lookup_after_refresh_from_db(
self,
) -> None:
"""
GIVEN:
- A document has several custom field values
- The document instance has just been refreshed from the database,
which drops any prefetched related objects (as
send_websocket_document_updated does before building overrides)
WHEN:
- DocumentMetadataOverrides.from_document() reads the document's
custom field values
THEN:
- The referenced CustomField objects are resolved with a single
query, not one query per custom field
"""
doc = DocumentFactory(mime_type="application/pdf")
for i in range(5):
CustomFieldInstance.objects.create(
document=doc,
field=CustomField.objects.create(
name=f"Test Custom Field {i}",
data_type=CustomField.FieldDataType.STRING,
),
value_text="value",
)
doc.refresh_from_db()
with CaptureQueriesContext(connection) as ctx:
overrides = DocumentMetadataOverrides.from_document(doc)
self.assertEqual(len(overrides.custom_fields), 5)
unbatched_field_lookups = [
query
for query in ctx.captured_queries
if 'FROM "documents_customfield" WHERE "documents_customfield"."id"'
in query["sql"]
]
self.assertEqual(
unbatched_field_lookups,
[],
"Expected CustomField data to come from the CustomFieldInstance "
"join, not a separate per-instance lookup, "
f"got: {unbatched_field_lookups}",
)
-19
View File
@@ -385,25 +385,6 @@ class TestTaskFailureHandler:
task_failure_handler(task_id=None, exception=ValueError("x"), traceback=None)
@pytest.mark.django_db
class TestApplyAiSuggestionsTracking:
def test_records_the_document_it_is_for(self) -> None:
"""
The action queues one task per document, so the tracked record notes
which document it is for -- otherwise a bulk run is an indistinguishable
wall of identical entries in the tasks list.
"""
task_id = send_publish(
"documents.tasks.apply_ai_suggestions",
(),
{"action_id": 1, "document_id": 42},
)
task = PaperlessTask.objects.get(task_id=task_id)
assert task.task_type == PaperlessTask.TaskType.APPLY_AI_SUGGESTIONS
assert task.input_data == {"document_id": 42}
@pytest.mark.django_db
class TestTaskRevokedHandler:
def test_marks_task_revoked(self, mocker: pytest_mock.MockerFixture) -> None:
-108
View File
@@ -14,7 +14,6 @@ from documents.models import Correspondent
from documents.models import Document
from documents.models import DocumentType
from documents.models import Tag
from documents.models import WorkflowAction
from documents.sanity_checker import SanityCheckFailedException
from documents.sanity_checker import SanityCheckMessages
from documents.tests.test_classifier import dummy_preprocess
@@ -448,110 +447,3 @@ class TestAIIndex(DirectoriesMixin, TestCase):
rebuild=False,
document_ids=doc_ids,
)
class TestApplyAISuggestionsTask(DirectoriesMixin, TestCase):
def setUp(self) -> None:
super().setUp()
self.doc = Document.objects.create(
title="doc",
content="content",
checksum="apply-ai-suggestions",
)
self.action = WorkflowAction.objects.create(
type=WorkflowAction.WorkflowActionType.APPLY_AI_SUGGESTIONS,
ai_suggestion_fields=[WorkflowAction.AISuggestionField.TITLE],
)
def test_reindexes_without_sending_document_updated(self) -> None:
"""
GIVEN:
- An apply AI suggestions action that changes the document
WHEN:
- The task runs
THEN:
- The search index and caches are refreshed directly, deliberately
not via the document_updated signal: that re-runs updated
workflows, which for this action means queueing another LLM
query for a document it just changed, forever
"""
with (
mock.patch(
"documents.workflows.ai.apply_ai_suggestions_to_document",
return_value=["title"],
),
mock.patch("documents.tasks.index_document") as index_document,
mock.patch("documents.tasks.clear_document_caches") as clear_caches,
mock.patch("documents.tasks.document_updated") as document_updated,
):
tasks.apply_ai_suggestions(self.action.pk, self.doc.pk)
index_document.delay.assert_called_once_with(self.doc.pk)
clear_caches.assert_called_once_with(self.doc.pk)
document_updated.send.assert_not_called()
def test_no_changes_skips_reindex(self) -> None:
"""
GIVEN:
- An apply AI suggestions action that changes nothing
WHEN:
- The task runs
THEN:
- No reindexing work is queued
"""
with (
mock.patch(
"documents.workflows.ai.apply_ai_suggestions_to_document",
return_value=[],
),
mock.patch("documents.tasks.index_document") as index_document,
):
tasks.apply_ai_suggestions(self.action.pk, self.doc.pk)
index_document.delay.assert_not_called()
@override_settings(AI_ENABLED=True, LLM_EMBEDDING_BACKEND="huggingface")
def test_updates_llm_index_when_enabled(self) -> None:
"""
GIVEN:
- An apply AI suggestions action that changes the document
- The LLM index is enabled
WHEN:
- The task runs
THEN:
- The document is updated in the LLM index too
"""
with (
mock.patch(
"documents.workflows.ai.apply_ai_suggestions_to_document",
return_value=["title"],
),
mock.patch("documents.tasks.index_document"),
mock.patch(
"documents.tasks.update_document_in_llm_index",
) as update_in_llm_index,
):
tasks.apply_ai_suggestions(self.action.pk, self.doc.pk)
update_in_llm_index.apply_async.assert_called_once()
def test_deleted_document_is_a_noop(self) -> None:
"""
GIVEN:
- A document that was deleted between the workflow running and the
queued task starting
WHEN:
- The task runs
THEN:
- It logs and exits rather than raising
"""
with (
mock.patch(
"documents.workflows.ai.apply_ai_suggestions_to_document",
) as apply_suggestions,
self.assertLogs("paperless.tasks", level="WARNING") as cm,
):
tasks.apply_ai_suggestions(self.action.pk, self.doc.pk + 1000)
apply_suggestions.assert_not_called()
self.assertIn("no longer exists", "".join(cm.output))
-413
View File
@@ -31,9 +31,7 @@ from documents.file_handling import create_source_path_directory
from documents.file_handling import generate_filename
from documents.file_handling import generate_unique_filename
from documents.signals.handlers import run_workflows
from documents.workflows.ai import apply_ai_suggestions_to_document
from documents.workflows.webhooks import send_webhook
from paperless_ai.exceptions import LLMTimeoutError
if TYPE_CHECKING:
from django.db.models import QuerySet
@@ -5490,414 +5488,3 @@ class TestRemoteOCRWorkflowAction(DirectoriesMixin, SampleDirMixin, APITestCase)
)
self.assertIn("only applies to consumption triggers", "".join(cm.output))
SUGGESTIONS = {
"title": "Suggested Title",
"tags": ["Existing Tag", "Suggested Tag"],
"correspondents": ["Existing Correspondent", "Suggested Correspondent"],
"document_types": ["Suggested Document Type"],
"storage_paths": ["Suggested Storage Path"],
"dates": ["2024-03-05"],
}
ALL_SUGGESTION_FIELDS = [
WorkflowAction.AISuggestionField.TITLE,
WorkflowAction.AISuggestionField.TAGS,
WorkflowAction.AISuggestionField.CORRESPONDENT,
WorkflowAction.AISuggestionField.DOCUMENT_TYPE,
WorkflowAction.AISuggestionField.STORAGE_PATH,
WorkflowAction.AISuggestionField.CREATED,
]
@override_settings(AI_ENABLED=True)
class TestApplyAISuggestionsWorkflowAction(
DirectoriesMixin,
SampleDirMixin,
APITestCase,
):
def setUp(self) -> None:
super().setUp()
self.user = User.objects.create(username="ai-user")
self.doc = Document.objects.create(
title="original.pdf",
content="the document content",
checksum="ai-suggestions-checksum",
mime_type="application/pdf",
created=datetime.date(2020, 1, 1),
owner=self.user,
)
def make_action(self, **kwargs) -> WorkflowAction:
return WorkflowAction.objects.create(
type=WorkflowAction.WorkflowActionType.APPLY_AI_SUGGESTIONS,
ai_suggestion_fields=kwargs.pop(
"ai_suggestion_fields",
ALL_SUGGESTION_FIELDS,
),
**kwargs,
)
def make_workflow(self, action: WorkflowAction, trigger_type) -> Workflow:
trigger = WorkflowTrigger.objects.create(type=trigger_type)
w = Workflow.objects.create(name="Apply AI suggestions", order=0)
w.triggers.add(trigger)
w.actions.add(action)
w.save()
return w
def apply(self, action: WorkflowAction) -> list[str]:
with mock.patch(
"documents.workflows.ai.get_ai_document_classification",
return_value=SUGGESTIONS,
):
changed = apply_ai_suggestions_to_document(action, self.doc)
self.doc.refresh_from_db()
return changed
def test_document_added_trigger_queues_task(self) -> None:
"""
GIVEN:
- A document added workflow with an apply AI suggestions action
WHEN:
- A matching document is added
THEN:
- The work is queued rather than run inline, so a slow LLM query
cannot stall the rest of the workflow run
"""
action = self.make_action()
self.make_workflow(action, WorkflowTrigger.WorkflowTriggerType.DOCUMENT_ADDED)
with mock.patch("documents.tasks.apply_ai_suggestions.delay") as delay:
run_workflows(
WorkflowTrigger.WorkflowTriggerType.DOCUMENT_ADDED,
self.doc,
)
delay.assert_called_once_with(action_id=action.pk, document_id=self.doc.pk)
def test_consumption_trigger_is_ignored(self) -> None:
"""
GIVEN:
- A workflow with an apply AI suggestions action and a consumption
trigger alongside a valid one
WHEN:
- The consumption trigger fires
THEN:
- The action is skipped, since the document has not been parsed
yet and so has no content to make suggestions from
"""
action = self.make_action()
w = self.make_workflow(
action,
WorkflowTrigger.WorkflowTriggerType.DOCUMENT_ADDED,
)
w.triggers.add(
WorkflowTrigger.objects.create(
type=WorkflowTrigger.WorkflowTriggerType.CONSUMPTION,
),
)
test_file = shutil.copy(
self.SAMPLE_DIR / "simple.pdf",
self.dirs.scratch_dir / "simple.pdf",
)
with (
mock.patch("documents.tasks.apply_ai_suggestions.delay") as delay,
self.assertLogs("paperless.handlers", level="DEBUG") as cm,
):
run_workflows(
WorkflowTrigger.WorkflowTriggerType.CONSUMPTION,
ConsumableDocument(
source=DocumentSource.ConsumeFolder,
original_file=test_file,
),
overrides=DocumentMetadataOverrides(),
)
delay.assert_not_called()
self.assertIn("does not apply to consumption triggers", "".join(cm.output))
def test_no_selected_fields_does_nothing(self) -> None:
"""
GIVEN:
- An action with no suggestion fields selected
WHEN:
- The action is applied
THEN:
- Nothing is changed and it is logged
"""
action = self.make_action(ai_suggestion_fields=[])
with self.assertLogs("paperless.workflows.ai", level="WARNING") as cm:
changed = self.apply(action)
self.assertEqual(changed, [])
self.assertIn("no AI suggestion fields selected", "".join(cm.output))
@override_settings(AI_ENABLED=False)
def test_ai_disabled_does_nothing(self) -> None:
"""
GIVEN:
- An action on an install where AI has since been disabled
WHEN:
- The action is applied
THEN:
- Nothing is changed and it is logged
"""
action = self.make_action()
with self.assertLogs("paperless.workflows.ai", level="ERROR") as cm:
changed = self.apply(action)
self.assertEqual(changed, [])
self.assertIn("AI is not enabled", "".join(cm.output))
def test_invalid_configuration_leaves_document_untouched(self) -> None:
"""
GIVEN:
- An AI backend that is misconfigured
WHEN:
- The action is applied
THEN:
- The failure is logged and the document is left alone. It is not
re-raised, because retrying will not fix a bad configuration
"""
action = self.make_action()
with (
mock.patch(
"documents.workflows.ai.get_ai_document_classification",
side_effect=ValueError("nope"),
),
self.assertLogs("paperless.workflows.ai", level="ERROR") as cm,
):
changed = apply_ai_suggestions_to_document(action, self.doc)
self.assertEqual(changed, [])
self.doc.refresh_from_db()
self.assertEqual(self.doc.title, "original.pdf")
self.assertIn("Invalid AI configuration", "".join(cm.output))
def test_transient_llm_failure_is_raised_for_retry(self) -> None:
"""
GIVEN:
- An LLM backend that times out, or rate limits the request
WHEN:
- The action is applied
THEN:
- The error propagates so the queued task can back off and retry,
rather than silently dropping this document's suggestions
"""
action = self.make_action()
with (
mock.patch(
"documents.workflows.ai.get_ai_document_classification",
side_effect=LLMTimeoutError(),
),
self.assertRaises(LLMTimeoutError),
):
apply_ai_suggestions_to_document(action, self.doc)
self.doc.refresh_from_db()
self.assertEqual(self.doc.title, "original.pdf")
def test_only_matching_objects_are_applied(self) -> None:
"""
GIVEN:
- An action without create missing, and only some of the suggested
objects existing
WHEN:
- The action is applied
THEN:
- Only the existing objects are assigned, unmatched suggestions are
dropped rather than creating anything
"""
tag = Tag.objects.create(name="Existing Tag", owner=self.user)
correspondent = Correspondent.objects.create(
name="Existing Correspondent",
owner=self.user,
)
action = self.make_action(ai_overwrite_existing=True)
changed = self.apply(action)
self.assertEqual(self.doc.correspondent, correspondent)
self.assertEqual(list(self.doc.tags.all()), [tag])
# Nothing matched for these and create missing is off
self.assertIsNone(self.doc.document_type)
self.assertIsNone(self.doc.storage_path)
self.assertNotIn("document_type", changed)
self.assertEqual(Tag.objects.count(), 1)
self.assertEqual(Correspondent.objects.count(), 1)
def test_create_missing_creates_objects_owned_by_document_owner(self) -> None:
"""
GIVEN:
- An action with create missing enabled
WHEN:
- The action is applied and suggestions match nothing
THEN:
- Tags, correspondents and document types are created, owned by the
document owner so they stay private to them
- Storage paths are never created, since a path template cannot be
inferred from a name
"""
action = self.make_action(
ai_create_missing=True,
ai_overwrite_existing=True,
)
changed = self.apply(action)
self.assertEqual(
sorted(t.name for t in self.doc.tags.all()),
["Existing Tag", "Suggested Tag"],
)
self.assertEqual(self.doc.correspondent.name, "Existing Correspondent")
self.assertEqual(self.doc.correspondent.owner, self.user)
self.assertEqual(self.doc.document_type.name, "Suggested Document Type")
self.assertEqual(self.doc.document_type.owner, self.user)
self.assertIsNone(self.doc.storage_path)
self.assertFalse(StoragePath.objects.exists())
self.assertNotIn("storage_path", changed)
def test_overwrite_disabled_keeps_existing_values(self) -> None:
"""
GIVEN:
- An action without overwrite existing
- A document that already has a title, created date and
correspondent
WHEN:
- The action is applied
THEN:
- The existing values are kept, only the empty document type is
filled in
"""
existing = Correspondent.objects.create(name="Mine", owner=self.user)
self.doc.correspondent = existing
self.doc.save()
action = self.make_action(ai_create_missing=True)
changed = self.apply(action)
self.assertEqual(self.doc.title, "original.pdf")
self.assertEqual(self.doc.created, datetime.date(2020, 1, 1))
self.assertEqual(self.doc.correspondent, existing)
self.assertEqual(self.doc.document_type.name, "Suggested Document Type")
self.assertNotIn("title", changed)
self.assertNotIn("correspondent", changed)
def test_overwrite_enabled_replaces_existing_values(self) -> None:
"""
GIVEN:
- An action with overwrite existing
- A document that already has a title and created date
WHEN:
- The action is applied
THEN:
- The suggested values replace them
"""
action = self.make_action(
ai_create_missing=True,
ai_overwrite_existing=True,
)
changed = self.apply(action)
self.assertEqual(self.doc.title, "Suggested Title")
self.assertEqual(self.doc.created, datetime.date(2024, 3, 5))
self.assertIn("title", changed)
self.assertIn("created", changed)
def test_tags_are_added_not_replaced(self) -> None:
"""
GIVEN:
- A document that already has a tag unrelated to the suggestions
WHEN:
- The action is applied with overwrite existing enabled
THEN:
- The existing tag is kept, since suggested tags are always
additive regardless of the overwrite setting
"""
kept = Tag.objects.create(name="Do Not Remove", owner=self.user)
self.doc.tags.add(kept)
Tag.objects.create(name="Existing Tag", owner=self.user)
action = self.make_action(ai_overwrite_existing=True)
self.apply(action)
self.assertEqual(
sorted(t.name for t in self.doc.tags.all()),
["Do Not Remove", "Existing Tag"],
)
def test_unselected_fields_are_untouched(self) -> None:
"""
GIVEN:
- An action that only selects the title
WHEN:
- The action is applied
THEN:
- Only the title changes, even though the LLM suggested everything
"""
action = self.make_action(
ai_suggestion_fields=[WorkflowAction.AISuggestionField.TITLE],
ai_create_missing=True,
ai_overwrite_existing=True,
)
changed = self.apply(action)
self.assertEqual(changed, ["title"])
self.assertEqual(self.doc.title, "Suggested Title")
self.assertEqual(self.doc.tags.count(), 0)
self.assertIsNone(self.doc.correspondent)
self.assertEqual(self.doc.created, datetime.date(2020, 1, 1))
def test_another_users_private_objects_are_not_matched(self) -> None:
"""
GIVEN:
- A suggested tag name that exists, but is owned by someone else
WHEN:
- The action is applied
THEN:
- It is not assigned, because the document owner cannot see it
"""
other = User.objects.create(username="someone-else")
Tag.objects.create(name="Existing Tag", owner=other)
action = self.make_action(
ai_suggestion_fields=[WorkflowAction.AISuggestionField.TAGS],
)
self.apply(action)
self.assertEqual(self.doc.tags.count(), 0)
def test_unparsable_dates_are_skipped(self) -> None:
"""
GIVEN:
- Suggested dates that are not all valid
WHEN:
- The action is applied
THEN:
- The first usable date is applied and the rest ignored
"""
action = self.make_action(
ai_suggestion_fields=[WorkflowAction.AISuggestionField.CREATED],
ai_overwrite_existing=True,
)
with mock.patch(
"documents.workflows.ai.get_ai_document_classification",
return_value={**SUGGESTIONS, "dates": ["not a date", "2019-07-04"]},
):
changed = apply_ai_suggestions_to_document(action, self.doc)
self.doc.refresh_from_db()
self.assertEqual(changed, ["created"])
self.assertEqual(self.doc.created, datetime.date(2019, 7, 4))
+23 -16
View File
@@ -178,7 +178,7 @@ from documents.permissions import has_perms_owner_aware
from documents.permissions import has_system_status_permission
from documents.permissions import permitted_document_ids
from documents.permissions import permitted_object_ids
from documents.permissions import set_permissions_for_object
from documents.permissions import set_permissions_for_objects
from documents.plugins.date_parsing import get_date_parser
from documents.schema import generate_object_with_permissions_schema
from documents.search import SearchHit
@@ -247,7 +247,6 @@ from paperless.serialisers import GroupSerializer
from paperless.serialisers import UserSerializer
from paperless.views import StandardPagination
from paperless_ai.ai_classifier import get_ai_document_classification
from paperless_ai.ai_classifier import get_llm_output_language
from paperless_ai.chat import stream_chat_with_documents
from paperless_ai.exceptions import LLMTimeoutError
from paperless_ai.matching import extract_unmatched_names
@@ -666,6 +665,20 @@ class TagViewSet(PermissionsAwareDocumentCountMixin, ModelViewSet[Tag]):
update_document_parent_tags(tag, new_parent)
def _get_llm_output_language(ai_config: AIConfig, request) -> str | None:
output_language = ai_config.llm_output_language
if (
not output_language
and hasattr(request.user, "ui_settings")
and isinstance(
request.user.ui_settings.settings,
dict,
)
):
output_language = request.user.ui_settings.settings.get("language")
return output_language
@extend_schema_view(**generate_object_with_permissions_schema(DocumentTypeSerializer))
class DocumentTypeViewSet(
PermissionsAwareDocumentCountMixin,
@@ -1528,10 +1541,7 @@ class DocumentViewSet(
if not ai_config.ai_enabled:
return HttpResponseBadRequest("AI is required for this feature")
output_language = get_llm_output_language(
ai_config=ai_config,
user=request.user,
)
output_language = _get_llm_output_language(ai_config=ai_config, request=request)
llm_cache_backend = ":".join(
part
for part in (
@@ -2320,10 +2330,7 @@ class ChatStreamingView(GenericAPIView[Any]):
id__in=permitted_document_ids(request.user),
)
output_language = get_llm_output_language(
ai_config=ai_config,
user=request.user,
)
output_language = _get_llm_output_language(ai_config=ai_config, request=request)
response = StreamingHttpResponse(
stream_chat_with_documents(
@@ -4910,12 +4917,12 @@ class BulkEditObjectsView(PassUserMixin):
qs_owner_update.update(owner=owner)
if "permissions" in serializer.validated_data:
for obj in qs:
set_permissions_for_object(
permissions=permissions,
object=obj,
merge=merge,
)
set_permissions_for_objects(
permissions=permissions,
model=object_class,
pks=qs.values_list("pk", flat=True),
merge=merge,
)
except Exception as e:
logger.warning(
-241
View File
@@ -1,241 +0,0 @@
import logging
from datetime import date
from datetime import datetime
from django.contrib.auth.models import User
from documents.models import Correspondent
from documents.models import Document
from documents.models import DocumentType
from documents.models import StoragePath
from documents.models import Tag
from documents.models import WorkflowAction
from paperless.config import AIConfig
from paperless_ai.ai_classifier import get_ai_document_classification
from paperless_ai.ai_classifier import get_llm_output_language
from paperless_ai.matching import extract_unmatched_names
from paperless_ai.matching import match_correspondents_by_name
from paperless_ai.matching import match_document_types_by_name
from paperless_ai.matching import match_storage_paths_by_name
from paperless_ai.matching import match_tags_by_name
logger = logging.getLogger("paperless.workflows.ai")
AISuggestionField = WorkflowAction.AISuggestionField
# Tags use m2m relation instead
DIRECT_FIELDS: dict[str, str] = {
AISuggestionField.TITLE: "title",
AISuggestionField.CORRESPONDENT: "correspondent",
AISuggestionField.DOCUMENT_TYPE: "document_type",
AISuggestionField.STORAGE_PATH: "storage_path",
AISuggestionField.CREATED: "created",
}
def resolve_date(dates: list[str]) -> date | None:
"""
First usable date out of the suggestions, which are expected as
YYYY-MM-DD. Document.created is a DateField, so only one can be applied.
"""
for value in dates:
try:
return datetime.strptime(value, "%Y-%m-%d").date()
except (TypeError, ValueError):
logger.debug("Ignoring unparsable suggested date %s", value)
return None
def resolve_object(
model,
names: list[str],
matched: list,
*,
create_missing: bool,
owner: User | None,
):
"""
Single object from a suggestion list. The best match if there was one, else
optionally a newly-created object. StoragePaths are excluded.
"""
if matched:
return matched[0]
if not create_missing or model is StoragePath:
return None
unmatched = extract_unmatched_names(names, matched)
if not unmatched:
return None
# (name, owner) is what MatchingModel is unique on
obj, created = model.objects.get_or_create(
name=unmatched[0][:128],
owner=owner,
)
if created:
logger.info("Created %s '%s' from AI suggestion", model.__name__, obj.name)
return obj
def resolve_tags(
names: list[str],
matched: list[Tag],
*,
create_missing: bool,
owner: User | None,
) -> list[Tag]:
"""
Matched tags, plus newly created ones if create_missing is set.
"""
tags = list(matched)
if not create_missing:
return tags
for name in extract_unmatched_names(names, matched):
tag, created = Tag.objects.get_or_create(
name=name[:128],
owner=owner,
)
if created:
logger.info("Created tag '%s' from AI suggestion", tag.name)
tags.append(tag)
return tags
def apply_ai_suggestions_to_document(
action: WorkflowAction,
document: Document,
logging_group=None,
) -> list[str]:
"""
Get suggestions about `document` and write the chosen fields.
Returns the names of the fields that were actually changed.
"""
selected = set(action.ai_suggestion_fields or [])
if not selected:
logger.warning(
"Workflow action %s has no AI suggestion fields selected, skipping",
action.pk,
extra={"group": logging_group},
)
return []
ai_config = AIConfig()
if not ai_config.ai_enabled:
logger.error(
"AI is not enabled, cannot apply AI suggestions for document %s",
document.pk,
extra={"group": logging_group},
)
return []
# Workflows run without a user, so we use the document owner
owner = document.owner
try:
suggestions = get_ai_document_classification(
document,
owner,
get_llm_output_language(ai_config, owner),
)
except ValueError:
# A bad AI config will not fix itself, so swallow it rather than
# letting the caller retry. Timeouts, rate limits, network errors etc
# propagate so the queued task can back off and try again.
logger.exception(
"Invalid AI configuration, cannot get suggestions for document %s",
document.pk,
extra={"group": logging_group},
)
return []
overwrite = action.ai_overwrite_existing
create_missing = action.ai_create_missing
updated_fields: list[str] = []
def should_set(field: str) -> bool:
# The field is selected and (overwrite or it's empty)
return field in selected and (
overwrite or getattr(document, DIRECT_FIELDS[field]) in (None, "")
)
if should_set(AISuggestionField.TITLE):
title = (suggestions.get("title") or "").strip()
if title:
# title is capped at 128 characters
document.title = title[:128]
updated_fields.append("title")
if should_set(AISuggestionField.CORRESPONDENT):
names = suggestions.get("correspondents", [])
correspondent = resolve_object(
Correspondent,
names,
match_correspondents_by_name(names, owner),
create_missing=create_missing,
owner=owner,
)
if correspondent:
document.correspondent = correspondent
updated_fields.append("correspondent")
if should_set(AISuggestionField.DOCUMENT_TYPE):
names = suggestions.get("document_types", [])
document_type = resolve_object(
DocumentType,
names,
match_document_types_by_name(names, owner),
create_missing=create_missing,
owner=owner,
)
if document_type:
document.document_type = document_type
updated_fields.append("document_type")
if should_set(AISuggestionField.STORAGE_PATH):
names = suggestions.get("storage_paths", [])
storage_path = resolve_object(
StoragePath,
names,
match_storage_paths_by_name(names, owner),
create_missing=create_missing,
owner=owner,
)
if storage_path:
document.storage_path = storage_path
updated_fields.append("storage_path")
if should_set(AISuggestionField.CREATED):
created = resolve_date(suggestions.get("dates", []))
if created:
document.created = created
updated_fields.append("created")
if updated_fields:
# save fields and update modified
document.save(update_fields=[*updated_fields, "modified"])
if AISuggestionField.TAGS in selected:
names = suggestions.get("tags", [])
tags = resolve_tags(
names,
match_tags_by_name(names, owner),
create_missing=create_missing,
owner=owner,
)
if tags:
# Suggested tags are always added, so overwrite_existing
# does not really apply here
document.add_nested_tags(tags)
updated_fields.append("tags")
logger.info(
"Applied AI suggestions %s to document %s",
updated_fields or "(none)",
document.pk,
extra={"group": logging_group},
)
return updated_fields
-16
View File
@@ -47,22 +47,6 @@ def get_language_name(language_code: str) -> str:
return language_code
def get_llm_output_language(ai_config: AIConfig, user: User | None) -> str | None:
"""
Language to localize LLM output into: the configured language, falling back
to the user's own UI language when unset.
"""
output_language = ai_config.llm_output_language
if (
not output_language
and user is not None
and hasattr(user, "ui_settings")
and isinstance(user.ui_settings.settings, dict)
):
output_language = user.ui_settings.settings.get("language")
return output_language
def build_prompt_without_rag(
document: Document,
config: AIConfig,
+7 -17
View File
@@ -11,7 +11,7 @@ from documents.models import Correspondent
from documents.models import DocumentType
from documents.models import StoragePath
from documents.models import Tag
from documents.permissions import permitted_object_ids
from documents.permissions import get_objects_for_user_owner_aware
from documents.permissions import restrict_queryset_to_visible
MATCH_THRESHOLD = 0.8
@@ -63,40 +63,30 @@ def resolve_storage_path_ids(ids: list[int], user: User | None) -> list[StorageP
def _match_by_name(
names: list[str],
user: User | None,
user: User,
model: type[ModelT],
perm: str,
) -> list[ModelT]:
# A workflow may have no user. In that case permitted_object_ids limits
# matching to unowned objects, avoiding another user's private taxonomy.
queryset = model.objects.filter(
pk__in=permitted_object_ids(user, model, perm),
)
queryset = get_objects_for_user_owner_aware(user, [perm], model)
return _match_names_to_queryset(names, queryset)
def match_tags_by_name(names: list[str], user: User | None) -> list[Tag]:
def match_tags_by_name(names: list[str], user: User) -> list[Tag]:
return _match_by_name(names, user, Tag, "view_tag")
def match_correspondents_by_name(
names: list[str],
user: User | None,
user: User,
) -> list[Correspondent]:
return _match_by_name(names, user, Correspondent, "view_correspondent")
def match_document_types_by_name(
names: list[str],
user: User | None,
) -> list[DocumentType]:
def match_document_types_by_name(names: list[str], user: User) -> list[DocumentType]:
return _match_by_name(names, user, DocumentType, "view_documenttype")
def match_storage_paths_by_name(
names: list[str],
user: User | None,
) -> list[StoragePath]:
def match_storage_paths_by_name(names: list[str], user: User) -> list[StoragePath]:
return _match_by_name(names, user, StoragePath, "view_storagepath")
+19 -6
View File
@@ -1,4 +1,5 @@
from collections.abc import Callable
from unittest.mock import patch
import pytest
import pytest_mock
@@ -44,25 +45,33 @@ class TestAIMatching(TestCase):
self.storage_path1 = StoragePath.objects.create(name="Test Storage Path 1")
self.storage_path2 = StoragePath.objects.create(name="Test Storage Path 2")
def test_match_tags_by_name(self) -> None:
@patch("paperless_ai.matching.get_objects_for_user_owner_aware")
def test_match_tags_by_name(self, mock_get_objects) -> None:
mock_get_objects.return_value = Tag.objects.all()
names = ["Test Tag 1", "Nonexistent Tag"]
result = match_tags_by_name(names, user=None)
self.assertEqual(len(result), 1)
self.assertEqual(result[0].name, "Test Tag 1")
def test_match_correspondents_by_name(self) -> None:
@patch("paperless_ai.matching.get_objects_for_user_owner_aware")
def test_match_correspondents_by_name(self, mock_get_objects) -> None:
mock_get_objects.return_value = Correspondent.objects.all()
names = ["Test Correspondent 1", "Nonexistent Correspondent"]
result = match_correspondents_by_name(names, user=None)
self.assertEqual(len(result), 1)
self.assertEqual(result[0].name, "Test Correspondent 1")
def test_match_document_types_by_name(self) -> None:
@patch("paperless_ai.matching.get_objects_for_user_owner_aware")
def test_match_document_types_by_name(self, mock_get_objects) -> None:
mock_get_objects.return_value = DocumentType.objects.all()
names = ["Test Document Type 1", "Nonexistent Document Type"]
result = match_document_types_by_name(names, user=None)
self.assertEqual(len(result), 1)
self.assertEqual(result[0].name, "Test Document Type 1")
def test_match_storage_paths_by_name(self) -> None:
@patch("paperless_ai.matching.get_objects_for_user_owner_aware")
def test_match_storage_paths_by_name(self, mock_get_objects) -> None:
mock_get_objects.return_value = StoragePath.objects.all()
names = ["Test Storage Path 1", "Nonexistent Storage Path"]
result = match_storage_paths_by_name(names, user=None)
self.assertEqual(len(result), 1)
@@ -74,12 +83,16 @@ class TestAIMatching(TestCase):
unmatched_names = extract_unmatched_names(llm_names, matched_objects)
self.assertEqual(unmatched_names, ["Nonexistent Tag"])
def test_match_tags_by_name_with_empty_names(self) -> None:
@patch("paperless_ai.matching.get_objects_for_user_owner_aware")
def test_match_tags_by_name_with_empty_names(self, mock_get_objects) -> None:
mock_get_objects.return_value = Tag.objects.all()
names = [None, "", " "]
result = match_tags_by_name(names, user=None)
self.assertEqual(result, [])
def test_match_tags_with_fuzzy_matching(self) -> None:
@patch("paperless_ai.matching.get_objects_for_user_owner_aware")
def test_match_tags_with_fuzzy_matching(self, mock_get_objects) -> None:
mock_get_objects.return_value = Tag.objects.all()
names = ["Test Taag 1", "Teest Tag 2"]
result = match_tags_by_name(names, user=None)
self.assertEqual(len(result), 2)
Generated
+4 -14
View File
@@ -4,11 +4,11 @@ requires-python = ">=3.11"
resolution-markers = [
"python_full_version >= '3.15' and sys_platform == 'darwin'",
"python_full_version >= '3.15' and sys_platform == 'linux'",
"python_full_version == '3.14.*' and platform_machine == 'x86_64' and sys_platform == 'linux'",
"python_full_version == '3.14.*' and platform_machine == 'aarch64' and sys_platform == 'linux'",
"python_full_version >= '3.12' and python_full_version < '3.15' and sys_platform == 'darwin'",
"python_full_version == '3.12.*' and platform_machine == 'x86_64' and sys_platform == 'linux'",
"python_full_version == '3.12.*' and platform_machine == 'aarch64' and sys_platform == 'linux'",
"python_full_version == '3.14.*' and platform_machine == 'x86_64' and sys_platform == 'linux'",
"python_full_version == '3.14.*' and platform_machine == 'aarch64' and sys_platform == 'linux'",
"(python_full_version >= '3.12' and python_full_version < '3.15' and platform_machine != 'aarch64' and platform_machine != 'x86_64' and sys_platform == 'linux') or (python_full_version == '3.13.*' and platform_machine == 'aarch64' and sys_platform == 'linux') or (python_full_version == '3.13.*' and platform_machine == 'x86_64' and sys_platform == 'linux')",
"python_full_version < '3.12' and sys_platform == 'darwin'",
"python_full_version < '3.12' and sys_platform == 'linux'",
@@ -1175,14 +1175,6 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/9d/76/d08f5c79f7643dbff4512605c28b75481966ed6c8cc9b397c9dd2ee91cd1/drf_spectacular_sidecar-2026.7.1-py3-none-any.whl", hash = "sha256:bc6d50c9b64660e45e09296d39553b3e759eedd825fc41d631ee5b3f88e0c5de", size = 2617384, upload-time = "2026-07-01T13:39:03.787Z" },
]
[[package]]
name = "drf-writable-nested"
version = "0.7.2"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/e8/57/df87d92fbfc3f0f2ef1a49c47f2a83389a4a13b7acf62b8bf7b223627d82/drf_writable_nested-0.7.2-py3-none-any.whl", hash = "sha256:4a3d2737c1cbfafa690e30236b169112e5b23cfe3d288f3992b0651a1b828c4d", size = 10570, upload-time = "2025-03-10T19:59:05.482Z" },
]
[[package]]
name = "execnet"
version = "2.1.2"
@@ -2895,7 +2887,6 @@ dependencies = [
{ name = "djangorestframework" },
{ name = "drf-spectacular" },
{ name = "drf-spectacular-sidecar" },
{ name = "drf-writable-nested" },
{ name = "filelock" },
{ name = "flower" },
{ name = "gotenberg-client" },
@@ -3045,7 +3036,6 @@ requires-dist = [
{ name = "djangorestframework", specifier = "~=3.16" },
{ name = "drf-spectacular", specifier = "~=0.30" },
{ name = "drf-spectacular-sidecar", specifier = "~=2026.7.1" },
{ name = "drf-writable-nested", specifier = "~=0.7.1" },
{ name = "filelock", specifier = "~=3.32.0" },
{ name = "flower", specifier = "~=2.0.1" },
{ name = "gotenberg-client", specifier = "~=0.14.0" },
@@ -5014,10 +5004,10 @@ version = "2.13.0+cpu"
source = { registry = "https://download.pytorch.org/whl/cpu" }
resolution-markers = [
"python_full_version >= '3.15' and sys_platform == 'linux'",
"python_full_version == '3.12.*' and platform_machine == 'x86_64' and sys_platform == 'linux'",
"python_full_version == '3.12.*' and platform_machine == 'aarch64' and sys_platform == 'linux'",
"python_full_version == '3.14.*' and platform_machine == 'x86_64' and sys_platform == 'linux'",
"python_full_version == '3.14.*' and platform_machine == 'aarch64' and sys_platform == 'linux'",
"python_full_version == '3.12.*' and platform_machine == 'x86_64' and sys_platform == 'linux'",
"python_full_version == '3.12.*' and platform_machine == 'aarch64' and sys_platform == 'linux'",
"(python_full_version >= '3.12' and python_full_version < '3.15' and platform_machine != 'aarch64' and platform_machine != 'x86_64' and sys_platform == 'linux') or (python_full_version == '3.13.*' and platform_machine == 'aarch64' and sys_platform == 'linux') or (python_full_version == '3.13.*' and platform_machine == 'x86_64' and sys_platform == 'linux')",
"python_full_version < '3.12' and sys_platform == 'linux'",
]