diff --git a/docs/usage.md b/docs/usage.md index 20d4c1189..35316dfaf 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -667,6 +667,35 @@ 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/). diff --git a/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.html b/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.html index 3827c41e6..af8e9a72d 100644 --- a/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.html +++ b/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.html @@ -462,6 +462,45 @@ } + @case (WorkflowActionType.ApplyAiSuggestions) { +
+
+

The document will be sent to the configured AI service for suggestions. Consider costs and privacy.

+ +
+
+
+
+ +
+
+ +
+
+ } } diff --git a/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.spec.ts b/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.spec.ts index dfe713c6d..d9c365673 100644 --- a/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.spec.ts +++ b/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.spec.ts @@ -22,6 +22,7 @@ import { } from 'src/app/data/matching-model' import { Workflow } from 'src/app/data/workflow' import { + AISuggestionField, WorkflowAction, WorkflowActionType, } from 'src/app/data/workflow-action' @@ -49,6 +50,7 @@ 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, @@ -239,14 +241,15 @@ describe('WorkflowEditDialogComponent', () => { SCHEDULE_DATE_FIELD_OPTIONS ) - // Email disabled + // Email, remote OCR and AI all 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.RemoteOcr && + a.id !== WorkflowActionType.ApplyAiSuggestions ) ) }) @@ -344,6 +347,125 @@ 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() diff --git a/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.ts b/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.ts index ed66e61dc..bb8cc58fa 100644 --- a/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.ts +++ b/src-ui/src/app/components/common/edit-dialog/workflow-edit-dialog/workflow-edit-dialog.component.ts @@ -30,6 +30,7 @@ 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' @@ -152,6 +153,37 @@ 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 { @@ -576,6 +608,24 @@ 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) @@ -1227,6 +1277,11 @@ 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 } ) @@ -1316,6 +1371,10 @@ 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) @@ -1369,6 +1428,9 @@ 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) diff --git a/src-ui/src/app/data/workflow-action.ts b/src-ui/src/app/data/workflow-action.ts index 09ef5418f..abcb590c1 100644 --- a/src-ui/src/app/data/workflow-action.ts +++ b/src-ui/src/app/data/workflow-action.ts @@ -8,6 +8,17 @@ 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 { @@ -102,4 +113,10 @@ export interface WorkflowAction extends ObjectWithId { webhook?: WorkflowActionWebhook passwords?: string[] + + ai_suggestion_fields?: AISuggestionField[] + + ai_create_missing?: boolean + + ai_overwrite_existing?: boolean } diff --git a/src/documents/migrations/0025_workflowaction_apply_ai_suggestions.py b/src/documents/migrations/0025_workflowaction_apply_ai_suggestions.py new file mode 100644 index 000000000..d00dc3817 --- /dev/null +++ b/src/documents/migrations/0025_workflowaction_apply_ai_suggestions.py @@ -0,0 +1,84 @@ +# 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", + ), + ), + ] diff --git a/src/documents/models.py b/src/documents/models.py index 2e46bfc57..a4e608720 100644 --- a/src/documents/models.py +++ b/src/documents/models.py @@ -766,6 +766,7 @@ 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, @@ -1674,6 +1675,18 @@ 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"), @@ -1912,6 +1925,33 @@ 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") diff --git a/src/documents/serialisers.py b/src/documents/serialisers.py index c9dc2845b..dcb84b527 100644 --- a/src/documents/serialisers.py +++ b/src/documents/serialisers.py @@ -3235,6 +3235,9 @@ class WorkflowActionSerializer(serializers.ModelSerializer[WorkflowAction]): "email", "webhook", "passwords", + "ai_suggestion_fields", + "ai_create_missing", + "ai_overwrite_existing", ] def validate(self, attrs): @@ -3292,6 +3295,23 @@ 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 @@ -3320,24 +3340,43 @@ 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: @@ -3345,6 +3384,14 @@ 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( diff --git a/src/documents/signals/handlers.py b/src/documents/signals/handlers.py index 87fa47643..d6359066b 100644 --- a/src/documents/signals/handlers.py +++ b/src/documents/signals/handlers.py @@ -984,6 +984,28 @@ 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 @@ -1039,6 +1061,7 @@ 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] = { @@ -1092,6 +1115,12 @@ 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 {} diff --git a/src/documents/tasks.py b/src/documents/tasks.py index c77d9fa30..ca53468eb 100644 --- a/src/documents/tasks.py +++ b/src/documents/tasks.py @@ -71,6 +71,7 @@ from paperless.config import RemoteOCRConfig from paperless.logging import consume_task_id from paperless.parsers import ParserContext from paperless.parsers.registry import get_parser_registry +from paperless_ai.exceptions import LLMTimeoutError from paperless_ai.indexing import llm_index_add_or_update_document from paperless_ai.indexing import llm_index_remove_document from paperless_ai.indexing import update_llm_index @@ -714,6 +715,45 @@ def llmindex_index( ) +@shared_task( + bind=True, + autoretry_for=(LLMTimeoutError,), + 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) diff --git a/src/documents/tests/test_api_workflows.py b/src/documents/tests/test_api_workflows.py index 788207751..78563de3f 100644 --- a/src/documents/tests/test_api_workflows.py +++ b/src/documents/tests/test_api_workflows.py @@ -641,6 +641,145 @@ 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: diff --git a/src/documents/tests/test_task_signals.py b/src/documents/tests/test_task_signals.py index 83652b5df..5aded2d17 100644 --- a/src/documents/tests/test_task_signals.py +++ b/src/documents/tests/test_task_signals.py @@ -385,6 +385,25 @@ 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: diff --git a/src/documents/tests/test_tasks.py b/src/documents/tests/test_tasks.py index 2954d6442..107c1dc6b 100644 --- a/src/documents/tests/test_tasks.py +++ b/src/documents/tests/test_tasks.py @@ -14,6 +14,7 @@ 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 @@ -447,3 +448,110 @@ 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)) diff --git a/src/documents/tests/test_workflows.py b/src/documents/tests/test_workflows.py index 5d4ba6598..e8fd5c41e 100644 --- a/src/documents/tests/test_workflows.py +++ b/src/documents/tests/test_workflows.py @@ -31,7 +31,10 @@ 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.base_model import ClassificationSuggestions +from paperless_ai.exceptions import LLMTimeoutError if TYPE_CHECKING: from django.db.models import QuerySet @@ -5488,3 +5491,481 @@ class TestRemoteOCRWorkflowAction(DirectoriesMixin, SampleDirMixin, APITestCase) ) self.assertIn("only applies to consumption triggers", "".join(cm.output)) + + +SUGGESTIONS: ClassificationSuggestions = { + "title": "Suggested Title", + "tags": { + "existing_ids": [], + "new_names": ["Existing Tag", "Suggested Tag"], + }, + "correspondents": { + "existing_ids": [], + "new_names": ["Existing Correspondent", "Suggested Correspondent"], + }, + "document_types": { + "existing_ids": [], + "new_names": ["Suggested Document Type"], + }, + "storage_paths": { + "existing_ids": [], + "new_names": ["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, + suggestions: ClassificationSuggestions = SUGGESTIONS, + ) -> 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_existing_id_suggestions_are_applied(self) -> None: + """ + GIVEN: + - AI suggestions that select existing taxonomy candidates by ID + WHEN: + - The suggestions are applied + THEN: + - Each selected object is assigned to the document + """ + tag = Tag.objects.create(name="Existing Tag", owner=self.user) + correspondent = Correspondent.objects.create( + name="Existing Correspondent", + owner=self.user, + ) + document_type = DocumentType.objects.create( + name="Existing Document Type", + owner=self.user, + ) + storage_path = StoragePath.objects.create( + name="Existing Storage Path", + path="{{ title }}", + owner=self.user, + ) + action = self.make_action(ai_overwrite_existing=True) + suggestions: ClassificationSuggestions = { + **SUGGESTIONS, + "tags": {"existing_ids": [tag.pk], "new_names": []}, + "correspondents": { + "existing_ids": [correspondent.pk], + "new_names": [], + }, + "document_types": { + "existing_ids": [document_type.pk], + "new_names": [], + }, + "storage_paths": { + "existing_ids": [storage_path.pk], + "new_names": [], + }, + } + + changed = self.apply(action, suggestions) + + self.assertEqual(list(self.doc.tags.all()), [tag]) + self.assertEqual(self.doc.correspondent, correspondent) + self.assertEqual(self.doc.document_type, document_type) + self.assertEqual(self.doc.storage_path, storage_path) + self.assertTrue( + {"tags", "correspondent", "document_type", "storage_path"} <= set(changed), + ) + + 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)) diff --git a/src/documents/views.py b/src/documents/views.py index 36a976624..bc77dc5eb 100644 --- a/src/documents/views.py +++ b/src/documents/views.py @@ -247,6 +247,7 @@ 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 @@ -665,20 +666,6 @@ 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, @@ -1542,7 +1529,10 @@ 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, request=request) + output_language = get_llm_output_language( + ai_config=ai_config, + user=request.user, + ) llm_cache_backend = ":".join( part for part in ( @@ -2331,7 +2321,10 @@ class ChatStreamingView(GenericAPIView[Any]): id__in=permitted_document_ids(request.user), ) - output_language = _get_llm_output_language(ai_config=ai_config, request=request) + output_language = get_llm_output_language( + ai_config=ai_config, + user=request.user, + ) response = StreamingHttpResponse( stream_chat_with_documents( diff --git a/src/documents/workflows/ai.py b/src/documents/workflows/ai.py new file mode 100644 index 000000000..386838782 --- /dev/null +++ b/src/documents/workflows/ai.py @@ -0,0 +1,259 @@ +import logging +from datetime import date +from datetime import datetime +from typing import TypeVar + +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 MatchingModel +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 +from paperless_ai.matching import resolve_correspondent_ids +from paperless_ai.matching import resolve_document_type_ids +from paperless_ai.matching import resolve_storage_path_ids +from paperless_ai.matching import resolve_tag_ids + +logger = logging.getLogger("paperless.workflows.ai") + +AISuggestionField = WorkflowAction.AISuggestionField +ObjT = TypeVar("ObjT", bound=MatchingModel) + +# 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: type[ObjT], + names: list[str], + matched: list[ObjT], + *, + create_missing: bool, + owner: User | None, +) -> ObjT | 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["title"].strip() + if title: + # title is capped at 128 characters + document.title = title[:128] + updated_fields.append("title") + + if should_set(AISuggestionField.CORRESPONDENT): + choice = suggestions["correspondents"] + names = choice["new_names"] + correspondent = resolve_object( + Correspondent, + names, + resolve_correspondent_ids(choice["existing_ids"], owner) + + 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): + choice = suggestions["document_types"] + names = choice["new_names"] + document_type = resolve_object( + DocumentType, + names, + resolve_document_type_ids(choice["existing_ids"], owner) + + 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): + choice = suggestions["storage_paths"] + names = choice["new_names"] + storage_path = resolve_object( + StoragePath, + names, + resolve_storage_path_ids(choice["existing_ids"], owner) + + 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["dates"]) + if created: + document.created = created + updated_fields.append("created") + + if AISuggestionField.TAGS in selected: + choice = suggestions["tags"] + names = choice["new_names"] + tags = resolve_tags( + names, + resolve_tag_ids(choice["existing_ids"], owner) + + 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") + + if updated_fields: + # save fields and update modified (excluding m2m tags from update_fields) + direct_updated_fields = [ + field for field in updated_fields if field in DIRECT_FIELDS.values() + ] + document.save(update_fields=[*direct_updated_fields, "modified"]) + + logger.info( + "Applied AI suggestions %s to document %s", + updated_fields or "(none)", + document.pk, + extra={"group": logging_group}, + ) + + return updated_fields diff --git a/src/paperless_ai/ai_classifier.py b/src/paperless_ai/ai_classifier.py index 14735e419..528300fc9 100644 --- a/src/paperless_ai/ai_classifier.py +++ b/src/paperless_ai/ai_classifier.py @@ -47,6 +47,22 @@ 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, diff --git a/src/paperless_ai/matching.py b/src/paperless_ai/matching.py index 0cadaf36c..b2e0ec955 100644 --- a/src/paperless_ai/matching.py +++ b/src/paperless_ai/matching.py @@ -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 get_objects_for_user_owner_aware +from documents.permissions import permitted_object_ids from documents.permissions import restrict_queryset_to_visible MATCH_THRESHOLD = 0.8 @@ -63,30 +63,40 @@ def resolve_storage_path_ids(ids: list[int], user: User | None) -> list[StorageP def _match_by_name( names: list[str], - user: User, + user: User | None, model: type[ModelT], perm: str, ) -> list[ModelT]: - queryset = get_objects_for_user_owner_aware(user, [perm], model) + # 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), + ) return _match_names_to_queryset(names, queryset) -def match_tags_by_name(names: list[str], user: User) -> list[Tag]: +def match_tags_by_name(names: list[str], user: User | None) -> list[Tag]: return _match_by_name(names, user, Tag, "view_tag") def match_correspondents_by_name( names: list[str], - user: User, + user: User | None, ) -> list[Correspondent]: return _match_by_name(names, user, Correspondent, "view_correspondent") -def match_document_types_by_name(names: list[str], user: User) -> list[DocumentType]: +def match_document_types_by_name( + names: list[str], + user: User | None, +) -> list[DocumentType]: return _match_by_name(names, user, DocumentType, "view_documenttype") -def match_storage_paths_by_name(names: list[str], user: User) -> list[StoragePath]: +def match_storage_paths_by_name( + names: list[str], + user: User | None, +) -> list[StoragePath]: return _match_by_name(names, user, StoragePath, "view_storagepath") diff --git a/src/paperless_ai/tests/test_matching.py b/src/paperless_ai/tests/test_matching.py index 4dd974fe6..bbf972ad9 100644 --- a/src/paperless_ai/tests/test_matching.py +++ b/src/paperless_ai/tests/test_matching.py @@ -1,5 +1,4 @@ from collections.abc import Callable -from unittest.mock import patch import pytest import pytest_mock @@ -45,33 +44,25 @@ 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") - @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() + def test_match_tags_by_name(self) -> None: 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") - @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() + def test_match_correspondents_by_name(self) -> None: 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") - @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() + def test_match_document_types_by_name(self) -> None: 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") - @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() + def test_match_storage_paths_by_name(self) -> None: names = ["Test Storage Path 1", "Nonexistent Storage Path"] result = match_storage_paths_by_name(names, user=None) self.assertEqual(len(result), 1) @@ -83,16 +74,12 @@ class TestAIMatching(TestCase): unmatched_names = extract_unmatched_names(llm_names, matched_objects) self.assertEqual(unmatched_names, ["Nonexistent Tag"]) - @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() + def test_match_tags_by_name_with_empty_names(self) -> None: names = [None, "", " "] result = match_tags_by_name(names, user=None) self.assertEqual(result, []) - @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() + def test_match_tags_with_fuzzy_matching(self) -> None: names = ["Test Taag 1", "Teest Tag 2"] result = match_tags_by_name(names, user=None) self.assertEqual(len(result), 2)