Compare commits

...
Author SHA1 Message Date
Trenton HandGitHub 6a02b87dde Feature: Updates remote OCR parser to respect the OCR mode setting (#13408)
* Have the remote parser respect the provided produce_archive_file setting, as already determined via the consumer checks

* Updates the documentation to be correct about the respecting now

* merge conflict fixing
2026-08-11 19:23:46 +00:00
Trenton HandGitHub 59a2651804 Fix: pass document chat queries as a QuerySet instead of a materialized list (#13638)
In tracemalloc based profiling, not materializing the whole Document list
reduced memory to approximately 20% of the baseline, with a peak memory
that scaled with the library size.  Now, the lazt queryset is used and only
the needed pk value is actually contributing to memory
2026-08-11 15:25:08 +00:00
GitHub Actions a99f63e059 Auto translate strings 2026-08-11 14:11:20 +00:00
shamoonandGitHub 939cb52f6e Fix: add pagination to saved views management page (#13646) 2026-08-11 07:08:30 -07:00
shamoonandGitHub 855669ddf9 Fix: fixes for workflow assign custom field values (#13630) 2026-08-10 07:38:07 -07:00
18 changed files with 428 additions and 89 deletions
+5 -4
View File
@@ -948,10 +948,11 @@ for display in the web interface.
!!! note !!! note
The **remote OCR parser** (Azure AI) always produces a searchable The **remote OCR parser** (Azure AI) also honors this setting: when
PDF and stores it as the archive copy, regardless of this setting. no archive is requested (`never`, or `auto` with a born-digital PDF),
`ARCHIVE_FILE_GENERATION=never` has no effect when the remote the remote engine is skipped entirely and locally-extracted text is
parser handles a document. used instead, avoiding an unnecessary API call and a duplicate text
layer.
#### [`PAPERLESS_OCR_CLEAN=<mode>`](#PAPERLESS_OCR_CLEAN) {#PAPERLESS_OCR_CLEAN} #### [`PAPERLESS_OCR_CLEAN=<mode>`](#PAPERLESS_OCR_CLEAN) {#PAPERLESS_OCR_CLEAN}
+5 -4
View File
@@ -187,10 +187,11 @@ PAPERLESS_ARCHIVE_FILE_GENERATION=auto
### Remote OCR parser ### Remote OCR parser
If you use the **remote OCR parser** (Azure AI), note that it always produces a If you use the **remote OCR parser** (Azure AI), `ARCHIVE_FILE_GENERATION` is
searchable PDF and stores it as the archive copy. `ARCHIVE_FILE_GENERATION=never` honored the same way as for the local engine: when no archive is requested
has no effect for documents handled by the remote parser - the archive is produced (`never`, or `auto` with a born-digital PDF), the remote engine is skipped
unconditionally by the remote engine. entirely and locally-extracted text is used instead, avoiding an unnecessary
API call and a duplicate text layer.
## Search Index (Whoosh -> Tantivy) ## Search Index (Whoosh -> Tantivy)
+3 -1
View File
@@ -576,7 +576,9 @@ The following workflow action types are available:
- Tags, correspondent, document type and storage path - Tags, correspondent, document type and storage path
- Document owner - Document owner
- View and / or edit permissions to users or groups - View and / or edit permissions to users or groups
- Custom fields. Note that no value for the field will be set - Custom fields, optionally with a value. If no value is set, the field is only added to the
document and any value it may already have is left untouched. If a value is set, it will
overwrite an existing value of that field on the document.
##### Removal {#workflow-action-removal} ##### Removal {#workflow-action-removal}
+8 -8
View File
@@ -599,7 +599,7 @@
</context-group> </context-group>
<context-group purpose="location"> <context-group purpose="location">
<context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.html</context> <context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.html</context>
<context context-type="linenumber">84,85</context> <context context-type="linenumber">85,86</context>
</context-group> </context-group>
</trans-unit> </trans-unit>
<trans-unit id="3768927257183755959" datatype="html"> <trans-unit id="3768927257183755959" datatype="html">
@@ -670,7 +670,7 @@
</context-group> </context-group>
<context-group purpose="location"> <context-group purpose="location">
<context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.html</context> <context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.html</context>
<context context-type="linenumber">85,86</context> <context context-type="linenumber">86,87</context>
</context-group> </context-group>
</trans-unit> </trans-unit>
<trans-unit id="5079885666748292382" datatype="html"> <trans-unit id="5079885666748292382" datatype="html">
@@ -9907,7 +9907,7 @@
</context-group> </context-group>
<context-group purpose="location"> <context-group purpose="location">
<context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context> <context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context>
<context context-type="linenumber">296</context> <context context-type="linenumber">314</context>
</context-group> </context-group>
</trans-unit> </trans-unit>
<trans-unit id="2620006875434695386" datatype="html"> <trans-unit id="2620006875434695386" datatype="html">
@@ -10234,7 +10234,7 @@
</context-group> </context-group>
<context-group purpose="location"> <context-group purpose="location">
<context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context> <context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context>
<context context-type="linenumber">290</context> <context context-type="linenumber">308</context>
</context-group> </context-group>
</trans-unit> </trans-unit>
<trans-unit id="3501895737484542570" datatype="html"> <trans-unit id="3501895737484542570" datatype="html">
@@ -10346,28 +10346,28 @@
<source>Saved view &quot;<x id="PH" equiv-text="savedView.name"/>&quot; deleted.</source> <source>Saved view &quot;<x id="PH" equiv-text="savedView.name"/>&quot; deleted.</source>
<context-group purpose="location"> <context-group purpose="location">
<context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context> <context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context>
<context context-type="linenumber">160</context> <context context-type="linenumber">178</context>
</context-group> </context-group>
</trans-unit> </trans-unit>
<trans-unit id="1660419335376265526" datatype="html"> <trans-unit id="1660419335376265526" datatype="html">
<source>Views saved successfully.</source> <source>Views saved successfully.</source>
<context-group purpose="location"> <context-group purpose="location">
<context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context> <context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context>
<context context-type="linenumber">237</context> <context context-type="linenumber">255</context>
</context-group> </context-group>
</trans-unit> </trans-unit>
<trans-unit id="1699877326523238632" datatype="html"> <trans-unit id="1699877326523238632" datatype="html">
<source>Error while saving views.</source> <source>Error while saving views.</source>
<context-group purpose="location"> <context-group purpose="location">
<context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context> <context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context>
<context context-type="linenumber">242</context> <context context-type="linenumber">260</context>
</context-group> </context-group>
</trans-unit> </trans-unit>
<trans-unit id="4919025779187821586" datatype="html"> <trans-unit id="4919025779187821586" datatype="html">
<source>Note: Sharing saved views does not share the underlying documents.</source> <source>Note: Sharing saved views does not share the underlying documents.</source>
<context-group purpose="location"> <context-group purpose="location">
<context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context> <context context-type="sourcefile">src/app/components/manage/saved-views/saved-views.component.ts</context>
<context context-type="linenumber">278</context> <context context-type="linenumber">296</context>
</context-group> </context-group>
</trans-unit> </trans-unit>
<trans-unit id="1229748338333965418" datatype="html"> <trans-unit id="1229748338333965418" datatype="html">
@@ -52,10 +52,10 @@ describe('CustomFieldsValuesComponent', () => {
}) })
it('should set selectedFields and map values correctly', () => { it('should set selectedFields and map values correctly', () => {
component.value = { 1: 'value1' } component.value = { 1: 'value1', 3: 0, 4: false }
component.selectedFields = [1, 2] component.selectedFields = [1, 2, 3, 4]
expect(component.selectedFields).toEqual([1, 2]) expect(component.selectedFields).toEqual([1, 2, 3, 4])
expect(component.value).toEqual({ 1: 'value1', 2: null }) expect(component.value).toEqual({ 1: 'value1', 2: null, 3: 0, 4: false })
}) })
it('should return the correct custom field by id', () => { it('should return the correct custom field by id', () => {
@@ -77,7 +77,7 @@ export class CustomFieldsValuesComponent extends AbstractInputComponent<Object>
this._selectedFields = newFields this._selectedFields = newFields
// map the selected fields to an object with field_id as key and value as value // map the selected fields to an object with field_id as key and value as value
this.value = newFields.reduce((acc, fieldId) => { this.value = newFields.reduce((acc, fieldId) => {
acc[fieldId] = this.value?.[fieldId] || null acc[fieldId] = this.value?.[fieldId] ?? null
return acc return acc
}, {}) }, {})
this.onChange(this.value) this.onChange(this.value)
@@ -7,7 +7,7 @@
</pngx-page-header> </pngx-page-header>
<form [formGroup]="savedViewsForm" (ngSubmit)="save()"> <form [formGroup]="savedViewsForm" (ngSubmit)="save()">
<ul class="list-group mb-3" formGroupName="savedViews"> <ul class="list-group mb-3" formGroupName="savedViews">
@for (view of savedViews(); track view) { @for (view of pagedSavedViews(); track view) {
<li class="list-group-item py-3"> <li class="list-group-item py-3">
<div [formGroupName]="view.id"> <div [formGroupName]="view.id">
<div class="row"> <div class="row">
@@ -81,6 +81,11 @@
} }
</ul> </ul>
<button type="button" (click)="reset()" class="btn btn-outline-secondary mb-2" [disabled]="(isDirty$ | async) === false" i18n>Cancel</button> <div class="d-flex align-items-center mb-3">
<button type="submit" class="btn btn-primary ms-2 mb-2" [disabled]="(isDirty$ | async) === false" i18n>Save</button> <button type="button" (click)="reset()" class="btn btn-outline-secondary mb-2" [disabled]="(isDirty$ | async) === false" i18n>Cancel</button>
<button type="submit" class="btn btn-primary ms-2 mb-2" [disabled]="(isDirty$ | async) === false" i18n>Save</button>
@if (savedViews()?.length > pageSize) {
<ngb-pagination class="ms-auto" [pageSize]="pageSize" [collectionSize]="savedViews().length" [page]="page()" [maxSize]="5" (pageChange)="page.set($event)" size="sm" aria-label="Pagination"></ngb-pagination>
}
</div>
</form> </form>
@@ -4,6 +4,7 @@ import { provideHttpClientTesting } from '@angular/common/http/testing'
import { signal } from '@angular/core' import { signal } from '@angular/core'
import { ComponentFixture, TestBed } from '@angular/core/testing' import { ComponentFixture, TestBed } from '@angular/core/testing'
import { FormsModule, ReactiveFormsModule } from '@angular/forms' import { FormsModule, ReactiveFormsModule } from '@angular/forms'
import { By } from '@angular/platform-browser'
import { NgbModal, NgbModule } from '@ng-bootstrap/ng-bootstrap' import { NgbModal, NgbModule } from '@ng-bootstrap/ng-bootstrap'
import { NgxBootstrapIconsModule, allIcons } from 'ngx-bootstrap-icons' import { NgxBootstrapIconsModule, allIcons } from 'ngx-bootstrap-icons'
import { Subject, of, throwError } from 'rxjs' import { Subject, of, throwError } from 'rxjs'
@@ -222,6 +223,44 @@ describe('SavedViewsComponent', () => {
).toEqual(view.show_on_dashboard) ).toEqual(view.show_on_dashboard)
}) })
it('should page saved views, clamp the page if views are removed', () => {
const manyViews = Array.from({ length: 30 }, (_, i) => ({
id: i + 1,
name: `view${i + 1}`,
})) as SavedView[]
const listSpy = jest.spyOn(savedViewService, 'list').mockReturnValue(
of({
all: manyViews.map((v) => v.id),
count: manyViews.length,
results: manyViews.concat([]),
})
)
component.ngOnInit()
fixture.detectChanges()
expect(listSpy).toHaveBeenCalledWith(1, 100000, null, false, {
full_perms: true,
})
expect(component.pagedSavedViews()).toHaveLength(25)
expect(fixture.debugElement.query(By.css('ngb-pagination'))).not.toBeNull()
// all views have controls, not just the current page
expect(
Object.keys(component.savedViewsForm.get('savedViews').value)
).toHaveLength(30)
component.page.set(2)
expect(component.pagedSavedViews()).toHaveLength(5)
listSpy.mockReturnValue(
of({
all: manyViews.slice(0, 25).map((v) => v.id),
count: 25,
results: manyViews.slice(0, 25),
})
)
component.ngOnInit()
expect(component.page()).toEqual(1)
})
it('should support editing permissions', () => { it('should support editing permissions', () => {
const confirmClicked = new Subject<any>() const confirmClicked = new Subject<any>()
const modalRef = { const modalRef = {
@@ -1,12 +1,19 @@
import { AsyncPipe } from '@angular/common' import { AsyncPipe } from '@angular/common'
import { Component, OnDestroy, OnInit, inject, signal } from '@angular/core' import {
Component,
OnDestroy,
OnInit,
computed,
inject,
signal,
} from '@angular/core'
import { import {
FormControl, FormControl,
FormGroup, FormGroup,
FormsModule, FormsModule,
ReactiveFormsModule, ReactiveFormsModule,
} from '@angular/forms' } from '@angular/forms'
import { NgbModal } from '@ng-bootstrap/ng-bootstrap' import { NgbModal, NgbPaginationModule } from '@ng-bootstrap/ng-bootstrap'
import { dirtyCheck } from '@ngneat/dirty-check-forms' import { dirtyCheck } from '@ngneat/dirty-check-forms'
import { NgxBootstrapIconsModule } from 'ngx-bootstrap-icons' import { NgxBootstrapIconsModule } from 'ngx-bootstrap-icons'
import { BehaviorSubject, Observable, of, switchMap, takeUntil } from 'rxjs' import { BehaviorSubject, Observable, of, switchMap, takeUntil } from 'rxjs'
@@ -42,6 +49,7 @@ import { LoadingComponentWithPermissions } from '../../loading-component/loading
FormsModule, FormsModule,
ReactiveFormsModule, ReactiveFormsModule,
AsyncPipe, AsyncPipe,
NgbPaginationModule,
NgxBootstrapIconsModule, NgxBootstrapIconsModule,
], ],
}) })
@@ -58,6 +66,14 @@ export class SavedViewsComponent
DisplayMode = DisplayMode DisplayMode = DisplayMode
readonly savedViews = signal<SavedView[]>(undefined) readonly savedViews = signal<SavedView[]>(undefined)
readonly page = signal(1)
public readonly pageSize = 25
// All views are loaded at init, so paging is only for display
readonly pagedSavedViews = computed(() => {
const start = (this.page() - 1) * this.pageSize
return this.savedViews()?.slice(start, start + this.pageSize)
})
private savedViewsGroup = new FormGroup({}) private savedViewsGroup = new FormGroup({})
public savedViewsForm: FormGroup = new FormGroup({ public savedViewsForm: FormGroup = new FormGroup({
savedViews: this.savedViewsGroup, savedViews: this.savedViewsGroup,
@@ -84,9 +100,11 @@ export class SavedViewsComponent
private reloadViews(): void { private reloadViews(): void {
this.loading.set(true) this.loading.set(true)
this.savedViewService this.savedViewService
.list(null, null, null, false, { full_perms: true }) .list(1, 100000, null, false, { full_perms: true })
.subscribe((r) => { .subscribe((r) => {
this.savedViews.set(r.results) this.savedViews.set(r.results)
const pageCount = Math.ceil(r.results.length / this.pageSize)
this.page.update((page) => Math.min(page, Math.max(1, pageCount)))
this.initialize() this.initialize()
}) })
} }
+7
View File
@@ -3213,6 +3213,13 @@ class WorkflowActionSerializer(serializers.ModelSerializer[WorkflowAction]):
{"assign_title": f'Invalid f-string detected: "{e.args[0]}"'}, {"assign_title": f'Invalid f-string detected: "{e.args[0]}"'},
) )
if attrs.get("assign_custom_fields_values"):
# Empty strings treated as None to avoid unexpected behavior
attrs["assign_custom_fields_values"] = {
field_id: (None if value == "" else value)
for field_id, value in attrs["assign_custom_fields_values"].items()
}
if ( if (
"type" in attrs "type" in attrs
and attrs["type"] == WorkflowAction.WorkflowActionType.EMAIL and attrs["type"] == WorkflowAction.WorkflowActionType.EMAIL
@@ -422,6 +422,11 @@ class TestApiWorkflows(DirectoriesMixin, APITestCase):
json.dumps( json.dumps(
{ {
"assign_title": "", "assign_title": "",
"assign_custom_fields": [self.cf1.id, self.cf2.id],
"assign_custom_fields_values": {
str(self.cf1.id): "",
str(self.cf2.id): 0,
},
}, },
), ),
content_type="application/json", content_type="application/json",
@@ -429,6 +434,10 @@ class TestApiWorkflows(DirectoriesMixin, APITestCase):
self.assertEqual(response.status_code, status.HTTP_201_CREATED) self.assertEqual(response.status_code, status.HTTP_201_CREATED)
action = WorkflowAction.objects.get(id=response.data["id"]) action = WorkflowAction.objects.get(id=response.data["id"])
self.assertIsNone(action.assign_title) self.assertIsNone(action.assign_title)
self.assertEqual(
action.assign_custom_fields_values,
{str(self.cf1.id): None, str(self.cf2.id): 0},
)
response = self.client.post( response = self.client.post(
self.ENDPOINT_TRIGGERS, self.ENDPOINT_TRIGGERS,
+49
View File
@@ -2000,6 +2000,55 @@ class TestWorkflows(
r"Doc added in \w{3,}", r"Doc added in \w{3,}",
) # Match any 3-letter month name ) # Match any 3-letter month name
def test_document_updated_workflow_existing_custom_field_empty_value(self) -> None:
"""
GIVEN:
- Existing workflow with UPDATED trigger and action that assigns a custom field
with an empty value
WHEN:
- Document is updated that already contains the field with a value
THEN:
- The existing value is left untouched, see GH #13627
"""
trigger = WorkflowTrigger.objects.create(
type=WorkflowTrigger.WorkflowTriggerType.DOCUMENT_UPDATED,
filter_has_document_type=self.dt,
)
action = WorkflowAction.objects.create()
action.assign_custom_fields.add(self.cf1)
action.assign_custom_fields_values = {self.cf1.pk: ""}
action.save()
w = Workflow.objects.create(
name="Workflow 1",
order=0,
)
w.triggers.add(trigger)
w.actions.add(action)
w.save()
doc = Document.objects.create(
title="sample test",
correspondent=self.c,
original_filename="sample.pdf",
)
CustomFieldInstance.objects.create(
document=doc,
field=self.cf1,
value_text="existing value",
)
superuser = User.objects.create_superuser("superuser")
self.client.force_authenticate(user=superuser)
self.client.patch(
f"/api/documents/{doc.id}/",
{"document_type": self.dt.id},
format="json",
)
doc.refresh_from_db()
self.assertEqual(doc.custom_fields.get(field=self.cf1).value, "existing value")
def test_document_updated_workflow_existing_custom_field(self) -> None: def test_document_updated_workflow_existing_custom_field(self) -> None:
""" """
GIVEN: GIVEN:
+1 -1
View File
@@ -2267,7 +2267,7 @@ class ChatStreamingView(GenericAPIView[Any]):
if not has_perms_owner_aware(request.user, "view_document", document): if not has_perms_owner_aware(request.user, "view_document", document):
return HttpResponseForbidden("Insufficient permissions") return HttpResponseForbidden("Insufficient permissions")
documents = [document] documents = Document.objects.filter(pk=document.pk)
else: else:
documents = Document.objects.filter( documents = Document.objects.filter(
id__in=permitted_document_ids(request.user), id__in=permitted_document_ids(request.user),
+2 -1
View File
@@ -105,7 +105,8 @@ def apply_assignment_to_document(
field=field, field=field,
document=document, document=document,
).first() ).first()
if instance and args[value_field_name] is not None: # empty string is indistinguishable from no value in the UI
if instance and args[value_field_name] not in (None, ""):
setattr(instance, value_field_name, args[value_field_name]) setattr(instance, value_field_name, args[value_field_name])
instance.save() instance.save()
elif not instance: elif not instance:
+33 -7
View File
@@ -3,7 +3,9 @@ Built-in remote-OCR document parser.
Handles documents by sending them to a configured remote OCR engine Handles documents by sending them to a configured remote OCR engine
(currently Azure AI Vision / Document Intelligence) and retrieving both (currently Azure AI Vision / Document Intelligence) and retrieving both
the extracted text and a searchable PDF with an embedded text layer. the extracted text and a searchable PDF with an embedded text layer. For
born-digital PDFs that need no archive copy, the remote call is skipped
entirely in favor of locally-extracted text (see ``RemoteDocumentParser.parse``).
When no engine is configured, ``score()`` returns ``None`` so the parser When no engine is configured, ``score()`` returns ``None`` so the parser
is effectively invisible to the registry the tesseract parser handles is effectively invisible to the registry the tesseract parser handles
@@ -22,6 +24,8 @@ from typing import Self
from django.conf import settings from django.conf import settings
from documents.parsers import ParseError from documents.parsers import ParseError
from paperless.parsers.utils import extract_pdf_text
from paperless.parsers.utils import post_process_text
from paperless.version import __full_version_str__ from paperless.version import __full_version_str__
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -70,8 +74,11 @@ class RemoteDocumentParser:
"""Parse documents via a remote OCR API (currently Azure AI Vision). """Parse documents via a remote OCR API (currently Azure AI Vision).
This parser sends documents to a remote engine that returns both This parser sends documents to a remote engine that returns both
extracted text and a searchable PDF with an embedded text layer. extracted text and a searchable PDF with an embedded text layer,
It does not depend on Tesseract or ocrmypdf. except when ``parse()`` is called with ``produce_archive=False`` for
a PDF, in which case the remote call is skipped and only locally
extracted text is returned (no archive). It does not depend on
Tesseract or ocrmypdf.
Class attributes Class attributes
---------------- ----------------
@@ -160,8 +167,11 @@ class RemoteDocumentParser:
Returns Returns
------- -------
bool bool
Always True the remote engine always returns a PDF with an Always True the remote engine is capable of returning a PDF
embedded text layer that serves as the archive copy. with an embedded text layer to serve as the archive copy.
Whether it actually does so for a given document depends on
``produce_archive`` passed to :meth:`parse` (see there for when
the remote engine call, and thus archive generation, is skipped).
""" """
return True return True
@@ -218,6 +228,12 @@ class RemoteDocumentParser:
) -> None: ) -> None:
"""Send the document to the remote engine and store results. """Send the document to the remote engine and store results.
When *produce_archive* is False for a PDF, the caller (via
``documents.consumer.should_produce_archive``) has already determined
that the document is born-digital and needs no archive skip the
remote engine entirely rather than re-OCRing it and creating a
duplicate text layer.
Parameters Parameters
---------- ----------
document_path: document_path:
@@ -225,8 +241,8 @@ class RemoteDocumentParser:
mime_type: mime_type:
Detected MIME type of the document. Detected MIME type of the document.
produce_archive: produce_archive:
Ignored the remote engine always returns a searchable PDF, Whether an archive copy is wanted. For PDFs, False skips the
which is stored as the archive copy regardless of this flag. remote engine and uses locally-extracted text instead.
""" """
config = RemoteEngineConfig( config = RemoteEngineConfig(
engine=settings.REMOTE_OCR_ENGINE, engine=settings.REMOTE_OCR_ENGINE,
@@ -241,6 +257,16 @@ class RemoteDocumentParser:
self._text = "" self._text = ""
return return
if not produce_archive and mime_type == "application/pdf":
logger.debug(
"Remote OCR: skipped — no archive requested, "
"using locally-extracted text",
)
self._text = (
post_process_text(extract_pdf_text(document_path, log=logger)) or ""
)
return
if config.engine == "azureai": if config.engine == "azureai":
self._text = self._azure_ai_vision_parse(document_path, config) self._text = self._azure_ai_vision_parse(document_path, config)
@@ -337,6 +337,117 @@ class TestRemoteParserParse:
assert remote_parser.get_date() is None assert remote_parser.get_date() is None
# ---------------------------------------------------------------------------
# parse() — produce_archive=False skips the remote engine (PDFs only)
# ---------------------------------------------------------------------------
class TestRemoteParserSkipsWhenNoArchiveWanted:
"""When the caller has already decided no archive is needed for a PDF
(documents.consumer.should_produce_archive), the remote engine call is
skipped entirely in favor of locally-extracted text.
"""
def test_pdf_skips_azure_when_no_archive_requested(
self,
remote_parser: RemoteDocumentParser,
simple_digital_pdf_file: Path,
azure_client: Mock,
) -> None:
"""
GIVEN: produce_archive=False for a PDF
WHEN: parse() is called
THEN: Azure is never invoked, no archive is produced, and text
comes from local pdftotext extraction
"""
remote_parser.parse(
simple_digital_pdf_file,
"application/pdf",
produce_archive=False,
)
azure_client.begin_analyze_document.assert_not_called()
assert remote_parser.get_archive_path() is None
assert remote_parser.get_text() != ""
def test_pdf_no_archive_requested_text_matches_local_extraction(
self,
remote_parser: RemoteDocumentParser,
simple_digital_pdf_file: Path,
azure_client: Mock,
mocker: MockerFixture,
) -> None:
"""
GIVEN: produce_archive=False for a PDF
WHEN: parse() is called
THEN: the returned text is exactly the locally-extracted text,
not anything from the (unused) Azure mock
"""
mocker.patch(
"paperless.parsers.remote.extract_pdf_text",
return_value="Local digital text.",
)
remote_parser.parse(
simple_digital_pdf_file,
"application/pdf",
produce_archive=False,
)
assert remote_parser.get_text() == "Local digital text."
def test_pdf_no_archive_requested_closes_no_client(
self,
remote_parser: RemoteDocumentParser,
simple_digital_pdf_file: Path,
azure_client: Mock,
) -> None:
remote_parser.parse(
simple_digital_pdf_file,
"application/pdf",
produce_archive=False,
)
azure_client.close.assert_not_called()
def test_non_pdf_still_calls_azure_when_no_archive_requested(
self,
remote_parser: RemoteDocumentParser,
simple_digital_pdf_file: Path,
azure_client: Mock,
) -> None:
"""
Images have no local-text fallback, so produce_archive=False does
not skip the remote engine for non-PDF MIME types.
"""
remote_parser.parse(
simple_digital_pdf_file,
"image/png",
produce_archive=False,
)
azure_client.begin_analyze_document.assert_called_once()
assert remote_parser.get_text() == _DEFAULT_TEXT
@pytest.mark.usefixtures("no_engine_settings")
def test_unconfigured_engine_takes_precedence_over_skip(
self,
remote_parser: RemoteDocumentParser,
simple_digital_pdf_file: Path,
) -> None:
"""An unconfigured engine still short-circuits before the
produce_archive check, returning empty text as before.
"""
remote_parser.parse(
simple_digital_pdf_file,
"application/pdf",
produce_archive=False,
)
assert remote_parser.get_text() == ""
assert remote_parser.get_archive_path() is None
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# parse() — Azure failure path # parse() — Azure failure path
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
+21 -6
View File
@@ -2,6 +2,8 @@ import json
import logging import logging
import sys import sys
from django.db.models import QuerySet
from documents.models import Document from documents.models import Document
from paperless.config import AIConfig from paperless.config import AIConfig
from paperless_ai.client import AIClient from paperless_ai.client import AIClient
@@ -82,10 +84,21 @@ def _build_document_reference(
def _get_document_references( def _get_document_references(
documents: list[Document], documents: QuerySet[Document],
top_nodes: list, top_nodes: list,
) -> list[dict[str, int | str]]: ) -> list[dict[str, int | str]]:
allowed_documents = {doc.pk: doc for doc in documents} candidate_ids: set[int] = set()
for node in top_nodes:
try:
candidate_ids.add(int(node.metadata["document_id"]))
except (KeyError, TypeError, ValueError): # pragma: no cover
continue
if not candidate_ids:
return []
allowed_documents = {doc.pk: doc for doc in documents.filter(pk__in=candidate_ids)}
references: list[dict[str, int | str]] = [] references: list[dict[str, int | str]] = []
seen_document_ids: set[int] = set() seen_document_ids: set[int] = set()
@@ -119,7 +132,7 @@ def _format_chat_metadata_trailer(references: list[dict[str, int | str]]) -> str
def stream_chat_with_documents( def stream_chat_with_documents(
query_str: str, query_str: str,
documents: list[Document], documents: QuerySet[Document],
output_language: str | None = None, output_language: str | None = None,
): ):
try: try:
@@ -135,10 +148,10 @@ def stream_chat_with_documents(
def _stream_chat_with_documents( def _stream_chat_with_documents(
query_str: str, query_str: str,
documents: list[Document], documents: QuerySet[Document],
output_language: str | None = None, output_language: str | None = None,
): ):
if not documents: if not documents.exists():
yield CHAT_NO_CONTENT_MESSAGE yield CHAT_NO_CONTENT_MESSAGE
return return
@@ -148,7 +161,9 @@ def _stream_chat_with_documents(
from llama_index.core.retrievers import VectorIndexRetriever from llama_index.core.retrievers import VectorIndexRetriever
config = AIConfig() config = AIConfig()
filters = _document_id_filters(str(doc.pk) for doc in documents) filters = _document_id_filters(
str(pk) for pk in documents.values_list("pk", flat=True)
)
# Hold the shared read lock for the whole operation: the query engine # Hold the shared read lock for the whole operation: the query engine
# retrieves from the vector store again during synthesis, so the connection # retrieves from the vector store again during synthesis, so the connection
+101 -46
View File
@@ -3,10 +3,12 @@ from unittest.mock import MagicMock
from unittest.mock import patch from unittest.mock import patch
import pytest import pytest
from django.db.models.signals import post_init
from llama_index.core import settings as llama_settings from llama_index.core import settings as llama_settings
from llama_index.core.embeddings.mock_embed_model import MockEmbedding from llama_index.core.embeddings.mock_embed_model import MockEmbedding
from llama_index.core.schema import TextNode from llama_index.core.schema import TextNode
from documents.models import Document
from documents.tests.factories import DocumentFactory from documents.tests.factories import DocumentFactory
from paperless_ai import chat from paperless_ai import chat
from paperless_ai import indexing from paperless_ai import indexing
@@ -36,16 +38,6 @@ def patch_embed_nodes():
yield mock_embed_nodes yield mock_embed_nodes
@pytest.fixture
def mock_document():
doc = MagicMock()
doc.pk = 1
doc.title = "Test Document"
doc.filename = "test_file.pdf"
doc.content = "This is the document content."
return doc
def assert_chat_output( def assert_chat_output(
output: list[str], output: list[str],
*, *,
@@ -61,6 +53,13 @@ def assert_chat_output(
} }
def _fake_documents_queryset(pks: list[int]) -> MagicMock:
qs = MagicMock()
qs.exists.return_value = bool(pks)
qs.values_list.return_value = pks
return qs
@pytest.mark.parametrize( @pytest.mark.parametrize(
("output_language", "expected_language_line"), ("output_language", "expected_language_line"),
[ [
@@ -107,9 +106,10 @@ def test_build_refine_prompt(
@pytest.mark.django_db @pytest.mark.django_db
def test_stream_chat_with_one_document_retrieval( def test_stream_chat_with_one_document_retrieval(
mock_document,
patch_embed_nodes, patch_embed_nodes,
) -> None: ) -> None:
document = DocumentFactory.create(title="Test Document", content="ignored")
documents = Document.objects.filter(pk=document.pk)
with ( with (
patch("paperless_ai.chat.AIClient") as mock_client_cls, patch("paperless_ai.chat.AIClient") as mock_client_cls,
patch("paperless_ai.chat.load_or_build_index") as mock_load_index, patch("paperless_ai.chat.load_or_build_index") as mock_load_index,
@@ -124,22 +124,19 @@ def test_stream_chat_with_one_document_retrieval(
mock_client_cls.return_value = mock_client mock_client_cls.return_value = mock_client
mock_client.llm = MagicMock() mock_client.llm = MagicMock()
mock_node = TextNode(
text="This is node content.",
metadata={"document_id": str(mock_document.pk), "title": "Test Document"},
)
mock_index = MagicMock() mock_index = MagicMock()
# Simulate get_nodes returning nodes (content exists) mock_index.vector_store.get_nodes.return_value = [
mock_index.vector_store.get_nodes.return_value = [mock_node] TextNode(
text="This is node content.",
metadata={"document_id": str(document.pk), "title": "Test Document"},
),
]
mock_load_index.return_value = mock_index mock_load_index.return_value = mock_index
mock_retriever_instance = MagicMock() mock_retriever_instance = MagicMock()
mock_retriever_instance.retrieve.return_value = [ mock_retriever_instance.retrieve.return_value = [
MagicMock( MagicMock(
metadata={ metadata={"document_id": str(document.pk), "title": "Test Document"},
"document_id": str(mock_document.pk),
"title": "Test Document",
},
), ),
] ]
@@ -153,7 +150,7 @@ def test_stream_chat_with_one_document_retrieval(
"llama_index.core.retrievers.VectorIndexRetriever", "llama_index.core.retrievers.VectorIndexRetriever",
return_value=mock_retriever_instance, return_value=mock_retriever_instance,
): ):
output = list(stream_chat_with_documents("What is this?", [mock_document])) output = list(stream_chat_with_documents("What is this?", documents))
mock_query_engine.query.assert_called_once_with("What is this?") mock_query_engine.query.assert_called_once_with("What is this?")
synthesizer_kwargs = mock_get_response_synthesizer.call_args.kwargs synthesizer_kwargs = mock_get_response_synthesizer.call_args.kwargs
@@ -166,13 +163,16 @@ def test_stream_chat_with_one_document_retrieval(
output, output,
expected_chunks=["chunk1", "chunk2"], expected_chunks=["chunk1", "chunk2"],
expected_references=[ expected_references=[
{"id": mock_document.pk, "title": "Test Document"}, {"id": document.pk, "title": "Test Document"},
], ],
) )
@pytest.mark.django_db @pytest.mark.django_db
def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> None: def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> None:
doc1 = DocumentFactory.create(title="Document 1", content="ignored")
doc2 = DocumentFactory.create(title="Document 2", content="ignored")
documents = Document.objects.filter(pk__in=[doc1.pk, doc2.pk])
with ( with (
patch("paperless_ai.chat.AIClient") as mock_client_cls, patch("paperless_ai.chat.AIClient") as mock_client_cls,
patch("paperless_ai.chat.load_or_build_index") as mock_load_index, patch("paperless_ai.chat.load_or_build_index") as mock_load_index,
@@ -184,23 +184,23 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non
mock_client_cls.return_value = mock_client mock_client_cls.return_value = mock_client
mock_client.llm = MagicMock() mock_client.llm = MagicMock()
mock_node1 = TextNode(
text="Content for doc 1.",
metadata={"document_id": "1", "title": "Document 1"},
)
mock_node2 = TextNode(
text="Content for doc 2.",
metadata={"document_id": "2", "title": "Document 2"},
)
mock_index = MagicMock() mock_index = MagicMock()
# Simulate get_nodes returning nodes (content exists) mock_index.vector_store.get_nodes.return_value = [
mock_index.vector_store.get_nodes.return_value = [mock_node1, mock_node2] TextNode(
text="Content for doc 1.",
metadata={"document_id": str(doc1.pk), "title": "Document 1"},
),
TextNode(
text="Content for doc 2.",
metadata={"document_id": str(doc2.pk), "title": "Document 2"},
),
]
mock_load_index.return_value = mock_index mock_load_index.return_value = mock_index
mock_retriever_instance = MagicMock() mock_retriever_instance = MagicMock()
mock_retriever_instance.retrieve.return_value = [ mock_retriever_instance.retrieve.return_value = [
MagicMock(metadata={"document_id": "1", "title": "Document 1"}), MagicMock(metadata={"document_id": str(doc1.pk), "title": "Document 1"}),
MagicMock(metadata={"document_id": "2", "title": "Document 2"}), MagicMock(metadata={"document_id": str(doc2.pk), "title": "Document 2"}),
] ]
mock_response_stream = MagicMock() mock_response_stream = MagicMock()
@@ -210,14 +210,11 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non
mock_query_engine_cls.return_value = mock_query_engine mock_query_engine_cls.return_value = mock_query_engine
mock_query_engine.query.return_value = mock_response_stream mock_query_engine.query.return_value = mock_response_stream
doc1 = MagicMock(pk=1, title="Document 1", filename="doc1.pdf")
doc2 = MagicMock(pk=2, title="Document 2", filename="doc2.pdf")
with patch( with patch(
"llama_index.core.retrievers.VectorIndexRetriever", "llama_index.core.retrievers.VectorIndexRetriever",
return_value=mock_retriever_instance, return_value=mock_retriever_instance,
): ):
output = list(stream_chat_with_documents("What's up?", [doc1, doc2])) output = list(stream_chat_with_documents("What's up?", documents))
mock_query_engine.query.assert_called_once_with("What's up?") mock_query_engine.query.assert_called_once_with("What's up?")
patch_embed_nodes.assert_not_called() patch_embed_nodes.assert_not_called()
@@ -225,15 +222,15 @@ def test_stream_chat_with_multiple_documents_retrieval(patch_embed_nodes) -> Non
output, output,
expected_chunks=["chunk1", "chunk2"], expected_chunks=["chunk1", "chunk2"],
expected_references=[ expected_references=[
{"id": 1, "title": "Document 1"}, {"id": doc1.pk, "title": "Document 1"},
{"id": 2, "title": "Document 2"}, {"id": doc2.pk, "title": "Document 2"},
], ],
) )
def test_stream_chat_empty_document_list() -> None: def test_stream_chat_empty_document_list() -> None:
with patch("paperless_ai.chat.load_or_build_index") as mock_load_index: with patch("paperless_ai.chat.load_or_build_index") as mock_load_index:
output = list(stream_chat_with_documents("Any info?", [])) output = list(stream_chat_with_documents("Any info?", Document.objects.none()))
mock_load_index.assert_not_called() mock_load_index.assert_not_called()
assert output == ["Sorry, I couldn't find any content to answer your question."] assert output == ["Sorry, I couldn't find any content to answer your question."]
@@ -253,7 +250,9 @@ def test_stream_chat_no_matching_nodes() -> None:
mock_index.vector_store.get_nodes.return_value = [] mock_index.vector_store.get_nodes.return_value = []
mock_load_index.return_value = mock_index mock_load_index.return_value = mock_index
output = list(stream_chat_with_documents("Any info?", [MagicMock(pk=1)])) output = list(
stream_chat_with_documents("Any info?", _fake_documents_queryset([1])),
)
assert output == ["Sorry, I couldn't find any content to answer your question."] assert output == ["Sorry, I couldn't find any content to answer your question."]
@@ -282,7 +281,9 @@ def test_stream_chat_unexpected_failure_returns_generic_error(caplog) -> None:
) )
mock_retriever_cls.return_value = mock_retriever mock_retriever_cls.return_value = mock_retriever
output = list(stream_chat_with_documents("Any info?", [MagicMock(pk=1)])) output = list(
stream_chat_with_documents("Any info?", _fake_documents_queryset([1])),
)
assert output == [CHAT_ERROR_MESSAGE] assert output == [CHAT_ERROR_MESSAGE]
assert "Failed to stream document chat response" in caplog.text assert "Failed to stream document chat response" in caplog.text
@@ -298,7 +299,12 @@ class TestStreamChatRetrieval:
) -> None: ) -> None:
doc = DocumentFactory.create(content="hello world") doc = DocumentFactory.create(content="hello world")
# Nothing indexed for this document yet. # Nothing indexed for this document yet.
out = list(chat.stream_chat_with_documents("question?", [doc])) out = list(
chat.stream_chat_with_documents(
"question?",
Document.objects.filter(pk=doc.pk),
),
)
assert chat.CHAT_NO_CONTENT_MESSAGE in out assert chat.CHAT_NO_CONTENT_MESSAGE in out
def test_chat_filter_contains_only_requested_document_ids( def test_chat_filter_contains_only_requested_document_ids(
@@ -332,7 +338,12 @@ class TestStreamChatRetrieval:
side_effect=capture_retriever, side_effect=capture_retriever,
) )
list(chat.stream_chat_with_documents("question?", [included])) list(
chat.stream_chat_with_documents(
"question?",
Document.objects.filter(pk=included.pk),
),
)
assert captured_filters, "VectorIndexRetriever was never constructed" assert captured_filters, "VectorIndexRetriever was never constructed"
filt = captured_filters[0] filt = captured_filters[0]
@@ -340,3 +351,47 @@ class TestStreamChatRetrieval:
filter_values = filt.filters[0].value filter_values = filt.filters[0].value
assert str(included.pk) in filter_values assert str(included.pk) in filter_values
assert str(excluded.pk) not in filter_values assert str(excluded.pk) not in filter_values
@pytest.mark.django_db
def test_get_document_references_only_queries_referenced_documents(
self,
django_assert_num_queries,
) -> None:
"""Building references must not hydrate every document the caller is
permitted to see -- only the (<= CHAT_RETRIEVER_TOP_K) documents that
the retriever actually returned nodes for.
"""
referenced = DocumentFactory.create(title="Referenced Document")
# Many more documents are "accessible" but never referenced by a node.
DocumentFactory.create_batch(200)
documents = Document.objects.all()
top_nodes = [
MagicMock(
metadata={
"document_id": str(referenced.pk),
"title": "Referenced Document",
},
),
]
hydrated_count = 0
def _count_hydration(sender, instance, **kwargs):
nonlocal hydrated_count
hydrated_count += 1
post_init.connect(_count_hydration, sender=Document)
try:
# One query: `documents.filter(pk__in=candidate_ids)` for the single
# referenced id. No query should scale with the 200 unreferenced documents.
with django_assert_num_queries(1):
references = chat._get_document_references(documents, top_nodes)
finally:
post_init.disconnect(_count_hydration, sender=Document)
# The bug this guards against: the old code hydrated all 201 accessible
# documents via `{doc.pk: doc for doc in documents}` before filtering by
# top_nodes. Only the referenced document should ever be constructed.
assert hydrated_count == 1
assert references == [{"id": referenced.pk, "title": "Referenced Document"}]