feat(security): add request size limits to API endpoints
- Add RequestSizeLimitMiddleware that checks Content-Length header before request body is read: non-multipart requests capped at MAX_REQUEST_BODY_SIZE (default 1 MB), multipart uploads capped at MAX_UPLOAD_SIZE (default 1 GB). Returns HTTP 413 on violation. - Register middleware in app/main.py - Add max_request_body_size setting to app/config.py - Fix ui_upload in files.py to check Content-Length early and read in 64 KB chunks (bounded memory usage), removing the post-write os.path.getsize check - Document MAX_REQUEST_BODY_SIZE in .env.demo and ConfigurationGuide.md - Mark SECURITY_AUDIT.md item #4 as resolved - Add 9 tests in test_request_size_limit.py - Update test_upload_file_too_large to use patch.object instead of the now-unused os.path.getsize mock Closes #173 Co-authored-by: christianlouis <361235+christianlouis@users.noreply.github.com>
This commit is contained in:
@@ -21,6 +21,12 @@ MAX_UPLOAD_SIZE=1073741824
|
|||||||
# Default: None (no splitting). Example: 104857600 for 100MB chunks
|
# Default: None (no splitting). Example: 104857600 for 100MB chunks
|
||||||
# MAX_SINGLE_FILE_SIZE=104857600
|
# MAX_SINGLE_FILE_SIZE=104857600
|
||||||
|
|
||||||
|
# **Request Body Size Limit** (Security - see SECURITY_AUDIT.md)
|
||||||
|
# Maximum request body size in bytes for non-file-upload requests (JSON, form data, etc.).
|
||||||
|
# Default: 1MB (1048576 bytes). File uploads are governed by MAX_UPLOAD_SIZE above.
|
||||||
|
# Prevents memory exhaustion from oversized JSON/form payloads.
|
||||||
|
# MAX_REQUEST_BODY_SIZE=1048576
|
||||||
|
|
||||||
# **Security Headers** (see SECURITY_AUDIT.md and docs/DeploymentGuide.md)
|
# **Security Headers** (see SECURITY_AUDIT.md and docs/DeploymentGuide.md)
|
||||||
# Disabled by default since most deployments use a reverse proxy (Traefik, Nginx, etc.)
|
# Disabled by default since most deployments use a reverse proxy (Traefik, Nginx, etc.)
|
||||||
# that already adds these headers. Set to true only if deploying directly without a reverse proxy.
|
# that already adds these headers. Set to true only if deploying directly without a reverse proxy.
|
||||||
|
|||||||
+3
-1
@@ -274,6 +274,8 @@ ftp = ftplib.FTP() # nosec B321 - Plaintext FTP intentional when configured
|
|||||||
- ✅ Unique filenames with UUID to prevent conflicts and overwrites
|
- ✅ Unique filenames with UUID to prevent conflicts and overwrites
|
||||||
- ✅ File upload size limits with configurable maximum (default: 1GB)
|
- ✅ File upload size limits with configurable maximum (default: 1GB)
|
||||||
- ✅ Optional file splitting for large PDFs (when max_single_file_size is configured)
|
- ✅ Optional file splitting for large PDFs (when max_single_file_size is configured)
|
||||||
|
- ✅ Request body size limits via `RequestSizeLimitMiddleware` (non-upload: 1MB default; uploads: governed by MAX_UPLOAD_SIZE)
|
||||||
|
- ✅ Streaming file reads in upload endpoint to prevent memory exhaustion
|
||||||
- ⏳ **TODO:** Implement rate limiting on API endpoints
|
- ⏳ **TODO:** Implement rate limiting on API endpoints
|
||||||
- ⏳ **TODO:** Add CSRF protection for state-changing operations
|
- ⏳ **TODO:** Add CSRF protection for state-changing operations
|
||||||
- ⏳ **TODO:** Add comprehensive input sanitization for all user inputs ([#172](https://github.com/christianlouis/DocuElevate/issues/172))
|
- ⏳ **TODO:** Add comprehensive input sanitization for all user inputs ([#172](https://github.com/christianlouis/DocuElevate/issues/172))
|
||||||
@@ -302,7 +304,7 @@ ftp = ftplib.FTP() # nosec B321 - Plaintext FTP intentional when configured
|
|||||||
1. ~~**Enable CodeQL scanning**~~ ✅ Already implemented - Two CodeQL workflows active
|
1. ~~**Enable CodeQL scanning**~~ ✅ Already implemented - Two CodeQL workflows active
|
||||||
2. **Implement rate limiting** - Prevent abuse and DoS attacks (consider slowapi or fastapi-limiter)
|
2. **Implement rate limiting** - Prevent abuse and DoS attacks (consider slowapi or fastapi-limiter)
|
||||||
3. **Add comprehensive input validation** - Prevent injection attacks ([#172](https://github.com/christianlouis/DocuElevate/issues/172))
|
3. **Add comprehensive input validation** - Prevent injection attacks ([#172](https://github.com/christianlouis/DocuElevate/issues/172))
|
||||||
4. **Add request size limits** - Prevent memory exhaustion from large uploads ([#173](https://github.com/christianlouis/DocuElevate/issues/173))
|
4. ~~**Add request size limits**~~ ✅ Implemented - `RequestSizeLimitMiddleware` enforces `MAX_REQUEST_BODY_SIZE` (default 1 MB) for non-file requests and `MAX_UPLOAD_SIZE` (default 1 GB) for multipart uploads; file uploads also use streaming reads to bound memory usage ([#173](https://github.com/christianlouis/DocuElevate/issues/173))
|
||||||
5. **Implement CSRF protection** - Protect state-changing operations
|
5. **Implement CSRF protection** - Protect state-changing operations
|
||||||
|
|
||||||
### Medium Priority
|
### Medium Priority
|
||||||
|
|||||||
+313
-95
@@ -45,7 +45,8 @@ def list_files_api(
|
|||||||
page: int = Query(1, ge=1, description="Page number"),
|
page: int = Query(1, ge=1, description="Page number"),
|
||||||
per_page: int = Query(50, ge=1, le=200, description="Items per page"),
|
per_page: int = Query(50, ge=1, le=200, description="Items per page"),
|
||||||
sort_by: str = Query(
|
sort_by: str = Query(
|
||||||
"created_at", description="Sort field: id, original_filename, file_size, mime_type, created_at, status"
|
"created_at",
|
||||||
|
description="Sort field: id, original_filename, file_size, mime_type, created_at, status",
|
||||||
),
|
),
|
||||||
sort_order: str = Query("desc", description="Sort order: asc or desc"),
|
sort_order: str = Query("desc", description="Sort order: asc or desc"),
|
||||||
search: Optional[str] = Query(None, description="Search in filename"),
|
search: Optional[str] = Query(None, description="Search in filename"),
|
||||||
@@ -128,7 +129,13 @@ def list_files_api(
|
|||||||
"mime_type": f.mime_type,
|
"mime_type": f.mime_type,
|
||||||
"created_at": f.created_at.isoformat() if f.created_at else None,
|
"created_at": f.created_at.isoformat() if f.created_at else None,
|
||||||
"processing_status": statuses.get(
|
"processing_status": statuses.get(
|
||||||
f.id, {"status": "pending", "last_step": None, "has_errors": False, "total_steps": 0}
|
f.id,
|
||||||
|
{
|
||||||
|
"status": "pending",
|
||||||
|
"last_step": None,
|
||||||
|
"has_errors": False,
|
||||||
|
"total_steps": 0,
|
||||||
|
},
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
@@ -138,7 +145,12 @@ def list_files_api(
|
|||||||
|
|
||||||
return {
|
return {
|
||||||
"files": result,
|
"files": result,
|
||||||
"pagination": {"page": page, "per_page": per_page, "total_items": total_items, "total_pages": total_pages},
|
"pagination": {
|
||||||
|
"page": page,
|
||||||
|
"per_page": per_page,
|
||||||
|
"total_items": total_items,
|
||||||
|
"total_pages": total_pages,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -162,11 +174,16 @@ def get_file_details(request: Request, file_id: int, db: DbSession):
|
|||||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||||
|
|
||||||
if not file_record:
|
if not file_record:
|
||||||
raise HTTPException(status_code=404, detail=f"File record with ID {file_id} not found")
|
raise HTTPException(
|
||||||
|
status_code=404, detail=f"File record with ID {file_id} not found"
|
||||||
|
)
|
||||||
|
|
||||||
# Get processing logs
|
# Get processing logs
|
||||||
logs = (
|
logs = (
|
||||||
db.query(ProcessingLog).filter(ProcessingLog.file_id == file_id).order_by(ProcessingLog.timestamp.desc()).all()
|
db.query(ProcessingLog)
|
||||||
|
.filter(ProcessingLog.file_id == file_id)
|
||||||
|
.order_by(ProcessingLog.timestamp.desc())
|
||||||
|
.all()
|
||||||
)
|
)
|
||||||
|
|
||||||
# Build log list
|
# Build log list
|
||||||
@@ -187,7 +204,13 @@ def get_file_details(request: Request, file_id: int, db: DbSession):
|
|||||||
processing_status = _get_file_processing_status(db, file_id)
|
processing_status = _get_file_processing_status(db, file_id)
|
||||||
|
|
||||||
# Check if files exist on disk
|
# Check if files exist on disk
|
||||||
files_on_disk = {"original": os.path.exists(file_record.local_filename) if file_record.local_filename else False}
|
files_on_disk = {
|
||||||
|
"original": (
|
||||||
|
os.path.exists(file_record.local_filename)
|
||||||
|
if file_record.local_filename
|
||||||
|
else False
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"file": {
|
"file": {
|
||||||
@@ -197,7 +220,9 @@ def get_file_details(request: Request, file_id: int, db: DbSession):
|
|||||||
"local_filename": file_record.local_filename,
|
"local_filename": file_record.local_filename,
|
||||||
"file_size": file_record.file_size,
|
"file_size": file_record.file_size,
|
||||||
"mime_type": file_record.mime_type,
|
"mime_type": file_record.mime_type,
|
||||||
"created_at": file_record.created_at.isoformat() if file_record.created_at else None,
|
"created_at": (
|
||||||
|
file_record.created_at.isoformat() if file_record.created_at else None
|
||||||
|
),
|
||||||
},
|
},
|
||||||
"processing_status": processing_status,
|
"processing_status": processing_status,
|
||||||
"logs": log_list,
|
"logs": log_list,
|
||||||
@@ -214,30 +239,41 @@ def delete_file_record(request: Request, file_id: int, db: DbSession):
|
|||||||
"""
|
"""
|
||||||
# Check if file deletion is allowed
|
# Check if file deletion is allowed
|
||||||
if not settings.allow_file_delete:
|
if not settings.allow_file_delete:
|
||||||
raise HTTPException(status_code=403, detail="File deletion is disabled in the configuration")
|
raise HTTPException(
|
||||||
|
status_code=403, detail="File deletion is disabled in the configuration"
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Find the file record
|
# Find the file record
|
||||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||||
|
|
||||||
if not file_record:
|
if not file_record:
|
||||||
raise HTTPException(status_code=404, detail=f"File record with ID {file_id} not found")
|
raise HTTPException(
|
||||||
|
status_code=404, detail=f"File record with ID {file_id} not found"
|
||||||
|
)
|
||||||
|
|
||||||
# Log the deletion
|
# Log the deletion
|
||||||
logger.info(f"Deleting file record: ID={file_id}, Filename={file_record.original_filename}")
|
logger.info(
|
||||||
|
f"Deleting file record: ID={file_id}, Filename={file_record.original_filename}"
|
||||||
|
)
|
||||||
|
|
||||||
# Delete the record
|
# Delete the record
|
||||||
db.delete(file_record)
|
db.delete(file_record)
|
||||||
db.commit()
|
db.commit()
|
||||||
|
|
||||||
return {"status": "success", "message": f"File record {file_id} deleted successfully"}
|
return {
|
||||||
|
"status": "success",
|
||||||
|
"message": f"File record {file_id} deleted successfully",
|
||||||
|
}
|
||||||
|
|
||||||
except HTTPException:
|
except HTTPException:
|
||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
db.rollback()
|
db.rollback()
|
||||||
logger.exception(f"Error deleting file record {file_id}: {str(e)}")
|
logger.exception(f"Error deleting file record {file_id}: {str(e)}")
|
||||||
raise HTTPException(status_code=500, detail=f"Error deleting file record: {str(e)}")
|
raise HTTPException(
|
||||||
|
status_code=500, detail=f"Error deleting file record: {str(e)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@router.post("/files/bulk-delete")
|
@router.post("/files/bulk-delete")
|
||||||
@@ -249,14 +285,18 @@ def bulk_delete_files(request: Request, file_ids: List[int], db: DbSession):
|
|||||||
"""
|
"""
|
||||||
# Check if file deletion is allowed
|
# Check if file deletion is allowed
|
||||||
if not settings.allow_file_delete:
|
if not settings.allow_file_delete:
|
||||||
raise HTTPException(status_code=403, detail="File deletion is disabled in the configuration")
|
raise HTTPException(
|
||||||
|
status_code=403, detail="File deletion is disabled in the configuration"
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Find all file records
|
# Find all file records
|
||||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||||
|
|
||||||
if not file_records:
|
if not file_records:
|
||||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
raise HTTPException(
|
||||||
|
status_code=404, detail="No files found with the provided IDs"
|
||||||
|
)
|
||||||
|
|
||||||
deleted_count = len(file_records)
|
deleted_count = len(file_records)
|
||||||
deleted_ids = [f.id for f in file_records]
|
deleted_ids = [f.id for f in file_records]
|
||||||
@@ -281,7 +321,9 @@ def bulk_delete_files(request: Request, file_ids: List[int], db: DbSession):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
db.rollback()
|
db.rollback()
|
||||||
logger.exception(f"Error bulk deleting file records: {str(e)}")
|
logger.exception(f"Error bulk deleting file records: {str(e)}")
|
||||||
raise HTTPException(status_code=500, detail=f"Error bulk deleting file records: {str(e)}")
|
raise HTTPException(
|
||||||
|
status_code=500, detail=f"Error bulk deleting file records: {str(e)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@router.post("/files/bulk-reprocess")
|
@router.post("/files/bulk-reprocess")
|
||||||
@@ -295,7 +337,9 @@ def bulk_reprocess_files(request: Request, file_ids: List[int], db: DbSession):
|
|||||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||||
|
|
||||||
if not file_records:
|
if not file_records:
|
||||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
raise HTTPException(
|
||||||
|
status_code=404, detail="No files found with the provided IDs"
|
||||||
|
)
|
||||||
|
|
||||||
task_ids = []
|
task_ids = []
|
||||||
processed_files = []
|
processed_files = []
|
||||||
@@ -304,7 +348,9 @@ def bulk_reprocess_files(request: Request, file_ids: List[int], db: DbSession):
|
|||||||
for file_record in file_records:
|
for file_record in file_records:
|
||||||
try:
|
try:
|
||||||
# Check if local file exists
|
# Check if local file exists
|
||||||
if not file_record.local_filename or not os.path.exists(file_record.local_filename):
|
if not file_record.local_filename or not os.path.exists(
|
||||||
|
file_record.local_filename
|
||||||
|
):
|
||||||
errors.append(
|
errors.append(
|
||||||
{
|
{
|
||||||
"file_id": file_record.id,
|
"file_id": file_record.id,
|
||||||
@@ -315,10 +361,16 @@ def bulk_reprocess_files(request: Request, file_ids: List[int], db: DbSession):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
# Queue the file for processing, passing file_id to skip duplicate check
|
# Queue the file for processing, passing file_id to skip duplicate check
|
||||||
task = process_document.delay(file_record.local_filename, file_id=file_record.id)
|
task = process_document.delay(
|
||||||
|
file_record.local_filename, file_id=file_record.id
|
||||||
|
)
|
||||||
task_ids.append(task.id)
|
task_ids.append(task.id)
|
||||||
processed_files.append(
|
processed_files.append(
|
||||||
{"file_id": file_record.id, "filename": file_record.original_filename, "task_id": task.id}
|
{
|
||||||
|
"file_id": file_record.id,
|
||||||
|
"filename": file_record.original_filename,
|
||||||
|
"task_id": task.id,
|
||||||
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -328,7 +380,13 @@ def bulk_reprocess_files(request: Request, file_ids: List[int], db: DbSession):
|
|||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Error reprocessing file {file_record.id}: {str(e)}")
|
logger.exception(f"Error reprocessing file {file_record.id}: {str(e)}")
|
||||||
errors.append({"file_id": file_record.id, "filename": file_record.original_filename, "error": str(e)})
|
errors.append(
|
||||||
|
{
|
||||||
|
"file_id": file_record.id,
|
||||||
|
"filename": file_record.original_filename,
|
||||||
|
"error": str(e),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"status": "success" if processed_files else "error",
|
"status": "success" if processed_files else "error",
|
||||||
@@ -342,7 +400,9 @@ def bulk_reprocess_files(request: Request, file_ids: List[int], db: DbSession):
|
|||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Error bulk reprocessing files: {str(e)}")
|
logger.exception(f"Error bulk reprocessing files: {str(e)}")
|
||||||
raise HTTPException(status_code=500, detail=f"Error bulk reprocessing files: {str(e)}")
|
raise HTTPException(
|
||||||
|
status_code=500, detail=f"Error bulk reprocessing files: {str(e)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@router.post("/files/{file_id}/reprocess")
|
@router.post("/files/{file_id}/reprocess")
|
||||||
@@ -362,15 +422,24 @@ def reprocess_single_file(request: Request, file_id: int, db: DbSession):
|
|||||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||||
|
|
||||||
if not file_record:
|
if not file_record:
|
||||||
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
|
raise HTTPException(
|
||||||
|
status_code=404, detail=f"File with ID {file_id} not found"
|
||||||
|
)
|
||||||
|
|
||||||
# Check if local file exists
|
# Check if local file exists
|
||||||
if not file_record.local_filename or not os.path.exists(file_record.local_filename):
|
if not file_record.local_filename or not os.path.exists(
|
||||||
raise HTTPException(status_code=400, detail="Local file not found on disk. Cannot reprocess.")
|
file_record.local_filename
|
||||||
|
):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail="Local file not found on disk. Cannot reprocess.",
|
||||||
|
)
|
||||||
|
|
||||||
# Queue the file for processing, passing file_id to skip duplicate check
|
# Queue the file for processing, passing file_id to skip duplicate check
|
||||||
task = process_document.delay(
|
task = process_document.delay(
|
||||||
file_record.local_filename, original_filename=file_record.original_filename, file_id=file_record.id
|
file_record.local_filename,
|
||||||
|
original_filename=file_record.original_filename,
|
||||||
|
file_id=file_record.id,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -389,7 +458,9 @@ def reprocess_single_file(request: Request, file_id: int, db: DbSession):
|
|||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Error reprocessing file {file_id}: {str(e)}")
|
logger.exception(f"Error reprocessing file {file_id}: {str(e)}")
|
||||||
raise HTTPException(status_code=500, detail=f"Error reprocessing file: {str(e)}")
|
raise HTTPException(
|
||||||
|
status_code=500, detail=f"Error reprocessing file: {str(e)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@router.post("/files/{file_id}/reprocess-with-cloud-ocr")
|
@router.post("/files/{file_id}/reprocess-with-cloud-ocr")
|
||||||
@@ -413,24 +484,34 @@ def reprocess_with_cloud_ocr(request: Request, file_id: int, db: DbSession):
|
|||||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||||
|
|
||||||
if not file_record:
|
if not file_record:
|
||||||
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
|
raise HTTPException(
|
||||||
|
status_code=404, detail=f"File with ID {file_id} not found"
|
||||||
|
)
|
||||||
|
|
||||||
# Prefer using the original_file_path if available, otherwise fall back to local_filename
|
# Prefer using the original_file_path if available, otherwise fall back to local_filename
|
||||||
source_file = None
|
source_file = None
|
||||||
if file_record.original_file_path and os.path.exists(file_record.original_file_path):
|
if file_record.original_file_path and os.path.exists(
|
||||||
|
file_record.original_file_path
|
||||||
|
):
|
||||||
source_file = file_record.original_file_path
|
source_file = file_record.original_file_path
|
||||||
logger.info(f"Using original file for Cloud OCR reprocessing: {source_file}")
|
logger.info(
|
||||||
|
f"Using original file for Cloud OCR reprocessing: {source_file}"
|
||||||
|
)
|
||||||
elif file_record.local_filename and os.path.exists(file_record.local_filename):
|
elif file_record.local_filename and os.path.exists(file_record.local_filename):
|
||||||
source_file = file_record.local_filename
|
source_file = file_record.local_filename
|
||||||
logger.info(f"Using local file for Cloud OCR reprocessing: {source_file}")
|
logger.info(f"Using local file for Cloud OCR reprocessing: {source_file}")
|
||||||
else:
|
else:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=400, detail="Neither original nor local file found on disk. Cannot reprocess."
|
status_code=400,
|
||||||
|
detail="Neither original nor local file found on disk. Cannot reprocess.",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Queue the file for processing with force_cloud_ocr=True
|
# Queue the file for processing with force_cloud_ocr=True
|
||||||
task = process_document.delay(
|
task = process_document.delay(
|
||||||
source_file, original_filename=file_record.original_filename, file_id=file_record.id, force_cloud_ocr=True
|
source_file,
|
||||||
|
original_filename=file_record.original_filename,
|
||||||
|
file_id=file_record.id,
|
||||||
|
force_cloud_ocr=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -451,7 +532,9 @@ def reprocess_with_cloud_ocr(request: Request, file_id: int, db: DbSession):
|
|||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Error reprocessing file {file_id} with Cloud OCR: {str(e)}")
|
logger.exception(f"Error reprocessing file {file_id} with Cloud OCR: {str(e)}")
|
||||||
raise HTTPException(status_code=500, detail=f"Error reprocessing file with Cloud OCR: {str(e)}")
|
raise HTTPException(
|
||||||
|
status_code=500, detail=f"Error reprocessing file with Cloud OCR: {str(e)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _extract_text_from_pdf(file_path: str) -> str:
|
def _extract_text_from_pdf(file_path: str) -> str:
|
||||||
@@ -496,24 +579,40 @@ def _retry_pipeline_step(file_record: FileRecord, step_name: str, db: Session) -
|
|||||||
|
|
||||||
if step_name == "process_document":
|
if step_name == "process_document":
|
||||||
# Full reprocessing with duplicate check bypass
|
# Full reprocessing with duplicate check bypass
|
||||||
logger.info(f"Retrying process_document for file {file_id}: local_filename={file_record.local_filename!r}")
|
logger.info(
|
||||||
|
f"Retrying process_document for file {file_id}: local_filename={file_record.local_filename!r}"
|
||||||
|
)
|
||||||
if not file_record.local_filename:
|
if not file_record.local_filename:
|
||||||
logger.error(f"process_document retry failed for file {file_id}: local_filename is None")
|
logger.error(
|
||||||
raise HTTPException(status_code=400, detail="Local file path is None. Cannot retry.")
|
f"process_document retry failed for file {file_id}: local_filename is None"
|
||||||
|
)
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400, detail="Local file path is None. Cannot retry."
|
||||||
|
)
|
||||||
|
|
||||||
exists = os.path.exists(file_record.local_filename)
|
exists = os.path.exists(file_record.local_filename)
|
||||||
logger.info(f"Checking local_filename: {file_record.local_filename!r}, exists={exists}")
|
logger.info(
|
||||||
|
f"Checking local_filename: {file_record.local_filename!r}, exists={exists}"
|
||||||
|
)
|
||||||
if not exists:
|
if not exists:
|
||||||
error_message = f"Local file not found on disk. Cannot retry. Path checked: local_filename={file_record.local_filename!r} (exists=False)"
|
error_message = f"Local file not found on disk. Cannot retry. Path checked: local_filename={file_record.local_filename!r} (exists=False)"
|
||||||
logger.error(f"process_document retry failed for file {file_id}: {error_message}")
|
logger.error(
|
||||||
|
f"process_document retry failed for file {file_id}: {error_message}"
|
||||||
|
)
|
||||||
raise HTTPException(status_code=400, detail=error_message)
|
raise HTTPException(status_code=400, detail=error_message)
|
||||||
|
|
||||||
logger.info(f"Found file for process_document retry at: {file_record.local_filename!r}")
|
logger.info(
|
||||||
|
f"Found file for process_document retry at: {file_record.local_filename!r}"
|
||||||
|
)
|
||||||
task = process_document.delay(
|
task = process_document.delay(
|
||||||
file_record.local_filename, original_filename=file_record.original_filename, file_id=file_id
|
file_record.local_filename,
|
||||||
|
original_filename=file_record.original_filename,
|
||||||
|
file_id=file_id,
|
||||||
)
|
)
|
||||||
elif step_name == "process_with_azure_document_intelligence":
|
elif step_name == "process_with_azure_document_intelligence":
|
||||||
from app.tasks.process_with_azure_document_intelligence import process_with_azure_document_intelligence
|
from app.tasks.process_with_azure_document_intelligence import (
|
||||||
|
process_with_azure_document_intelligence,
|
||||||
|
)
|
||||||
|
|
||||||
# OCR needs the file in workdir/tmp
|
# OCR needs the file in workdir/tmp
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -522,10 +621,14 @@ def _retry_pipeline_step(file_record: FileRecord, step_name: str, db: Session) -
|
|||||||
)
|
)
|
||||||
if not file_record.local_filename:
|
if not file_record.local_filename:
|
||||||
logger.error(f"OCR retry failed for file {file_id}: local_filename is None")
|
logger.error(f"OCR retry failed for file {file_id}: local_filename is None")
|
||||||
raise HTTPException(status_code=400, detail="Local file path is None. Cannot retry OCR.")
|
raise HTTPException(
|
||||||
|
status_code=400, detail="Local file path is None. Cannot retry OCR."
|
||||||
|
)
|
||||||
|
|
||||||
exists = os.path.exists(file_record.local_filename)
|
exists = os.path.exists(file_record.local_filename)
|
||||||
logger.info(f"Checking local_filename: {file_record.local_filename!r}, exists={exists}")
|
logger.info(
|
||||||
|
f"Checking local_filename: {file_record.local_filename!r}, exists={exists}"
|
||||||
|
)
|
||||||
if not exists:
|
if not exists:
|
||||||
error_message = f"Local file not found on disk. Cannot retry OCR. Path checked: local_filename={file_record.local_filename!r} (exists=False)"
|
error_message = f"Local file not found on disk. Cannot retry OCR. Path checked: local_filename={file_record.local_filename!r} (exists=False)"
|
||||||
logger.error(f"OCR retry failed for file {file_id}: {error_message}")
|
logger.error(f"OCR retry failed for file {file_id}: {error_message}")
|
||||||
@@ -541,22 +644,35 @@ def _retry_pipeline_step(file_record: FileRecord, step_name: str, db: Session) -
|
|||||||
f"Retrying extract_metadata_with_gpt for file {file_id}: local_filename={file_record.local_filename!r}"
|
f"Retrying extract_metadata_with_gpt for file {file_id}: local_filename={file_record.local_filename!r}"
|
||||||
)
|
)
|
||||||
if not file_record.local_filename:
|
if not file_record.local_filename:
|
||||||
logger.error(f"Metadata extraction retry failed for file {file_id}: local_filename is None")
|
logger.error(
|
||||||
raise HTTPException(status_code=400, detail="Local file path is None. Cannot retry metadata extraction.")
|
f"Metadata extraction retry failed for file {file_id}: local_filename is None"
|
||||||
|
)
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail="Local file path is None. Cannot retry metadata extraction.",
|
||||||
|
)
|
||||||
|
|
||||||
exists = os.path.exists(file_record.local_filename)
|
exists = os.path.exists(file_record.local_filename)
|
||||||
logger.info(f"Checking local_filename: {file_record.local_filename!r}, exists={exists}")
|
logger.info(
|
||||||
|
f"Checking local_filename: {file_record.local_filename!r}, exists={exists}"
|
||||||
|
)
|
||||||
if not exists:
|
if not exists:
|
||||||
error_message = f"Local file not found on disk. Cannot retry metadata extraction. Path checked: local_filename={file_record.local_filename!r} (exists=False)"
|
error_message = f"Local file not found on disk. Cannot retry metadata extraction. Path checked: local_filename={file_record.local_filename!r} (exists=False)"
|
||||||
logger.error(f"Metadata extraction retry failed for file {file_id}: {error_message}")
|
logger.error(
|
||||||
|
f"Metadata extraction retry failed for file {file_id}: {error_message}"
|
||||||
|
)
|
||||||
raise HTTPException(status_code=400, detail=error_message)
|
raise HTTPException(status_code=400, detail=error_message)
|
||||||
|
|
||||||
logger.info(f"Found file for metadata extraction retry at: {file_record.local_filename!r}")
|
logger.info(
|
||||||
|
f"Found file for metadata extraction retry at: {file_record.local_filename!r}"
|
||||||
|
)
|
||||||
extracted_text = _extract_text_from_pdf(file_record.local_filename)
|
extracted_text = _extract_text_from_pdf(file_record.local_filename)
|
||||||
filename = os.path.basename(file_record.local_filename)
|
filename = os.path.basename(file_record.local_filename)
|
||||||
task = extract_metadata_with_gpt.delay(filename, extracted_text, file_id)
|
task = extract_metadata_with_gpt.delay(filename, extracted_text, file_id)
|
||||||
elif step_name == "embed_metadata_into_pdf":
|
elif step_name == "embed_metadata_into_pdf":
|
||||||
from app.tasks.extract_metadata_with_gpt import extract_metadata_with_gpt as extract_metadata_task
|
from app.tasks.extract_metadata_with_gpt import (
|
||||||
|
extract_metadata_with_gpt as extract_metadata_task,
|
||||||
|
)
|
||||||
|
|
||||||
# Retrying embed requires re-running metadata extraction first, because
|
# Retrying embed requires re-running metadata extraction first, because
|
||||||
# embed_metadata_into_pdf needs the actual metadata dict (not empty).
|
# embed_metadata_into_pdf needs the actual metadata dict (not empty).
|
||||||
@@ -581,24 +697,36 @@ def _retry_pipeline_step(file_record: FileRecord, step_name: str, db: Session) -
|
|||||||
# Check 1: local_filename (original tmp location)
|
# Check 1: local_filename (original tmp location)
|
||||||
if file_record.local_filename:
|
if file_record.local_filename:
|
||||||
exists = os.path.exists(file_record.local_filename)
|
exists = os.path.exists(file_record.local_filename)
|
||||||
checked_paths.append(f"local_filename={file_record.local_filename!r} (exists={exists})")
|
checked_paths.append(
|
||||||
logger.info(f"Checking local_filename: {file_record.local_filename!r}, exists={exists}")
|
f"local_filename={file_record.local_filename!r} (exists={exists})"
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"Checking local_filename: {file_record.local_filename!r}, exists={exists}"
|
||||||
|
)
|
||||||
if exists:
|
if exists:
|
||||||
file_path = file_record.local_filename
|
file_path = file_record.local_filename
|
||||||
|
|
||||||
# Check 2: processed_file_path
|
# Check 2: processed_file_path
|
||||||
if not file_path and file_record.processed_file_path:
|
if not file_path and file_record.processed_file_path:
|
||||||
exists = os.path.exists(file_record.processed_file_path)
|
exists = os.path.exists(file_record.processed_file_path)
|
||||||
checked_paths.append(f"processed_file_path={file_record.processed_file_path!r} (exists={exists})")
|
checked_paths.append(
|
||||||
logger.info(f"Checking processed_file_path: {file_record.processed_file_path!r}, exists={exists}")
|
f"processed_file_path={file_record.processed_file_path!r} (exists={exists})"
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"Checking processed_file_path: {file_record.processed_file_path!r}, exists={exists}"
|
||||||
|
)
|
||||||
if exists:
|
if exists:
|
||||||
file_path = file_record.processed_file_path
|
file_path = file_record.processed_file_path
|
||||||
|
|
||||||
# Check 3: original_file_path (immutable copy in workdir/original/)
|
# Check 3: original_file_path (immutable copy in workdir/original/)
|
||||||
if not file_path and file_record.original_file_path:
|
if not file_path and file_record.original_file_path:
|
||||||
exists = os.path.exists(file_record.original_file_path)
|
exists = os.path.exists(file_record.original_file_path)
|
||||||
checked_paths.append(f"original_file_path={file_record.original_file_path!r} (exists={exists})")
|
checked_paths.append(
|
||||||
logger.info(f"Checking original_file_path: {file_record.original_file_path!r}, exists={exists}")
|
f"original_file_path={file_record.original_file_path!r} (exists={exists})"
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"Checking original_file_path: {file_record.original_file_path!r}, exists={exists}"
|
||||||
|
)
|
||||||
if exists:
|
if exists:
|
||||||
file_path = file_record.original_file_path
|
file_path = file_record.original_file_path
|
||||||
|
|
||||||
@@ -606,17 +734,25 @@ def _retry_pipeline_step(file_record: FileRecord, step_name: str, db: Session) -
|
|||||||
if not file_path and file_record.local_filename:
|
if not file_path and file_record.local_filename:
|
||||||
workdir = settings.workdir
|
workdir = settings.workdir
|
||||||
tmp_dir = os.path.join(workdir, "tmp")
|
tmp_dir = os.path.join(workdir, "tmp")
|
||||||
fallback_path = os.path.join(tmp_dir, os.path.basename(file_record.local_filename))
|
fallback_path = os.path.join(
|
||||||
|
tmp_dir, os.path.basename(file_record.local_filename)
|
||||||
|
)
|
||||||
exists = os.path.exists(fallback_path)
|
exists = os.path.exists(fallback_path)
|
||||||
checked_paths.append(f"workdir_tmp_fallback={fallback_path!r} (exists={exists})")
|
checked_paths.append(
|
||||||
logger.info(f"Checking workdir/tmp fallback: {fallback_path!r}, exists={exists}")
|
f"workdir_tmp_fallback={fallback_path!r} (exists={exists})"
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"Checking workdir/tmp fallback: {fallback_path!r}, exists={exists}"
|
||||||
|
)
|
||||||
if exists:
|
if exists:
|
||||||
file_path = fallback_path
|
file_path = fallback_path
|
||||||
|
|
||||||
if not file_path:
|
if not file_path:
|
||||||
paths_detail = "; ".join(checked_paths)
|
paths_detail = "; ".join(checked_paths)
|
||||||
error_message = f"File not found on disk. Cannot retry metadata embedding. Paths checked: {paths_detail}"
|
error_message = f"File not found on disk. Cannot retry metadata embedding. Paths checked: {paths_detail}"
|
||||||
logger.error(f"embed_metadata_into_pdf retry failed for file {file_id}: {error_message}")
|
logger.error(
|
||||||
|
f"embed_metadata_into_pdf retry failed for file {file_id}: {error_message}"
|
||||||
|
)
|
||||||
raise HTTPException(status_code=400, detail=error_message)
|
raise HTTPException(status_code=400, detail=error_message)
|
||||||
|
|
||||||
logger.info(f"Found file for embed_metadata_into_pdf retry at: {file_path!r}")
|
logger.info(f"Found file for embed_metadata_into_pdf retry at: {file_path!r}")
|
||||||
@@ -624,9 +760,13 @@ def _retry_pipeline_step(file_record: FileRecord, step_name: str, db: Session) -
|
|||||||
# Pass the full path to the task so it can locate the file
|
# Pass the full path to the task so it can locate the file
|
||||||
task = extract_metadata_task.delay(file_path, extracted_text, file_id)
|
task = extract_metadata_task.delay(file_path, extracted_text, file_id)
|
||||||
else:
|
else:
|
||||||
raise HTTPException(status_code=400, detail=f"Unsupported pipeline step: {step_name}")
|
raise HTTPException(
|
||||||
|
status_code=400, detail=f"Unsupported pipeline step: {step_name}"
|
||||||
|
)
|
||||||
|
|
||||||
logger.info(f"Retrying pipeline step: FileID={file_record.id}, Step={step_name}, TaskID={task.id}")
|
logger.info(
|
||||||
|
f"Retrying pipeline step: FileID={file_record.id}, Step={step_name}, TaskID={task.id}"
|
||||||
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"status": "success",
|
"status": "success",
|
||||||
@@ -644,7 +784,8 @@ def retry_subtask(
|
|||||||
file_id: int,
|
file_id: int,
|
||||||
db: DbSession,
|
db: DbSession,
|
||||||
subtask_name: str = Query(
|
subtask_name: str = Query(
|
||||||
..., description="Name of the subtask to retry (e.g., 'upload_to_dropbox', 'extract_metadata_with_gpt')"
|
...,
|
||||||
|
description="Name of the subtask to retry (e.g., 'upload_to_dropbox', 'extract_metadata_with_gpt')",
|
||||||
),
|
),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -666,7 +807,9 @@ def retry_subtask(
|
|||||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||||
|
|
||||||
if not file_record:
|
if not file_record:
|
||||||
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
|
raise HTTPException(
|
||||||
|
status_code=404, detail=f"File with ID {file_id} not found"
|
||||||
|
)
|
||||||
|
|
||||||
# Pipeline processing steps that can be retried from the failed step
|
# Pipeline processing steps that can be retried from the failed step
|
||||||
pipeline_step_names = {
|
pipeline_step_names = {
|
||||||
@@ -749,7 +892,9 @@ def retry_subtask(
|
|||||||
upload_task = task_map[subtask_name]
|
upload_task = task_map[subtask_name]
|
||||||
task = upload_task.delay(file_path, file_id)
|
task = upload_task.delay(file_path, file_id)
|
||||||
|
|
||||||
logger.info(f"Retrying upload subtask: FileID={file_record.id}, Subtask={subtask_name}, TaskID={task.id}")
|
logger.info(
|
||||||
|
f"Retrying upload subtask: FileID={file_record.id}, Subtask={subtask_name}, TaskID={task.id}"
|
||||||
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"status": "success",
|
"status": "success",
|
||||||
@@ -762,7 +907,9 @@ def retry_subtask(
|
|||||||
except HTTPException:
|
except HTTPException:
|
||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Error retrying subtask {subtask_name} for file {file_id}: {str(e)}")
|
logger.exception(
|
||||||
|
f"Error retrying subtask {subtask_name} for file {file_id}: {str(e)}"
|
||||||
|
)
|
||||||
raise HTTPException(status_code=500, detail=f"Error retrying subtask: {str(e)}")
|
raise HTTPException(status_code=500, detail=f"Error retrying subtask: {str(e)}")
|
||||||
|
|
||||||
|
|
||||||
@@ -791,12 +938,18 @@ def get_file_preview(
|
|||||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||||
|
|
||||||
if not file_record:
|
if not file_record:
|
||||||
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
|
raise HTTPException(
|
||||||
|
status_code=404, detail=f"File with ID {file_id} not found"
|
||||||
|
)
|
||||||
|
|
||||||
if version == "original":
|
if version == "original":
|
||||||
# Return the original file from tmp
|
# Return the original file from tmp
|
||||||
if not file_record.local_filename or not os.path.exists(file_record.local_filename):
|
if not file_record.local_filename or not os.path.exists(
|
||||||
raise HTTPException(status_code=404, detail="Original file not found on disk")
|
file_record.local_filename
|
||||||
|
):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=404, detail="Original file not found on disk"
|
||||||
|
)
|
||||||
|
|
||||||
file_path = file_record.local_filename
|
file_path = file_record.local_filename
|
||||||
|
|
||||||
@@ -822,20 +975,27 @@ def get_file_preview(
|
|||||||
if not file_path:
|
if not file_path:
|
||||||
raise HTTPException(status_code=404, detail="Processed file not found")
|
raise HTTPException(status_code=404, detail="Processed file not found")
|
||||||
else:
|
else:
|
||||||
raise HTTPException(status_code=400, detail="Invalid version parameter. Use 'original' or 'processed'")
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail="Invalid version parameter. Use 'original' or 'processed'",
|
||||||
|
)
|
||||||
|
|
||||||
# Return the file
|
# Return the file
|
||||||
return FileResponse(
|
return FileResponse(
|
||||||
path=file_path,
|
path=file_path,
|
||||||
media_type=file_record.mime_type or "application/pdf",
|
media_type=file_record.mime_type or "application/pdf",
|
||||||
headers={"Content-Disposition": f'inline; filename="{file_record.original_filename}"'},
|
headers={
|
||||||
|
"Content-Disposition": f'inline; filename="{file_record.original_filename}"'
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
except HTTPException:
|
except HTTPException:
|
||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Error retrieving file preview: {str(e)}")
|
logger.exception(f"Error retrieving file preview: {str(e)}")
|
||||||
raise HTTPException(status_code=500, detail=f"Error retrieving file preview: {str(e)}")
|
raise HTTPException(
|
||||||
|
status_code=500, detail=f"Error retrieving file preview: {str(e)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@router.get("/files/{file_id}/download")
|
@router.get("/files/{file_id}/download")
|
||||||
@@ -863,12 +1023,18 @@ def download_file(
|
|||||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||||
|
|
||||||
if not file_record:
|
if not file_record:
|
||||||
raise HTTPException(status_code=404, detail=f"File with ID {file_id} not found")
|
raise HTTPException(
|
||||||
|
status_code=404, detail=f"File with ID {file_id} not found"
|
||||||
|
)
|
||||||
|
|
||||||
if version == "original":
|
if version == "original":
|
||||||
# Return the original file from tmp
|
# Return the original file from tmp
|
||||||
if not file_record.local_filename or not os.path.exists(file_record.local_filename):
|
if not file_record.local_filename or not os.path.exists(
|
||||||
raise HTTPException(status_code=404, detail="Original file not found on disk")
|
file_record.local_filename
|
||||||
|
):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=404, detail="Original file not found on disk"
|
||||||
|
)
|
||||||
|
|
||||||
file_path = file_record.local_filename
|
file_path = file_record.local_filename
|
||||||
|
|
||||||
@@ -894,13 +1060,18 @@ def download_file(
|
|||||||
if not file_path:
|
if not file_path:
|
||||||
raise HTTPException(status_code=404, detail="Processed file not found")
|
raise HTTPException(status_code=404, detail="Processed file not found")
|
||||||
else:
|
else:
|
||||||
raise HTTPException(status_code=400, detail="Invalid version parameter. Use 'original' or 'processed'")
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail="Invalid version parameter. Use 'original' or 'processed'",
|
||||||
|
)
|
||||||
|
|
||||||
# Return the file with attachment disposition to trigger download
|
# Return the file with attachment disposition to trigger download
|
||||||
return FileResponse(
|
return FileResponse(
|
||||||
path=file_path,
|
path=file_path,
|
||||||
media_type=file_record.mime_type or "application/pdf",
|
media_type=file_record.mime_type or "application/pdf",
|
||||||
headers={"Content-Disposition": f'attachment; filename="{file_record.original_filename}"'},
|
headers={
|
||||||
|
"Content-Disposition": f'attachment; filename="{file_record.original_filename}"'
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
except HTTPException:
|
except HTTPException:
|
||||||
@@ -916,6 +1087,21 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
|||||||
"""Endpoint to accept a user-uploaded file and enqueue it for processing."""
|
"""Endpoint to accept a user-uploaded file and enqueue it for processing."""
|
||||||
workdir = settings.workdir
|
workdir = settings.workdir
|
||||||
|
|
||||||
|
# Early size check: reject before reading the body if Content-Length is known
|
||||||
|
max_size = settings.max_upload_size
|
||||||
|
content_length_header = request.headers.get("content-length")
|
||||||
|
if content_length_header is not None:
|
||||||
|
try:
|
||||||
|
declared_size = int(content_length_header)
|
||||||
|
if declared_size > max_size:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=413,
|
||||||
|
detail=f"File too large: declared size {declared_size} bytes exceeds maximum "
|
||||||
|
f"{max_size} bytes. See SECURITY_AUDIT.md for configuration details.",
|
||||||
|
)
|
||||||
|
except ValueError:
|
||||||
|
pass # Malformed header; proceed and check actual size after reading
|
||||||
|
|
||||||
# Extract just the filename without any path components to prevent path traversal
|
# Extract just the filename without any path components to prevent path traversal
|
||||||
# First, use basename to remove any directory components
|
# First, use basename to remove any directory components
|
||||||
base_filename = os.path.basename(file.filename)
|
base_filename = os.path.basename(file.filename)
|
||||||
@@ -934,27 +1120,37 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
|||||||
# Store both the safe original name and the unique name
|
# Store both the safe original name and the unique name
|
||||||
target_path = os.path.join(workdir, target_filename)
|
target_path = os.path.join(workdir, target_filename)
|
||||||
|
|
||||||
|
# Read file in chunks to avoid loading the entire body into memory at once,
|
||||||
|
# enforcing the size limit during the read so memory usage stays bounded.
|
||||||
try:
|
try:
|
||||||
|
written_size = 0
|
||||||
with open(target_path, "wb") as f:
|
with open(target_path, "wb") as f:
|
||||||
content = await file.read()
|
chunk_size = 65536 # 64 KB chunks
|
||||||
f.write(content)
|
while True:
|
||||||
|
chunk = await file.read(chunk_size)
|
||||||
|
if not chunk:
|
||||||
|
break
|
||||||
|
written_size += len(chunk)
|
||||||
|
if written_size > max_size:
|
||||||
|
# Exceeded limit mid-stream; clean up and reject
|
||||||
|
f.close()
|
||||||
|
os.remove(target_path)
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=413,
|
||||||
|
detail=f"File too large: exceeded {max_size} bytes during upload. "
|
||||||
|
f"See SECURITY_AUDIT.md for configuration details.",
|
||||||
|
)
|
||||||
|
f.write(chunk)
|
||||||
|
except HTTPException:
|
||||||
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
if os.path.exists(target_path):
|
||||||
|
os.remove(target_path)
|
||||||
raise HTTPException(status_code=500, detail=f"Failed to save file: {e}")
|
raise HTTPException(status_code=500, detail=f"Failed to save file: {e}")
|
||||||
|
|
||||||
# Log the mapping between original and safe filename
|
# Log the mapping between original and safe filename
|
||||||
logger.info(f"Saved uploaded file '{safe_filename}' as '{target_filename}'")
|
logger.info(f"Saved uploaded file '{safe_filename}' as '{target_filename}'")
|
||||||
|
file_size = written_size
|
||||||
# Check file size against configured maximum
|
|
||||||
file_size = os.path.getsize(target_path)
|
|
||||||
max_size = settings.max_upload_size
|
|
||||||
if file_size > max_size:
|
|
||||||
# Remove the file if it's too large
|
|
||||||
os.remove(target_path)
|
|
||||||
raise HTTPException(
|
|
||||||
status_code=413,
|
|
||||||
detail=f"File too large: {file_size} bytes (max {max_size} bytes). "
|
|
||||||
f"See SECURITY_AUDIT.md for configuration details.",
|
|
||||||
)
|
|
||||||
|
|
||||||
# Same set of allowed file types as in the IMAP task
|
# Same set of allowed file types as in the IMAP task
|
||||||
ALLOWED_MIME_TYPES = {
|
ALLOWED_MIME_TYPES = {
|
||||||
@@ -993,7 +1189,9 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
|||||||
# Check if file splitting is needed (only for PDFs)
|
# Check if file splitting is needed (only for PDFs)
|
||||||
from app.utils.file_splitting import should_split_file
|
from app.utils.file_splitting import should_split_file
|
||||||
|
|
||||||
should_split = is_pdf and should_split_file(target_path, settings.max_single_file_size)
|
should_split = is_pdf and should_split_file(
|
||||||
|
target_path, settings.max_single_file_size
|
||||||
|
)
|
||||||
|
|
||||||
if should_split:
|
if should_split:
|
||||||
# File needs to be split before processing
|
# File needs to be split before processing
|
||||||
@@ -1011,7 +1209,9 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
|||||||
task_ids = []
|
task_ids = []
|
||||||
for split_file in split_files:
|
for split_file in split_files:
|
||||||
split_filename = os.path.basename(split_file)
|
split_filename = os.path.basename(split_file)
|
||||||
task = process_document.delay(split_file, original_filename=split_filename)
|
task = process_document.delay(
|
||||||
|
split_file, original_filename=split_filename
|
||||||
|
)
|
||||||
task_ids.append(task.id)
|
task_ids.append(task.id)
|
||||||
logger.info(f"Enqueued split PDF part for processing: {split_file}")
|
logger.info(f"Enqueued split PDF part for processing: {split_file}")
|
||||||
|
|
||||||
@@ -1029,7 +1229,9 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Failed to split file {target_path}: {str(e)}")
|
logger.exception(f"Failed to split file {target_path}: {str(e)}")
|
||||||
# Fall back to processing the whole file
|
# Fall back to processing the whole file
|
||||||
logger.warning(f"Falling back to processing whole file due to split error: {str(e)}")
|
logger.warning(
|
||||||
|
f"Falling back to processing whole file due to split error: {str(e)}"
|
||||||
|
)
|
||||||
should_split = False
|
should_split = False
|
||||||
|
|
||||||
if is_pdf and not should_split:
|
if is_pdf and not should_split:
|
||||||
@@ -1037,21 +1239,37 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
|||||||
task = process_document.delay(target_path, original_filename=safe_filename)
|
task = process_document.delay(target_path, original_filename=safe_filename)
|
||||||
logger.info(f"Enqueued PDF for processing: {target_path}")
|
logger.info(f"Enqueued PDF for processing: {target_path}")
|
||||||
elif mime_type in IMAGE_MIME_TYPES or any(
|
elif mime_type in IMAGE_MIME_TYPES or any(
|
||||||
file_ext.endswith(ext) for ext in [".jpg", ".jpeg", ".png", ".gif", ".bmp", ".tiff", ".webp", ".svg"]
|
file_ext.endswith(ext)
|
||||||
|
for ext in [".jpg", ".jpeg", ".png", ".gif", ".bmp", ".tiff", ".webp", ".svg"]
|
||||||
):
|
):
|
||||||
# If it's an image, convert to PDF first
|
# If it's an image, convert to PDF first
|
||||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename)
|
task = convert_to_pdf.delay(target_path, original_filename=safe_filename)
|
||||||
logger.info(f"Enqueued image for PDF conversion: {target_path}")
|
logger.info(f"Enqueued image for PDF conversion: {target_path}")
|
||||||
elif mime_type in ALLOWED_MIME_TYPES or any(
|
elif mime_type in ALLOWED_MIME_TYPES or any(
|
||||||
file_ext.endswith(ext)
|
file_ext.endswith(ext)
|
||||||
for ext in [".doc", ".docx", ".xls", ".xlsx", ".ppt", ".pptx", ".odt", ".ods", ".odp", ".rtf", ".txt", ".csv"]
|
for ext in [
|
||||||
|
".doc",
|
||||||
|
".docx",
|
||||||
|
".xls",
|
||||||
|
".xlsx",
|
||||||
|
".ppt",
|
||||||
|
".pptx",
|
||||||
|
".odt",
|
||||||
|
".ods",
|
||||||
|
".odp",
|
||||||
|
".rtf",
|
||||||
|
".txt",
|
||||||
|
".csv",
|
||||||
|
]
|
||||||
):
|
):
|
||||||
# If it's an office document, convert to PDF first
|
# If it's an office document, convert to PDF first
|
||||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename)
|
task = convert_to_pdf.delay(target_path, original_filename=safe_filename)
|
||||||
logger.info(f"Enqueued office document for PDF conversion: {target_path}")
|
logger.info(f"Enqueued office document for PDF conversion: {target_path}")
|
||||||
else:
|
else:
|
||||||
# For any other file type, attempt conversion but log a warning
|
# For any other file type, attempt conversion but log a warning
|
||||||
logger.warning(f"Unsupported MIME type {mime_type} for {target_path}, attempting conversion")
|
logger.warning(
|
||||||
|
f"Unsupported MIME type {mime_type} for {target_path}, attempting conversion"
|
||||||
|
)
|
||||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename)
|
task = convert_to_pdf.delay(target_path, original_filename=safe_filename)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
|
|||||||
+65
-21
@@ -33,7 +33,9 @@ class Settings(BaseSettings):
|
|||||||
# Making Paperless optional
|
# Making Paperless optional
|
||||||
paperless_ngx_api_token: Optional[str] = None
|
paperless_ngx_api_token: Optional[str] = None
|
||||||
paperless_host: Optional[str] = None
|
paperless_host: Optional[str] = None
|
||||||
paperless_custom_field_absender: Optional[str] = None # Name of the "absender" custom field in Paperless
|
paperless_custom_field_absender: Optional[str] = (
|
||||||
|
None # Name of the "absender" custom field in Paperless
|
||||||
|
)
|
||||||
# JSON mapping of metadata field names to Paperless custom field names
|
# JSON mapping of metadata field names to Paperless custom field names
|
||||||
# Example: {"absender": "Sender", "empfaenger": "Recipient",
|
# Example: {"absender": "Sender", "empfaenger": "Recipient",
|
||||||
# "language": "Language", "correspondent": "Correspondent"}
|
# "language": "Language", "correspondent": "Correspondent"}
|
||||||
@@ -113,7 +115,9 @@ class Settings(BaseSettings):
|
|||||||
sftp_private_key_passphrase: Optional[str] = None
|
sftp_private_key_passphrase: Optional[str] = None
|
||||||
# Security: Host key verification is enabled by default for security
|
# Security: Host key verification is enabled by default for security
|
||||||
# In development/testing, set to True to disable verification (not recommended)
|
# In development/testing, set to True to disable verification (not recommended)
|
||||||
sftp_disable_host_key_verification: bool = False # Default enforces host key verification
|
sftp_disable_host_key_verification: bool = (
|
||||||
|
False # Default enforces host key verification
|
||||||
|
)
|
||||||
|
|
||||||
# Email settings
|
# Email settings
|
||||||
email_host: Optional[str] = None
|
email_host: Optional[str] = None
|
||||||
@@ -121,13 +125,17 @@ class Settings(BaseSettings):
|
|||||||
email_username: Optional[str] = None
|
email_username: Optional[str] = None
|
||||||
email_password: Optional[str] = None
|
email_password: Optional[str] = None
|
||||||
email_use_tls: bool = True
|
email_use_tls: bool = True
|
||||||
email_sender: Optional[str] = None # From address, defaults to email_username if not set
|
email_sender: Optional[str] = (
|
||||||
|
None # From address, defaults to email_username if not set
|
||||||
|
)
|
||||||
email_default_recipient: Optional[str] = None
|
email_default_recipient: Optional[str] = None
|
||||||
|
|
||||||
# OneDrive settings
|
# OneDrive settings
|
||||||
onedrive_client_id: Optional[str] = None
|
onedrive_client_id: Optional[str] = None
|
||||||
onedrive_client_secret: Optional[str] = None
|
onedrive_client_secret: Optional[str] = None
|
||||||
onedrive_tenant_id: Optional[str] = "common" # Default to "common" for personal accounts
|
onedrive_tenant_id: Optional[str] = (
|
||||||
|
"common" # Default to "common" for personal accounts
|
||||||
|
)
|
||||||
onedrive_refresh_token: Optional[str] = None # Required for personal accounts
|
onedrive_refresh_token: Optional[str] = None # Required for personal accounts
|
||||||
onedrive_folder_path: Optional[str] = None
|
onedrive_folder_path: Optional[str] = None
|
||||||
|
|
||||||
@@ -145,31 +153,43 @@ class Settings(BaseSettings):
|
|||||||
uptime_kuma_ping_interval: int = 5 # Default ping interval in minutes
|
uptime_kuma_ping_interval: int = 5 # Default ping interval in minutes
|
||||||
|
|
||||||
# HTTP request settings
|
# HTTP request settings
|
||||||
http_request_timeout: int = 120 # Default timeout for HTTP requests in seconds (handles large file operations)
|
http_request_timeout: int = (
|
||||||
|
120 # Default timeout for HTTP requests in seconds (handles large file operations)
|
||||||
|
)
|
||||||
|
|
||||||
# Feature flags
|
# Feature flags
|
||||||
allow_file_delete: bool = True # Default to allowing file deletion from database
|
allow_file_delete: bool = True # Default to allowing file deletion from database
|
||||||
|
|
||||||
# Batch processing settings
|
# Batch processing settings
|
||||||
processall_throttle_threshold: int = Field(
|
processall_throttle_threshold: int = Field(
|
||||||
default=20, description="Number of files above which throttling is applied in /processall endpoint"
|
default=20,
|
||||||
|
description="Number of files above which throttling is applied in /processall endpoint",
|
||||||
)
|
)
|
||||||
processall_throttle_delay: int = Field(
|
processall_throttle_delay: int = Field(
|
||||||
default=3, description="Delay in seconds between each task submission when throttling in /processall"
|
default=3,
|
||||||
|
description="Delay in seconds between each task submission when throttling in /processall",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Notification settings
|
# Notification settings
|
||||||
notification_urls: Union[List[str], str] = Field(
|
notification_urls: Union[List[str], str] = Field(
|
||||||
default_factory=list, description="List of Apprise notification URLs (e.g., discord://, telegram://, etc.)"
|
default_factory=list,
|
||||||
|
description="List of Apprise notification URLs (e.g., discord://, telegram://, etc.)",
|
||||||
|
)
|
||||||
|
notify_on_task_failure: bool = Field(
|
||||||
|
default=True, description="Send notifications when Celery tasks fail"
|
||||||
)
|
)
|
||||||
notify_on_task_failure: bool = Field(default=True, description="Send notifications when Celery tasks fail")
|
|
||||||
notify_on_credential_failure: bool = Field(
|
notify_on_credential_failure: bool = Field(
|
||||||
default=True, description="Send notifications when credential checks fail"
|
default=True, description="Send notifications when credential checks fail"
|
||||||
)
|
)
|
||||||
notify_on_startup: bool = Field(default=True, description="Send notifications when application starts")
|
notify_on_startup: bool = Field(
|
||||||
notify_on_shutdown: bool = Field(default=False, description="Send notifications when application shuts down")
|
default=True, description="Send notifications when application starts"
|
||||||
|
)
|
||||||
|
notify_on_shutdown: bool = Field(
|
||||||
|
default=False, description="Send notifications when application shuts down"
|
||||||
|
)
|
||||||
notify_on_file_processed: bool = Field(
|
notify_on_file_processed: bool = Field(
|
||||||
default=True, description="Send notifications when files are successfully processed"
|
default=True,
|
||||||
|
description="Send notifications when files are successfully processed",
|
||||||
)
|
)
|
||||||
|
|
||||||
# File upload size limits (for security - see SECURITY_AUDIT.md)
|
# File upload size limits (for security - see SECURITY_AUDIT.md)
|
||||||
@@ -184,6 +204,14 @@ class Settings(BaseSettings):
|
|||||||
" it will be split into smaller chunks for processing. Default: None (no splitting)."
|
" it will be split into smaller chunks for processing. Default: None (no splitting)."
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
max_request_body_size: int = Field(
|
||||||
|
default=1048576, # 1MB in bytes (1024 * 1024)
|
||||||
|
description=(
|
||||||
|
"Maximum request body size in bytes for non-file-upload requests. Default: 1MB."
|
||||||
|
" Prevents memory exhaustion attacks via oversized JSON/form payloads."
|
||||||
|
" File uploads are governed by MAX_UPLOAD_SIZE instead."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
# Deduplication settings - prevents processing of duplicate files
|
# Deduplication settings - prevents processing of duplicate files
|
||||||
enable_deduplication: bool = Field(
|
enable_deduplication: bool = Field(
|
||||||
@@ -228,7 +256,9 @@ class Settings(BaseSettings):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Content-Security-Policy (CSP) - Controls resource loading
|
# Content-Security-Policy (CSP) - Controls resource loading
|
||||||
security_header_csp_enabled: bool = Field(default=True, description="Enable CSP header.")
|
security_header_csp_enabled: bool = Field(
|
||||||
|
default=True, description="Enable CSP header."
|
||||||
|
)
|
||||||
security_header_csp_value: str = Field(
|
security_header_csp_value: str = Field(
|
||||||
default=(
|
default=(
|
||||||
"default-src 'self'; script-src 'self' 'unsafe-inline';"
|
"default-src 'self'; script-src 'self' 'unsafe-inline';"
|
||||||
@@ -238,14 +268,18 @@ class Settings(BaseSettings):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# X-Frame-Options - Prevents clickjacking
|
# X-Frame-Options - Prevents clickjacking
|
||||||
security_header_x_frame_options_enabled: bool = Field(default=True, description="Enable X-Frame-Options header.")
|
security_header_x_frame_options_enabled: bool = Field(
|
||||||
|
default=True, description="Enable X-Frame-Options header."
|
||||||
|
)
|
||||||
security_header_x_frame_options_value: str = Field(
|
security_header_x_frame_options_value: str = Field(
|
||||||
default="DENY", description="X-Frame-Options header value. Options: DENY, SAMEORIGIN, or ALLOW-FROM uri"
|
default="DENY",
|
||||||
|
description="X-Frame-Options header value. Options: DENY, SAMEORIGIN, or ALLOW-FROM uri",
|
||||||
)
|
)
|
||||||
|
|
||||||
# X-Content-Type-Options - Prevents MIME sniffing
|
# X-Content-Type-Options - Prevents MIME sniffing
|
||||||
security_header_x_content_type_options_enabled: bool = Field(
|
security_header_x_content_type_options_enabled: bool = Field(
|
||||||
default=True, description="Enable X-Content-Type-Options header (always set to 'nosniff')."
|
default=True,
|
||||||
|
description="Enable X-Content-Type-Options header (always set to 'nosniff').",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Audit Logging Configuration (see SECURITY_AUDIT.md – Infrastructure Security)
|
# Audit Logging Configuration (see SECURITY_AUDIT.md – Infrastructure Security)
|
||||||
@@ -301,7 +335,9 @@ class Settings(BaseSettings):
|
|||||||
if isinstance(data, dict):
|
if isinstance(data, dict):
|
||||||
for key, value in data.items():
|
for key, value in data.items():
|
||||||
if isinstance(value, str) and len(value) >= 2:
|
if isinstance(value, str) and len(value) >= 2:
|
||||||
if (value[0] == '"' and value[-1] == '"') or (value[0] == "'" and value[-1] == "'"):
|
if (value[0] == '"' and value[-1] == '"') or (
|
||||||
|
value[0] == "'" and value[-1] == "'"
|
||||||
|
):
|
||||||
data[key] = value[1:-1]
|
data[key] = value[1:-1]
|
||||||
return data
|
return data
|
||||||
|
|
||||||
@@ -336,7 +372,9 @@ class Settings(BaseSettings):
|
|||||||
return env_build_date
|
return env_build_date
|
||||||
|
|
||||||
# Then try to get build date from BUILD_DATE file
|
# Then try to get build date from BUILD_DATE file
|
||||||
build_date_file = os.path.join(os.path.dirname(os.path.dirname(__file__)), "BUILD_DATE")
|
build_date_file = os.path.join(
|
||||||
|
os.path.dirname(os.path.dirname(__file__)), "BUILD_DATE"
|
||||||
|
)
|
||||||
if os.path.exists(build_date_file):
|
if os.path.exists(build_date_file):
|
||||||
with open(build_date_file, "r") as f:
|
with open(build_date_file, "r") as f:
|
||||||
return f.read().strip()
|
return f.read().strip()
|
||||||
@@ -353,7 +391,9 @@ class Settings(BaseSettings):
|
|||||||
return env_version
|
return env_version
|
||||||
|
|
||||||
# Then try to get version from VERSION file
|
# Then try to get version from VERSION file
|
||||||
version_file = os.path.join(os.path.dirname(os.path.dirname(__file__)), "VERSION")
|
version_file = os.path.join(
|
||||||
|
os.path.dirname(os.path.dirname(__file__)), "VERSION"
|
||||||
|
)
|
||||||
if os.path.exists(version_file):
|
if os.path.exists(version_file):
|
||||||
with open(version_file, "r") as f:
|
with open(version_file, "r") as f:
|
||||||
return f.read().strip()
|
return f.read().strip()
|
||||||
@@ -370,7 +410,9 @@ class Settings(BaseSettings):
|
|||||||
return env_sha
|
return env_sha
|
||||||
|
|
||||||
# Then try to get from GIT_SHA file
|
# Then try to get from GIT_SHA file
|
||||||
git_sha_file = os.path.join(os.path.dirname(os.path.dirname(__file__)), "GIT_SHA")
|
git_sha_file = os.path.join(
|
||||||
|
os.path.dirname(os.path.dirname(__file__)), "GIT_SHA"
|
||||||
|
)
|
||||||
if os.path.exists(git_sha_file):
|
if os.path.exists(git_sha_file):
|
||||||
with open(git_sha_file, "r") as f:
|
with open(git_sha_file, "r") as f:
|
||||||
return f.read().strip()
|
return f.read().strip()
|
||||||
@@ -381,7 +423,9 @@ class Settings(BaseSettings):
|
|||||||
@property
|
@property
|
||||||
def runtime_info(self) -> str:
|
def runtime_info(self) -> str:
|
||||||
"""Get runtime information from file."""
|
"""Get runtime information from file."""
|
||||||
runtime_info_file = os.path.join(os.path.dirname(os.path.dirname(__file__)), "RUNTIME_INFO")
|
runtime_info_file = os.path.join(
|
||||||
|
os.path.dirname(os.path.dirname(__file__)), "RUNTIME_INFO"
|
||||||
|
)
|
||||||
if os.path.exists(runtime_info_file):
|
if os.path.exists(runtime_info_file):
|
||||||
with open(runtime_info_file, "r") as f:
|
with open(runtime_info_file, "r") as f:
|
||||||
return f.read().strip()
|
return f.read().strip()
|
||||||
|
|||||||
+33
-9
@@ -20,6 +20,7 @@ from app.config import settings
|
|||||||
from app.database import init_db
|
from app.database import init_db
|
||||||
from app.middleware.audit_log import AuditLogMiddleware
|
from app.middleware.audit_log import AuditLogMiddleware
|
||||||
from app.middleware.rate_limit import create_limiter, get_rate_limit_exceeded_handler
|
from app.middleware.rate_limit import create_limiter, get_rate_limit_exceeded_handler
|
||||||
|
from app.middleware.request_size_limit import RequestSizeLimitMiddleware
|
||||||
from app.middleware.security_headers import SecurityHeadersMiddleware
|
from app.middleware.security_headers import SecurityHeadersMiddleware
|
||||||
from app.utils.config_validator import check_all_configs
|
from app.utils.config_validator import check_all_configs
|
||||||
from app.utils.notification import init_apprise, notify_shutdown, notify_startup
|
from app.utils.notification import init_apprise, notify_shutdown, notify_startup
|
||||||
@@ -40,7 +41,8 @@ if settings.auth_enabled and not settings.session_secret:
|
|||||||
"Generate one with: python -c 'import secrets; print(secrets.token_hex(32))'"
|
"Generate one with: python -c 'import secrets; print(secrets.token_hex(32))'"
|
||||||
)
|
)
|
||||||
SESSION_SECRET = (
|
SESSION_SECRET = (
|
||||||
settings.session_secret or "INSECURE_DEFAULT_FOR_DEVELOPMENT_ONLY_DO_NOT_USE_IN_PRODUCTION_MINIMUM_32_CHARS"
|
settings.session_secret
|
||||||
|
or "INSECURE_DEFAULT_FOR_DEVELOPMENT_ONLY_DO_NOT_USE_IN_PRODUCTION_MINIMUM_32_CHARS"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -79,11 +81,15 @@ async def lifespan(app: FastAPI):
|
|||||||
len(issues) > 0 for provider, issues in config_issues["storage"].items()
|
len(issues) > 0 for provider, issues in config_issues["storage"].items()
|
||||||
)
|
)
|
||||||
if has_issues:
|
if has_issues:
|
||||||
logging.warning("Application started with configuration issues - some features may be unavailable")
|
logging.warning(
|
||||||
|
"Application started with configuration issues - some features may be unavailable"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
logging.info("Application started with valid configuration")
|
logging.info("Application started with valid configuration")
|
||||||
|
|
||||||
logging.info("Router organization: Using refactored API routers from app/api/ directory")
|
logging.info(
|
||||||
|
"Router organization: Using refactored API routers from app/api/ directory"
|
||||||
|
)
|
||||||
|
|
||||||
# Initialize notification system
|
# Initialize notification system
|
||||||
init_apprise()
|
init_apprise()
|
||||||
@@ -104,7 +110,9 @@ async def lifespan(app: FastAPI):
|
|||||||
app = FastAPI(title="DocuElevate", lifespan=lifespan)
|
app = FastAPI(title="DocuElevate", lifespan=lifespan)
|
||||||
|
|
||||||
# Initialize rate limiter and attach to app state
|
# Initialize rate limiter and attach to app state
|
||||||
limiter = create_limiter(redis_url=settings.redis_url, enabled=settings.rate_limiting_enabled)
|
limiter = create_limiter(
|
||||||
|
redis_url=settings.redis_url, enabled=settings.rate_limiting_enabled
|
||||||
|
)
|
||||||
app.state.limiter = limiter
|
app.state.limiter = limiter
|
||||||
app.add_exception_handler(RateLimitExceeded, get_rate_limit_exceeded_handler())
|
app.add_exception_handler(RateLimitExceeded, get_rate_limit_exceeded_handler())
|
||||||
|
|
||||||
@@ -116,6 +124,12 @@ app.add_exception_handler(RateLimitExceeded, get_rate_limit_exceeded_handler())
|
|||||||
# Set to False if reverse proxy (Traefik, Nginx) handles security headers
|
# Set to False if reverse proxy (Traefik, Nginx) handles security headers
|
||||||
app.add_middleware(SecurityHeadersMiddleware, config=settings)
|
app.add_middleware(SecurityHeadersMiddleware, config=settings)
|
||||||
|
|
||||||
|
# 2) Request Size Limit Middleware - enforces body size limits before reading
|
||||||
|
# MAX_REQUEST_BODY_SIZE: limit for non-file requests (default 1 MB)
|
||||||
|
# MAX_UPLOAD_SIZE: limit for multipart/form-data uploads (default 1 GB)
|
||||||
|
# See SECURITY_AUDIT.md – Code Security section
|
||||||
|
app.add_middleware(RequestSizeLimitMiddleware, config=settings)
|
||||||
|
|
||||||
# 2) Audit Logging Middleware - logs all requests with sensitive data masking
|
# 2) Audit Logging Middleware - logs all requests with sensitive data masking
|
||||||
# Configure via AUDIT_LOGGING_ENABLED environment variable
|
# Configure via AUDIT_LOGGING_ENABLED environment variable
|
||||||
# See SECURITY_AUDIT.md – Infrastructure Security section
|
# See SECURITY_AUDIT.md – Infrastructure Security section
|
||||||
@@ -128,14 +142,19 @@ app.add_middleware(SessionMiddleware, secret_key=SESSION_SECRET)
|
|||||||
app.add_middleware(ProxyHeadersMiddleware, trusted_hosts="*")
|
app.add_middleware(ProxyHeadersMiddleware, trusted_hosts="*")
|
||||||
|
|
||||||
# 5) Restrict valid hosts to prevent Host header attacks
|
# 5) Restrict valid hosts to prevent Host header attacks
|
||||||
app.add_middleware(TrustedHostMiddleware, allowed_hosts=[settings.external_hostname, "localhost", "127.0.0.1"])
|
app.add_middleware(
|
||||||
|
TrustedHostMiddleware,
|
||||||
|
allowed_hosts=[settings.external_hostname, "localhost", "127.0.0.1"],
|
||||||
|
)
|
||||||
|
|
||||||
# Mount the static files directory
|
# Mount the static files directory
|
||||||
static_dir = pathlib.Path(__file__).parents[1] / "frontend" / "static"
|
static_dir = pathlib.Path(__file__).parents[1] / "frontend" / "static"
|
||||||
if os.path.exists(static_dir):
|
if os.path.exists(static_dir):
|
||||||
app.mount("/static", StaticFiles(directory=str(static_dir)), name="static")
|
app.mount("/static", StaticFiles(directory=str(static_dir)), name="static")
|
||||||
else:
|
else:
|
||||||
print(f"WARNING: Static directory not found at {static_dir}. Static files will not be served.")
|
print(
|
||||||
|
f"WARNING: Static directory not found at {static_dir}. Static files will not be served."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# Custom exception handlers that return JSON for API routes and HTML for frontend routes
|
# Custom exception handlers that return JSON for API routes and HTML for frontend routes
|
||||||
@@ -154,7 +173,9 @@ async def http_exception_handler(request: Request, exc: HTTPException):
|
|||||||
|
|
||||||
# Handle 404 errors with a custom template
|
# Handle 404 errors with a custom template
|
||||||
if exc.status_code == 404:
|
if exc.status_code == 404:
|
||||||
return templates.TemplateResponse("404.html", {"request": request}, status_code=status.HTTP_404_NOT_FOUND)
|
return templates.TemplateResponse(
|
||||||
|
"404.html", {"request": request}, status_code=status.HTTP_404_NOT_FOUND
|
||||||
|
)
|
||||||
|
|
||||||
# For other HTTP errors, we could create specific templates or use a generic one
|
# For other HTTP errors, we could create specific templates or use a generic one
|
||||||
# For now, return a simple error page
|
# For now, return a simple error page
|
||||||
@@ -174,13 +195,16 @@ async def custom_500_handler(request: Request, exc: Exception):
|
|||||||
# For API routes, return JSON instead of HTML
|
# For API routes, return JSON instead of HTML
|
||||||
if request.url.path.startswith("/api/"):
|
if request.url.path.startswith("/api/"):
|
||||||
return JSONResponse(
|
return JSONResponse(
|
||||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, content={"detail": "Internal server error"}
|
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||||
|
content={"detail": "Internal server error"},
|
||||||
)
|
)
|
||||||
|
|
||||||
# Serve the 500 template for non-API routes
|
# Serve the 500 template for non-API routes
|
||||||
templates = Jinja2Templates(directory=str(static_dir.parent / "templates"))
|
templates = Jinja2Templates(directory=str(static_dir.parent / "templates"))
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
"500.html", {"request": request, "exc": exc}, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR
|
"500.html",
|
||||||
|
{"request": request, "exc": exc},
|
||||||
|
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,114 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
|
||||||
|
"""
|
||||||
|
Request Size Limit Middleware for DocuElevate.
|
||||||
|
|
||||||
|
This middleware enforces configurable size limits on incoming HTTP request bodies
|
||||||
|
to prevent memory exhaustion and Denial-of-Service (DoS) attacks.
|
||||||
|
|
||||||
|
Two independent limits are enforced:
|
||||||
|
- ``MAX_REQUEST_BODY_SIZE``: applied to all non-multipart requests (JSON, form data, etc.).
|
||||||
|
Default: 1 MB. Configurable via the ``MAX_REQUEST_BODY_SIZE`` environment variable.
|
||||||
|
- ``MAX_UPLOAD_SIZE``: applied to multipart/form-data (file upload) requests.
|
||||||
|
Default: 1 GB. Configurable via the ``MAX_UPLOAD_SIZE`` environment variable.
|
||||||
|
|
||||||
|
When a request exceeds the applicable limit the middleware immediately returns
|
||||||
|
``HTTP 413 Request Entity Too Large`` without reading the full body, which keeps
|
||||||
|
memory usage bounded.
|
||||||
|
|
||||||
|
See SECURITY_AUDIT.md – Code Security section for background.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
|
||||||
|
from fastapi import Request
|
||||||
|
from fastapi.responses import JSONResponse
|
||||||
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class RequestSizeLimitMiddleware(BaseHTTPMiddleware):
|
||||||
|
"""
|
||||||
|
Middleware that rejects requests whose body exceeds a configured size limit.
|
||||||
|
|
||||||
|
File-upload requests (``Content-Type: multipart/form-data``) are checked
|
||||||
|
against ``config.max_upload_size``; all other requests are checked against
|
||||||
|
``config.max_request_body_size``.
|
||||||
|
|
||||||
|
The check is performed on the ``Content-Length`` header before the body is
|
||||||
|
read, so oversized requests are rejected without buffering the payload into
|
||||||
|
memory. If the client omits the ``Content-Length`` header the request is
|
||||||
|
passed through to the normal handler (where endpoint-level checks still
|
||||||
|
apply for file uploads).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, app, config):
|
||||||
|
"""
|
||||||
|
Initialize the middleware.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
app: The ASGI application to wrap.
|
||||||
|
config: Application settings object with ``max_request_body_size``
|
||||||
|
and ``max_upload_size`` attributes.
|
||||||
|
"""
|
||||||
|
super().__init__(app)
|
||||||
|
self.max_body_size = config.max_request_body_size
|
||||||
|
self.max_upload_size = config.max_upload_size
|
||||||
|
logger.info(
|
||||||
|
f"Request size limit middleware enabled – "
|
||||||
|
f"body limit: {self.max_body_size} bytes, "
|
||||||
|
f"upload limit: {self.max_upload_size} bytes"
|
||||||
|
)
|
||||||
|
|
||||||
|
async def dispatch(self, request: Request, call_next):
|
||||||
|
"""
|
||||||
|
Check the ``Content-Length`` header and reject oversized requests early.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
request: Incoming HTTP request.
|
||||||
|
call_next: Next middleware or route handler.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
HTTP 413 response if the request is too large, otherwise the
|
||||||
|
downstream response.
|
||||||
|
"""
|
||||||
|
content_length_header = request.headers.get("content-length")
|
||||||
|
if content_length_header is not None:
|
||||||
|
try:
|
||||||
|
content_length = int(content_length_header)
|
||||||
|
except ValueError:
|
||||||
|
# Malformed header – let downstream handle it
|
||||||
|
return await call_next(request)
|
||||||
|
|
||||||
|
content_type = request.headers.get("content-type", "")
|
||||||
|
is_multipart = "multipart/form-data" in content_type
|
||||||
|
|
||||||
|
if is_multipart:
|
||||||
|
limit = self.max_upload_size
|
||||||
|
limit_description = "file upload"
|
||||||
|
config_var = "MAX_UPLOAD_SIZE"
|
||||||
|
else:
|
||||||
|
limit = self.max_body_size
|
||||||
|
limit_description = "request body"
|
||||||
|
config_var = "MAX_REQUEST_BODY_SIZE"
|
||||||
|
|
||||||
|
if content_length > limit:
|
||||||
|
logger.warning(
|
||||||
|
f"Rejected oversized {limit_description}: "
|
||||||
|
f"{content_length} bytes > {limit} bytes limit "
|
||||||
|
f"(configure with {config_var})"
|
||||||
|
)
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=413,
|
||||||
|
content={
|
||||||
|
"detail": (
|
||||||
|
f"Request body too large: {content_length} bytes "
|
||||||
|
f"(maximum allowed: {limit} bytes). "
|
||||||
|
f"Adjust the {config_var} environment variable to change this limit. "
|
||||||
|
f"See SECURITY_AUDIT.md for details."
|
||||||
|
)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
return await call_next(request)
|
||||||
@@ -39,12 +39,14 @@ Control how the `/processall` endpoint handles large batches of files to prevent
|
|||||||
|---------------------------|--------------------------------------------------------------------------------------------------------------|---------------|
|
|---------------------------|--------------------------------------------------------------------------------------------------------------|---------------|
|
||||||
| `MAX_UPLOAD_SIZE` | Maximum file upload size in bytes. Files exceeding this limit are rejected. | `1073741824` (1GB) |
|
| `MAX_UPLOAD_SIZE` | Maximum file upload size in bytes. Files exceeding this limit are rejected. | `1073741824` (1GB) |
|
||||||
| `MAX_SINGLE_FILE_SIZE` | Optional: Maximum size for a single file chunk in bytes. Files exceeding this are split into smaller parts. | `None` (no splitting) |
|
| `MAX_SINGLE_FILE_SIZE` | Optional: Maximum size for a single file chunk in bytes. Files exceeding this are split into smaller parts. | `None` (no splitting) |
|
||||||
|
| `MAX_REQUEST_BODY_SIZE` | Maximum request body size in bytes for non-file-upload requests (JSON, form data, etc.). File uploads use `MAX_UPLOAD_SIZE` instead. | `1048576` (1MB) |
|
||||||
|
|
||||||
**Configuration Examples:**
|
**Configuration Examples:**
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
# Default: Allow up to 1GB uploads, no splitting
|
# Default: Allow up to 1GB uploads, no splitting, 1MB JSON/form body limit
|
||||||
MAX_UPLOAD_SIZE=1073741824
|
MAX_UPLOAD_SIZE=1073741824
|
||||||
|
MAX_REQUEST_BODY_SIZE=1048576
|
||||||
|
|
||||||
# Conservative: 100MB max, split files over 50MB
|
# Conservative: 100MB max, split files over 50MB
|
||||||
MAX_UPLOAD_SIZE=104857600
|
MAX_UPLOAD_SIZE=104857600
|
||||||
|
|||||||
+152
-46
@@ -32,17 +32,24 @@ def mock_celery_tasks():
|
|||||||
mock_process_task.delay.return_value = mock_task
|
mock_process_task.delay.return_value = mock_task
|
||||||
mock_convert_task.delay.return_value = mock_task
|
mock_convert_task.delay.return_value = mock_task
|
||||||
|
|
||||||
yield {"process_document": mock_process_task.delay, "convert_to_pdf": mock_convert_task.delay}
|
yield {
|
||||||
|
"process_document": mock_process_task.delay,
|
||||||
|
"convert_to_pdf": mock_convert_task.delay,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
class TestValidFileUploads:
|
class TestValidFileUploads:
|
||||||
"""Tests for successful file uploads with various valid file types."""
|
"""Tests for successful file uploads with various valid file types."""
|
||||||
|
|
||||||
def test_upload_valid_pdf(self, client: TestClient, sample_pdf_path: str, mock_celery_tasks):
|
def test_upload_valid_pdf(
|
||||||
|
self, client: TestClient, sample_pdf_path: str, mock_celery_tasks
|
||||||
|
):
|
||||||
"""Test uploading a valid PDF file."""
|
"""Test uploading a valid PDF file."""
|
||||||
with open(sample_pdf_path, "rb") as f:
|
with open(sample_pdf_path, "rb") as f:
|
||||||
response = client.post("/api/ui-upload", files={"file": ("document.pdf", f, "application/pdf")})
|
response = client.post(
|
||||||
|
"/api/ui-upload", files={"file": ("document.pdf", f, "application/pdf")}
|
||||||
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
@@ -67,7 +74,8 @@ class TestValidFileUploads:
|
|||||||
"""Test uploading a valid text file."""
|
"""Test uploading a valid text file."""
|
||||||
text_content = b"This is a test text file.\nWith multiple lines."
|
text_content = b"This is a test text file.\nWith multiple lines."
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": ("document.txt", io.BytesIO(text_content), "text/plain")}
|
"/api/ui-upload",
|
||||||
|
files={"file": ("document.txt", io.BytesIO(text_content), "text/plain")},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -87,7 +95,10 @@ class TestValidFileUploads:
|
|||||||
b"\xff\xd9"
|
b"\xff\xd9"
|
||||||
)
|
)
|
||||||
|
|
||||||
response = client.post("/api/ui-upload", files={"file": ("image.jpg", io.BytesIO(jpeg_content), "image/jpeg")})
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
files={"file": ("image.jpg", io.BytesIO(jpeg_content), "image/jpeg")},
|
||||||
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
@@ -106,7 +117,8 @@ class TestValidFileUploads:
|
|||||||
)
|
)
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": ("screenshot.png", io.BytesIO(png_content), "image/png")}
|
"/api/ui-upload",
|
||||||
|
files={"file": ("screenshot.png", io.BytesIO(png_content), "image/png")},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -143,7 +155,10 @@ class TestValidFileUploads:
|
|||||||
"""Test uploading a CSV file."""
|
"""Test uploading a CSV file."""
|
||||||
csv_content = b"name,age,city\nJohn,30,NYC\nJane,25,LA\n"
|
csv_content = b"name,age,city\nJohn,30,NYC\nJane,25,LA\n"
|
||||||
|
|
||||||
response = client.post("/api/ui-upload", files={"file": ("data.csv", io.BytesIO(csv_content), "text/csv")})
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
files={"file": ("data.csv", io.BytesIO(csv_content), "text/csv")},
|
||||||
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
@@ -159,15 +174,17 @@ class TestInvalidFileUploads:
|
|||||||
"""Test that files exceeding MAX_UPLOAD_SIZE are rejected."""
|
"""Test that files exceeding MAX_UPLOAD_SIZE are rejected."""
|
||||||
from app.config import settings
|
from app.config import settings
|
||||||
|
|
||||||
# Create a large file content (mock it to avoid memory issues)
|
# Temporarily lower the upload limit so a tiny file exceeds it,
|
||||||
large_content = b"x" * 1024 # 1KB for testing
|
# avoiding the need to allocate a real 1 GB payload in memory.
|
||||||
|
small_limit = 100 # 100 bytes
|
||||||
with patch("os.path.getsize") as mock_getsize:
|
small_content = b"x" * (small_limit + 1)
|
||||||
# Mock the file size to be over the configured limit
|
|
||||||
mock_getsize.return_value = settings.max_upload_size + 1
|
|
||||||
|
|
||||||
|
with patch.object(settings, "max_upload_size", small_limit):
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": ("huge.pdf", io.BytesIO(large_content), "application/pdf")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": ("huge.pdf", io.BytesIO(small_content), "application/pdf")
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 413 # Request Entity Too Large
|
assert response.status_code == 413 # Request Entity Too Large
|
||||||
@@ -179,7 +196,14 @@ class TestInvalidFileUploads:
|
|||||||
exe_content = b"MZ\x90\x00" # PE header
|
exe_content = b"MZ\x90\x00" # PE header
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": ("program.exe", io.BytesIO(exe_content), "application/x-msdownload")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": (
|
||||||
|
"program.exe",
|
||||||
|
io.BytesIO(exe_content),
|
||||||
|
"application/x-msdownload",
|
||||||
|
)
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
# Per the code, unsupported types get a warning but are still processed
|
# Per the code, unsupported types get a warning but are still processed
|
||||||
@@ -189,7 +213,10 @@ class TestInvalidFileUploads:
|
|||||||
|
|
||||||
def test_upload_empty_file(self, client: TestClient, mock_celery_tasks):
|
def test_upload_empty_file(self, client: TestClient, mock_celery_tasks):
|
||||||
"""Test uploading an empty file."""
|
"""Test uploading an empty file."""
|
||||||
response = client.post("/api/ui-upload", files={"file": ("empty.txt", io.BytesIO(b""), "text/plain")})
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
files={"file": ("empty.txt", io.BytesIO(b""), "text/plain")},
|
||||||
|
)
|
||||||
|
|
||||||
# Empty files are accepted and queued for processing
|
# Empty files are accepted and queued for processing
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -209,14 +236,19 @@ class TestInvalidFileUploads:
|
|||||||
class TestUploadSecurity:
|
class TestUploadSecurity:
|
||||||
"""Tests for security aspects of file uploads."""
|
"""Tests for security aspects of file uploads."""
|
||||||
|
|
||||||
def test_path_traversal_prevention_dotdot(self, client: TestClient, mock_celery_tasks):
|
def test_path_traversal_prevention_dotdot(
|
||||||
|
self, client: TestClient, mock_celery_tasks
|
||||||
|
):
|
||||||
"""Test that path traversal attempts are prevented."""
|
"""Test that path traversal attempts are prevented."""
|
||||||
# Try to upload a file with path traversal in filename
|
# Try to upload a file with path traversal in filename
|
||||||
malicious_filename = "../../etc/passwd.pdf"
|
malicious_filename = "../../etc/passwd.pdf"
|
||||||
pdf_content = b"%PDF-1.4\n%EOF"
|
pdf_content = b"%PDF-1.4\n%EOF"
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": (malicious_filename, io.BytesIO(pdf_content), "application/pdf")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": (malicious_filename, io.BytesIO(pdf_content), "application/pdf")
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -227,13 +259,18 @@ class TestUploadSecurity:
|
|||||||
assert ".." not in data["stored_filename"]
|
assert ".." not in data["stored_filename"]
|
||||||
assert "/" not in data["stored_filename"]
|
assert "/" not in data["stored_filename"]
|
||||||
|
|
||||||
def test_path_traversal_prevention_absolute(self, client: TestClient, mock_celery_tasks):
|
def test_path_traversal_prevention_absolute(
|
||||||
|
self, client: TestClient, mock_celery_tasks
|
||||||
|
):
|
||||||
"""Test that absolute path attempts are prevented."""
|
"""Test that absolute path attempts are prevented."""
|
||||||
malicious_filename = "/etc/shadow.pdf"
|
malicious_filename = "/etc/shadow.pdf"
|
||||||
pdf_content = b"%PDF-1.4\n%EOF"
|
pdf_content = b"%PDF-1.4\n%EOF"
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": (malicious_filename, io.BytesIO(pdf_content), "application/pdf")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": (malicious_filename, io.BytesIO(pdf_content), "application/pdf")
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -243,13 +280,18 @@ class TestUploadSecurity:
|
|||||||
assert data["original_filename"] == "shadow.pdf"
|
assert data["original_filename"] == "shadow.pdf"
|
||||||
assert not data["stored_filename"].startswith("/")
|
assert not data["stored_filename"].startswith("/")
|
||||||
|
|
||||||
def test_filename_with_special_characters(self, client: TestClient, mock_celery_tasks):
|
def test_filename_with_special_characters(
|
||||||
|
self, client: TestClient, mock_celery_tasks
|
||||||
|
):
|
||||||
"""Test handling of filenames with special characters."""
|
"""Test handling of filenames with special characters."""
|
||||||
special_filename = "file name with spaces & special!@#chars.pdf"
|
special_filename = "file name with spaces & special!@#chars.pdf"
|
||||||
pdf_content = b"%PDF-1.4\n%EOF"
|
pdf_content = b"%PDF-1.4\n%EOF"
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": (special_filename, io.BytesIO(pdf_content), "application/pdf")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": (special_filename, io.BytesIO(pdf_content), "application/pdf")
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -265,14 +307,19 @@ class TestUploadSecurity:
|
|||||||
# Stored filename should have UUID and extension
|
# Stored filename should have UUID and extension
|
||||||
assert data["stored_filename"].endswith(".pdf")
|
assert data["stored_filename"].endswith(".pdf")
|
||||||
|
|
||||||
def test_path_traversal_prevention_windows_style(self, client: TestClient, mock_celery_tasks):
|
def test_path_traversal_prevention_windows_style(
|
||||||
|
self, client: TestClient, mock_celery_tasks
|
||||||
|
):
|
||||||
"""Test that Windows-style path traversal attempts are prevented."""
|
"""Test that Windows-style path traversal attempts are prevented."""
|
||||||
# Try Windows-style path with backslashes
|
# Try Windows-style path with backslashes
|
||||||
malicious_filename = "..\\..\\..\\windows\\system32\\config.pdf"
|
malicious_filename = "..\\..\\..\\windows\\system32\\config.pdf"
|
||||||
pdf_content = b"%PDF-1.4\n%EOF"
|
pdf_content = b"%PDF-1.4\n%EOF"
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": (malicious_filename, io.BytesIO(pdf_content), "application/pdf")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": (malicious_filename, io.BytesIO(pdf_content), "application/pdf")
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -291,13 +338,18 @@ class TestUploadSecurity:
|
|||||||
assert ".." not in data["original_filename"]
|
assert ".." not in data["original_filename"]
|
||||||
assert "/" not in data["original_filename"]
|
assert "/" not in data["original_filename"]
|
||||||
|
|
||||||
def test_path_traversal_prevention_mixed_separators(self, client: TestClient, mock_celery_tasks):
|
def test_path_traversal_prevention_mixed_separators(
|
||||||
|
self, client: TestClient, mock_celery_tasks
|
||||||
|
):
|
||||||
"""Test handling of filenames with mixed path separators."""
|
"""Test handling of filenames with mixed path separators."""
|
||||||
malicious_filename = "../path\\to/file.pdf"
|
malicious_filename = "../path\\to/file.pdf"
|
||||||
pdf_content = b"%PDF-1.4\n%EOF"
|
pdf_content = b"%PDF-1.4\n%EOF"
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": (malicious_filename, io.BytesIO(pdf_content), "application/pdf")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": (malicious_filename, io.BytesIO(pdf_content), "application/pdf")
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -323,7 +375,10 @@ class TestUploadErrorHandling:
|
|||||||
pdf_content = b"%PDF-1.4\n%EOF"
|
pdf_content = b"%PDF-1.4\n%EOF"
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": ("test.pdf", io.BytesIO(pdf_content), "application/pdf")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": ("test.pdf", io.BytesIO(pdf_content), "application/pdf")
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 500
|
assert response.status_code == 500
|
||||||
@@ -331,14 +386,21 @@ class TestUploadErrorHandling:
|
|||||||
|
|
||||||
def test_upload_celery_task_failure(self, client: TestClient, mock_celery_tasks):
|
def test_upload_celery_task_failure(self, client: TestClient, mock_celery_tasks):
|
||||||
"""Test handling when Celery task queueing fails."""
|
"""Test handling when Celery task queueing fails."""
|
||||||
mock_celery_tasks["process_document"].side_effect = Exception("Celery connection failed")
|
mock_celery_tasks["process_document"].side_effect = Exception(
|
||||||
|
"Celery connection failed"
|
||||||
|
)
|
||||||
|
|
||||||
pdf_content = b"%PDF-1.4\n%EOF"
|
pdf_content = b"%PDF-1.4\n%EOF"
|
||||||
|
|
||||||
# The endpoint should still handle the error gracefully
|
# The endpoint should still handle the error gracefully
|
||||||
# In this case, the exception will propagate
|
# In this case, the exception will propagate
|
||||||
with pytest.raises(Exception):
|
with pytest.raises(Exception):
|
||||||
client.post("/api/ui-upload", files={"file": ("test.pdf", io.BytesIO(pdf_content), "application/pdf")})
|
client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": ("test.pdf", io.BytesIO(pdf_content), "application/pdf")
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
@@ -351,11 +413,13 @@ class TestUploadFilenameHandling:
|
|||||||
|
|
||||||
# Upload same file twice
|
# Upload same file twice
|
||||||
response1 = client.post(
|
response1 = client.post(
|
||||||
"/api/ui-upload", files={"file": ("same.pdf", io.BytesIO(pdf_content), "application/pdf")}
|
"/api/ui-upload",
|
||||||
|
files={"file": ("same.pdf", io.BytesIO(pdf_content), "application/pdf")},
|
||||||
)
|
)
|
||||||
|
|
||||||
response2 = client.post(
|
response2 = client.post(
|
||||||
"/api/ui-upload", files={"file": ("same.pdf", io.BytesIO(pdf_content), "application/pdf")}
|
"/api/ui-upload",
|
||||||
|
files={"file": ("same.pdf", io.BytesIO(pdf_content), "application/pdf")},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response1.status_code == 200
|
assert response1.status_code == 200
|
||||||
@@ -375,14 +439,20 @@ class TestUploadFilenameHandling:
|
|||||||
content = b"Some content"
|
content = b"Some content"
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": ("NOEXTENSION", io.BytesIO(content), "application/octet-stream")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": ("NOEXTENSION", io.BytesIO(content), "application/octet-stream")
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
assert data["original_filename"] == "NOEXTENSION"
|
assert data["original_filename"] == "NOEXTENSION"
|
||||||
# Stored filename should be just the UUID without extension
|
# Stored filename should be just the UUID without extension
|
||||||
assert "." not in data["stored_filename"] or data["stored_filename"].count(".") == 0
|
assert (
|
||||||
|
"." not in data["stored_filename"]
|
||||||
|
or data["stored_filename"].count(".") == 0
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
@@ -394,7 +464,10 @@ class TestUploadMimeTypeDetection:
|
|||||||
pdf_content = b"%PDF-1.4\n%EOF"
|
pdf_content = b"%PDF-1.4\n%EOF"
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": ("doc.pdf", io.BytesIO(pdf_content), "application/octet-stream")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": ("doc.pdf", io.BytesIO(pdf_content), "application/octet-stream")
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -407,7 +480,14 @@ class TestUploadMimeTypeDetection:
|
|||||||
image_content = b"\x00\x01\x02\x03"
|
image_content = b"\x00\x01\x02\x03"
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": ("photo.jpg", io.BytesIO(image_content), "application/octet-stream")}
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": (
|
||||||
|
"photo.jpg",
|
||||||
|
io.BytesIO(image_content),
|
||||||
|
"application/octet-stream",
|
||||||
|
)
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -419,7 +499,9 @@ class TestUploadMimeTypeDetection:
|
|||||||
class TestFileSplitting:
|
class TestFileSplitting:
|
||||||
"""Tests for file splitting functionality when MAX_SINGLE_FILE_SIZE is configured."""
|
"""Tests for file splitting functionality when MAX_SINGLE_FILE_SIZE is configured."""
|
||||||
|
|
||||||
def test_pdf_splitting_when_configured(self, client: TestClient, sample_pdf_path: str, mock_celery_tasks):
|
def test_pdf_splitting_when_configured(
|
||||||
|
self, client: TestClient, sample_pdf_path: str, mock_celery_tasks
|
||||||
|
):
|
||||||
"""Test that PDFs are split when they exceed MAX_SINGLE_FILE_SIZE."""
|
"""Test that PDFs are split when they exceed MAX_SINGLE_FILE_SIZE."""
|
||||||
from app.config import settings
|
from app.config import settings
|
||||||
|
|
||||||
@@ -434,9 +516,14 @@ class TestFileSplitting:
|
|||||||
]
|
]
|
||||||
|
|
||||||
# Mock should_split_file to return True
|
# Mock should_split_file to return True
|
||||||
with patch("app.utils.file_splitting.should_split_file", return_value=True):
|
with patch(
|
||||||
|
"app.utils.file_splitting.should_split_file", return_value=True
|
||||||
|
):
|
||||||
with open(sample_pdf_path, "rb") as f:
|
with open(sample_pdf_path, "rb") as f:
|
||||||
response = client.post("/api/ui-upload", files={"file": ("large.pdf", f, "application/pdf")})
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
files={"file": ("large.pdf", f, "application/pdf")},
|
||||||
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
@@ -452,14 +539,19 @@ class TestFileSplitting:
|
|||||||
# Verify each split file was queued for processing
|
# Verify each split file was queued for processing
|
||||||
assert mock_celery_tasks["process_document"].call_count == 3
|
assert mock_celery_tasks["process_document"].call_count == 3
|
||||||
|
|
||||||
def test_no_splitting_when_not_configured(self, client: TestClient, sample_pdf_path: str, mock_celery_tasks):
|
def test_no_splitting_when_not_configured(
|
||||||
|
self, client: TestClient, sample_pdf_path: str, mock_celery_tasks
|
||||||
|
):
|
||||||
"""Test that PDFs are not split when MAX_SINGLE_FILE_SIZE is None."""
|
"""Test that PDFs are not split when MAX_SINGLE_FILE_SIZE is None."""
|
||||||
from app.config import settings
|
from app.config import settings
|
||||||
|
|
||||||
# Ensure max_single_file_size is None (default)
|
# Ensure max_single_file_size is None (default)
|
||||||
with patch.object(settings, "max_single_file_size", None):
|
with patch.object(settings, "max_single_file_size", None):
|
||||||
with open(sample_pdf_path, "rb") as f:
|
with open(sample_pdf_path, "rb") as f:
|
||||||
response = client.post("/api/ui-upload", files={"file": ("document.pdf", f, "application/pdf")})
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
files={"file": ("document.pdf", f, "application/pdf")},
|
||||||
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
@@ -472,14 +564,19 @@ class TestFileSplitting:
|
|||||||
# Verify file was processed directly without splitting
|
# Verify file was processed directly without splitting
|
||||||
mock_celery_tasks["process_document"].assert_called_once()
|
mock_celery_tasks["process_document"].assert_called_once()
|
||||||
|
|
||||||
def test_no_splitting_for_small_files(self, client: TestClient, sample_pdf_path: str, mock_celery_tasks):
|
def test_no_splitting_for_small_files(
|
||||||
|
self, client: TestClient, sample_pdf_path: str, mock_celery_tasks
|
||||||
|
):
|
||||||
"""Test that small PDFs are not split even when MAX_SINGLE_FILE_SIZE is configured."""
|
"""Test that small PDFs are not split even when MAX_SINGLE_FILE_SIZE is configured."""
|
||||||
from app.config import settings
|
from app.config import settings
|
||||||
|
|
||||||
# Configure a very large limit
|
# Configure a very large limit
|
||||||
with patch.object(settings, "max_single_file_size", 1000000000): # 1GB limit
|
with patch.object(settings, "max_single_file_size", 1000000000): # 1GB limit
|
||||||
with open(sample_pdf_path, "rb") as f:
|
with open(sample_pdf_path, "rb") as f:
|
||||||
response = client.post("/api/ui-upload", files={"file": ("small.pdf", f, "application/pdf")})
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
files={"file": ("small.pdf", f, "application/pdf")},
|
||||||
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
@@ -491,16 +588,24 @@ class TestFileSplitting:
|
|||||||
# Verify file was processed directly
|
# Verify file was processed directly
|
||||||
mock_celery_tasks["process_document"].assert_called_once()
|
mock_celery_tasks["process_document"].assert_called_once()
|
||||||
|
|
||||||
def test_splitting_fallback_on_error(self, client: TestClient, sample_pdf_path: str, mock_celery_tasks):
|
def test_splitting_fallback_on_error(
|
||||||
|
self, client: TestClient, sample_pdf_path: str, mock_celery_tasks
|
||||||
|
):
|
||||||
"""Test that if splitting fails, the file is processed as a whole."""
|
"""Test that if splitting fails, the file is processed as a whole."""
|
||||||
from app.config import settings
|
from app.config import settings
|
||||||
|
|
||||||
with patch.object(settings, "max_single_file_size", 100): # Small limit
|
with patch.object(settings, "max_single_file_size", 100): # Small limit
|
||||||
with patch("app.utils.file_splitting.should_split_file", return_value=True):
|
with patch("app.utils.file_splitting.should_split_file", return_value=True):
|
||||||
# Mock split_pdf_by_size to raise an exception
|
# Mock split_pdf_by_size to raise an exception
|
||||||
with patch("app.utils.file_splitting.split_pdf_by_size", side_effect=Exception("Split failed")):
|
with patch(
|
||||||
|
"app.utils.file_splitting.split_pdf_by_size",
|
||||||
|
side_effect=Exception("Split failed"),
|
||||||
|
):
|
||||||
with open(sample_pdf_path, "rb") as f:
|
with open(sample_pdf_path, "rb") as f:
|
||||||
response = client.post("/api/ui-upload", files={"file": ("document.pdf", f, "application/pdf")})
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
files={"file": ("document.pdf", f, "application/pdf")},
|
||||||
|
)
|
||||||
|
|
||||||
# Should still succeed, falling back to processing the whole file
|
# Should still succeed, falling back to processing the whole file
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -526,7 +631,8 @@ class TestFileSplitting:
|
|||||||
)
|
)
|
||||||
|
|
||||||
response = client.post(
|
response = client.post(
|
||||||
"/api/ui-upload", files={"file": ("image.png", io.BytesIO(image_content), "image/png")}
|
"/api/ui-upload",
|
||||||
|
files={"file": ("image.png", io.BytesIO(image_content), "image/png")},
|
||||||
)
|
)
|
||||||
|
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
|
|||||||
@@ -0,0 +1,180 @@
|
|||||||
|
"""
|
||||||
|
Tests for the RequestSizeLimitMiddleware.
|
||||||
|
|
||||||
|
Validates that:
|
||||||
|
- Non-file requests exceeding MAX_REQUEST_BODY_SIZE are rejected with HTTP 413
|
||||||
|
- Multipart/form-data uploads exceeding MAX_UPLOAD_SIZE are rejected with HTTP 413
|
||||||
|
- Requests within the limits pass through normally
|
||||||
|
- Missing Content-Length header does not cause false rejections
|
||||||
|
"""
|
||||||
|
|
||||||
|
import io
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
class TestRequestSizeLimitMiddleware:
|
||||||
|
"""Unit tests for the RequestSizeLimitMiddleware dispatch logic."""
|
||||||
|
|
||||||
|
def test_middleware_rejects_oversized_json_body(self, client: TestClient):
|
||||||
|
"""Non-file request with Content-Length exceeding MAX_REQUEST_BODY_SIZE is rejected."""
|
||||||
|
from app.config import settings
|
||||||
|
|
||||||
|
oversized = settings.max_request_body_size + 1
|
||||||
|
response = client.post(
|
||||||
|
"/api/process-url",
|
||||||
|
content=b"x" * 10, # actual body doesn't matter; header is checked first
|
||||||
|
headers={
|
||||||
|
"Content-Length": str(oversized),
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert response.status_code == 413
|
||||||
|
detail = response.json()["detail"]
|
||||||
|
assert "MAX_REQUEST_BODY_SIZE" in detail
|
||||||
|
|
||||||
|
def test_middleware_allows_request_within_json_limit(self, client: TestClient):
|
||||||
|
"""Non-file request with Content-Length within limit is not rejected by middleware."""
|
||||||
|
from app.config import settings
|
||||||
|
|
||||||
|
# Send a body within limit; the endpoint may return 4xx for its own reasons,
|
||||||
|
# but the middleware must NOT return 413.
|
||||||
|
small = settings.max_request_body_size - 1
|
||||||
|
response = client.post(
|
||||||
|
"/api/process-url",
|
||||||
|
content=b"{}",
|
||||||
|
headers={"Content-Length": str(small), "Content-Type": "application/json"},
|
||||||
|
)
|
||||||
|
# The endpoint may return 400/422 (bad JSON or auth), but NOT 413 from middleware
|
||||||
|
assert response.status_code != 413
|
||||||
|
|
||||||
|
def test_middleware_rejects_oversized_multipart_upload(self, client: TestClient):
|
||||||
|
"""Multipart upload with Content-Length exceeding MAX_UPLOAD_SIZE is rejected."""
|
||||||
|
from app.config import settings
|
||||||
|
|
||||||
|
oversized = settings.max_upload_size + 1
|
||||||
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
content=b"x" * 10,
|
||||||
|
headers={
|
||||||
|
"Content-Length": str(oversized),
|
||||||
|
"Content-Type": "multipart/form-data; boundary=boundary",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert response.status_code == 413
|
||||||
|
detail = response.json()["detail"]
|
||||||
|
assert "MAX_UPLOAD_SIZE" in detail
|
||||||
|
|
||||||
|
def test_middleware_allows_multipart_within_upload_limit(self, client: TestClient):
|
||||||
|
"""Multipart upload with Content-Length within MAX_UPLOAD_SIZE passes middleware."""
|
||||||
|
from app.config import settings
|
||||||
|
|
||||||
|
# A Content-Length within the upload limit should NOT be rejected by the middleware.
|
||||||
|
# The endpoint itself will reject because the body is not a real multipart payload.
|
||||||
|
within_limit = min(1024, settings.max_upload_size - 1)
|
||||||
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
content=b"x" * 10,
|
||||||
|
headers={
|
||||||
|
"Content-Length": str(within_limit),
|
||||||
|
"Content-Type": "multipart/form-data; boundary=boundary",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
# Not rejected by middleware (may be 400/422 from endpoint)
|
||||||
|
assert response.status_code != 413
|
||||||
|
|
||||||
|
def test_middleware_allows_request_without_content_length(self, client: TestClient):
|
||||||
|
"""Requests without Content-Length header pass through middleware (no false rejection)."""
|
||||||
|
# Remove Content-Length header entirely; middleware must not reject
|
||||||
|
response = client.get("/api/files")
|
||||||
|
# May get 200/401/403 but not 413
|
||||||
|
assert response.status_code != 413
|
||||||
|
|
||||||
|
def test_middleware_error_message_contains_limit_and_config_hint(
|
||||||
|
self, client: TestClient
|
||||||
|
):
|
||||||
|
"""413 response body contains limit details and config variable name."""
|
||||||
|
from app.config import settings
|
||||||
|
|
||||||
|
oversized = settings.max_request_body_size + 1
|
||||||
|
response = client.post(
|
||||||
|
"/api/process-url",
|
||||||
|
content=b"{}",
|
||||||
|
headers={
|
||||||
|
"Content-Length": str(oversized),
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert response.status_code == 413
|
||||||
|
detail = response.json()["detail"]
|
||||||
|
assert str(settings.max_request_body_size) in detail
|
||||||
|
assert "SECURITY_AUDIT.md" in detail
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
class TestFileUploadSizeLimitStreaming:
|
||||||
|
"""Integration tests for streaming size enforcement in the ui-upload endpoint."""
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def mock_celery(self):
|
||||||
|
with (
|
||||||
|
patch("app.api.files.process_document") as mock_proc,
|
||||||
|
patch("app.api.files.convert_to_pdf") as mock_conv,
|
||||||
|
):
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
|
task = MagicMock()
|
||||||
|
task.id = "test-task-id"
|
||||||
|
mock_proc.delay.return_value = task
|
||||||
|
mock_conv.delay.return_value = task
|
||||||
|
yield
|
||||||
|
|
||||||
|
def test_upload_rejected_when_content_length_declared_too_large(
|
||||||
|
self, client: TestClient
|
||||||
|
):
|
||||||
|
"""Upload is rejected early via Content-Length check before reading data."""
|
||||||
|
from app.config import settings
|
||||||
|
|
||||||
|
oversized = settings.max_upload_size + 1
|
||||||
|
pdf_data = b"%PDF-1.4\n%EOF"
|
||||||
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
content=pdf_data,
|
||||||
|
headers={
|
||||||
|
"Content-Length": str(oversized),
|
||||||
|
"Content-Type": "multipart/form-data; boundary=boundary",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert response.status_code == 413
|
||||||
|
|
||||||
|
def test_upload_rejected_mid_stream_when_data_exceeds_limit(
|
||||||
|
self, client: TestClient
|
||||||
|
):
|
||||||
|
"""Upload is rejected mid-stream when actual data exceeds max_upload_size."""
|
||||||
|
from app.config import settings
|
||||||
|
|
||||||
|
# Temporarily reduce max_upload_size to a tiny value for this test
|
||||||
|
small_limit = 100 # 100 bytes
|
||||||
|
with patch.object(settings, "max_upload_size", small_limit):
|
||||||
|
large_content = b"x" * (small_limit + 1)
|
||||||
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
files={
|
||||||
|
"file": ("big.pdf", io.BytesIO(large_content), "application/pdf")
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert response.status_code == 413
|
||||||
|
assert "too large" in response.json()["detail"].lower()
|
||||||
|
|
||||||
|
def test_upload_succeeds_within_size_limit(self, client: TestClient):
|
||||||
|
"""Small, valid file upload completes successfully within size limits."""
|
||||||
|
pdf_content = b"%PDF-1.4\n1 0 obj\n<</Type /Catalog>>\nendobj\n%%EOF"
|
||||||
|
response = client.post(
|
||||||
|
"/api/ui-upload",
|
||||||
|
files={"file": ("small.pdf", io.BytesIO(pdf_content), "application/pdf")},
|
||||||
|
)
|
||||||
|
assert response.status_code == 200
|
||||||
|
assert response.json()["status"] == "queued"
|
||||||
Reference in New Issue
Block a user