added duplicate check and database
This commit is contained in:
+7
-2
@@ -1,14 +1,19 @@
|
|||||||
|
# app/database.py
|
||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
|
|
||||||
from sqlalchemy import create_engine, Column, String, Integer
|
from sqlalchemy import create_engine
|
||||||
from sqlalchemy.ext.declarative import declarative_base
|
from sqlalchemy.ext.declarative import declarative_base
|
||||||
from sqlalchemy.orm import sessionmaker
|
from sqlalchemy.orm import sessionmaker
|
||||||
from .config import settings
|
from app.config import settings
|
||||||
|
|
||||||
Base = declarative_base()
|
Base = declarative_base()
|
||||||
engine = create_engine(settings.database_url, connect_args={"check_same_thread": False})
|
engine = create_engine(settings.database_url, connect_args={"check_same_thread": False})
|
||||||
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
|
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
|
||||||
|
|
||||||
|
def init_db():
|
||||||
|
"""Call this once (e.g. on startup) to create tables if they don't exist."""
|
||||||
|
Base.metadata.create_all(bind=engine)
|
||||||
|
|
||||||
def get_db():
|
def get_db():
|
||||||
db = SessionLocal()
|
db = SessionLocal()
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
import os
|
import os
|
||||||
from fastapi import FastAPI, HTTPException, UploadFile, File
|
from fastapi import FastAPI, HTTPException, UploadFile, File
|
||||||
|
from app.database import init_db
|
||||||
from app.config import settings
|
from app.config import settings
|
||||||
from app.tasks.upload_to_s3 import upload_to_s3
|
from app.tasks.upload_to_s3 import upload_to_s3
|
||||||
from app.tasks.upload_to_dropbox import upload_to_dropbox
|
from app.tasks.upload_to_dropbox import upload_to_dropbox
|
||||||
@@ -12,6 +13,10 @@ from app.frontend import router as frontend_router
|
|||||||
|
|
||||||
app = FastAPI(title="Document Processing API")
|
app = FastAPI(title="Document Processing API")
|
||||||
|
|
||||||
|
@app.on_event("startup")
|
||||||
|
def on_startup():
|
||||||
|
init_db() # Create tables if they don't exist
|
||||||
|
|
||||||
@app.get("/")
|
@app.get("/")
|
||||||
def root():
|
def root():
|
||||||
return {"message": "Document Processing API"}
|
return {"message": "Document Processing API"}
|
||||||
|
|||||||
+36
-2
@@ -1,7 +1,9 @@
|
|||||||
|
# app/models.py
|
||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
|
|
||||||
from .database import Base
|
from sqlalchemy import Column, String, Integer, DateTime, func
|
||||||
from sqlalchemy import Column, String, Integer
|
from sqlalchemy.ext.declarative import declarative_base
|
||||||
|
from app.database import Base
|
||||||
|
|
||||||
class DocumentMetadata(Base):
|
class DocumentMetadata(Base):
|
||||||
__tablename__ = "documents"
|
__tablename__ = "documents"
|
||||||
@@ -12,3 +14,35 @@ class DocumentMetadata(Base):
|
|||||||
recipient = Column(String)
|
recipient = Column(String)
|
||||||
tags = Column(String)
|
tags = Column(String)
|
||||||
summary = Column(String)
|
summary = Column(String)
|
||||||
|
|
||||||
|
class FileRecord(Base):
|
||||||
|
__tablename__ = "files"
|
||||||
|
|
||||||
|
id = Column(Integer, primary_key=True, index=True)
|
||||||
|
|
||||||
|
# Hash of the file content (e.g. SHA-256)
|
||||||
|
filehash = Column(String, unique=True, index=True, nullable=False)
|
||||||
|
|
||||||
|
# The name of the file as it was originally uploaded (if known)
|
||||||
|
original_filename = Column(String)
|
||||||
|
|
||||||
|
# The name/path we store on disk (e.g. /workdir/tmp/<uuid>.pdf)
|
||||||
|
local_filename = Column(String, nullable=False)
|
||||||
|
|
||||||
|
# Size of the file in bytes
|
||||||
|
file_size = Column(Integer, nullable=False)
|
||||||
|
|
||||||
|
# MIME type or extension (optional)
|
||||||
|
mime_type = Column(String)
|
||||||
|
|
||||||
|
# Timestamp when we inserted this record
|
||||||
|
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||||
|
|
||||||
|
class ProcessingLog(Base):
|
||||||
|
__tablename__ = "processing_logs"
|
||||||
|
id = Column(Integer, primary_key=True, index=True)
|
||||||
|
file_id = Column(Integer, ForeignKey("files.id"))
|
||||||
|
step_name = Column(String) # e.g. "OCR", "convert_to_pdf", "upload_s3"
|
||||||
|
status = Column(String) # "success" / "failure"
|
||||||
|
message = Column(String) # error text or success note
|
||||||
|
timestamp = Column(DateTime(timezone=True), server_default=func.now())
|
||||||
|
|||||||
+56
-11
@@ -4,15 +4,20 @@ import os
|
|||||||
import uuid
|
import uuid
|
||||||
import boto3
|
import boto3
|
||||||
import shutil
|
import shutil
|
||||||
|
import mimetypes
|
||||||
import fitz # PyMuPDF for checking embedded text
|
import fitz # PyMuPDF for checking embedded text
|
||||||
|
|
||||||
from app.config import settings
|
from app.config import settings
|
||||||
from app.tasks.retry_config import BaseTaskWithRetry
|
from app.tasks.retry_config import BaseTaskWithRetry
|
||||||
from app.tasks.process_with_textract import process_with_textract
|
from app.tasks.process_with_textract import process_with_textract
|
||||||
from app.tasks.extract_metadata_with_gpt import extract_metadata_with_gpt
|
from app.tasks.extract_metadata_with_gpt import extract_metadata_with_gpt
|
||||||
|
|
||||||
# Import the shared Celery instance
|
|
||||||
from app.celery_app import celery
|
from app.celery_app import celery
|
||||||
|
|
||||||
|
# NEW imports for the DB
|
||||||
|
from app.database import SessionLocal
|
||||||
|
from app.models import FileRecord
|
||||||
|
from app.utils import hash_file
|
||||||
|
|
||||||
# Initialize S3 client
|
# Initialize S3 client
|
||||||
s3_client = boto3.client(
|
s3_client = boto3.client(
|
||||||
"s3",
|
"s3",
|
||||||
@@ -21,13 +26,20 @@ s3_client = boto3.client(
|
|||||||
region_name=settings.aws_region,
|
region_name=settings.aws_region,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@celery.task(base=BaseTaskWithRetry)
|
@celery.task(base=BaseTaskWithRetry)
|
||||||
def upload_to_s3(original_local_file: str):
|
def upload_to_s3(original_local_file: str):
|
||||||
"""
|
"""
|
||||||
Uploads a file to S3 with a UUID-based filename and triggers processing.
|
Uploads a file to S3 with a UUID-based filename and triggers processing.
|
||||||
- If the PDF already contains embedded text, skip Textract and extract text locally.
|
|
||||||
- Otherwise, upload to S3 and process with Textract.
|
Steps:
|
||||||
|
1. Check if we have a FileRecord entry (via SHA-256 hash). If found, skip re-processing.
|
||||||
|
2. If not found, insert a new DB row and continue with the pipeline:
|
||||||
|
- Copy file to /workdir/tmp
|
||||||
|
- Check for embedded text. If present, skip S3 and run local GPT extraction
|
||||||
|
- Otherwise, upload to S3 and queue Textract-based OCR
|
||||||
"""
|
"""
|
||||||
|
|
||||||
bucket_name = settings.s3_bucket_name
|
bucket_name = settings.s3_bucket_name
|
||||||
if not bucket_name:
|
if not bucket_name:
|
||||||
print("[ERROR] S3 bucket name not set.")
|
print("[ERROR] S3 bucket name not set.")
|
||||||
@@ -37,22 +49,54 @@ def upload_to_s3(original_local_file: str):
|
|||||||
print(f"[ERROR] File {original_local_file} not found.")
|
print(f"[ERROR] File {original_local_file} not found.")
|
||||||
return {"error": "File not found"}
|
return {"error": "File not found"}
|
||||||
|
|
||||||
# Generate UUID and create a new filename
|
# 0. Compute the file hash and check for duplicates
|
||||||
file_ext = os.path.splitext(original_local_file)[1] # Preserve original file extension
|
filehash = hash_file(original_local_file)
|
||||||
|
original_filename = os.path.basename(original_local_file)
|
||||||
|
file_size = os.path.getsize(original_local_file)
|
||||||
|
mime_type, _ = mimetypes.guess_type(original_local_file)
|
||||||
|
if not mime_type:
|
||||||
|
mime_type = "application/octet-stream"
|
||||||
|
|
||||||
|
# Acquire DB session in the task
|
||||||
|
with SessionLocal() as db:
|
||||||
|
existing = db.query(FileRecord).filter_by(filehash=filehash).one_or_none()
|
||||||
|
if existing:
|
||||||
|
print(f"[INFO] Duplicate file detected (hash={filehash[:10]}...) Skipping processing.")
|
||||||
|
return {
|
||||||
|
"status": "duplicate_file",
|
||||||
|
"file_id": existing.id,
|
||||||
|
"detail": "File already processed."
|
||||||
|
}
|
||||||
|
|
||||||
|
# Not a duplicate -> insert a new record
|
||||||
|
new_record = FileRecord(
|
||||||
|
filehash=filehash,
|
||||||
|
original_filename=original_filename,
|
||||||
|
local_filename="", # Will fill in after we move it
|
||||||
|
file_size=file_size,
|
||||||
|
mime_type=mime_type,
|
||||||
|
)
|
||||||
|
db.add(new_record)
|
||||||
|
db.commit()
|
||||||
|
db.refresh(new_record)
|
||||||
|
|
||||||
|
# 1. Generate a UUID-based filename and place it in /workdir/tmp
|
||||||
|
file_ext = os.path.splitext(original_local_file)[1]
|
||||||
file_uuid = str(uuid.uuid4())
|
file_uuid = str(uuid.uuid4())
|
||||||
new_filename = f"{file_uuid}{file_ext}"
|
new_filename = f"{file_uuid}{file_ext}"
|
||||||
|
|
||||||
# Construct the new local path using settings.workdir and a 'tmp' subdirectory
|
|
||||||
tmp_dir = os.path.join(settings.workdir, "tmp")
|
tmp_dir = os.path.join(settings.workdir, "tmp")
|
||||||
new_local_path = os.path.join(tmp_dir, new_filename)
|
|
||||||
|
|
||||||
# Ensure the target tmp directory exists
|
|
||||||
os.makedirs(tmp_dir, exist_ok=True)
|
os.makedirs(tmp_dir, exist_ok=True)
|
||||||
|
new_local_path = os.path.join(tmp_dir, new_filename)
|
||||||
|
|
||||||
# Copy the file instead of moving it
|
# Copy the file instead of moving it
|
||||||
shutil.copy(original_local_file, new_local_path)
|
shutil.copy(original_local_file, new_local_path)
|
||||||
|
|
||||||
# Check for embedded text
|
# Update the DB with final local filename
|
||||||
|
new_record.local_filename = new_local_path
|
||||||
|
db.commit()
|
||||||
|
|
||||||
|
# 2. Check for embedded text (outside the DB session to avoid long open transactions)
|
||||||
pdf_doc = fitz.open(new_local_path)
|
pdf_doc = fitz.open(new_local_path)
|
||||||
has_text = any(page.get_text() for page in pdf_doc)
|
has_text = any(page.get_text() for page in pdf_doc)
|
||||||
pdf_doc.close()
|
pdf_doc.close()
|
||||||
@@ -72,6 +116,7 @@ def upload_to_s3(original_local_file: str):
|
|||||||
|
|
||||||
return {"file": new_local_path, "status": "Text extracted locally"}
|
return {"file": new_local_path, "status": "Text extracted locally"}
|
||||||
|
|
||||||
|
# 3. If no embedded text, upload to S3 and queue Textract processing
|
||||||
try:
|
try:
|
||||||
print(f"[INFO] Uploading {new_local_path} to s3://{bucket_name}/{new_filename}...")
|
print(f"[INFO] Uploading {new_local_path} to s3://{bucket_name}/{new_filename}...")
|
||||||
s3_client.upload_file(new_local_path, bucket_name, new_filename)
|
s3_client.upload_file(new_local_path, bucket_name, new_filename)
|
||||||
|
|||||||
@@ -0,0 +1,16 @@
|
|||||||
|
# app/utils.py
|
||||||
|
import hashlib
|
||||||
|
|
||||||
|
def hash_file(filepath, chunk_size=65536):
|
||||||
|
"""
|
||||||
|
Returns the SHA-256 hash of the file at 'filepath'.
|
||||||
|
Reads the file in chunks to handle large files efficiently.
|
||||||
|
"""
|
||||||
|
sha256 = hashlib.sha256()
|
||||||
|
with open(filepath, "rb") as f:
|
||||||
|
while True:
|
||||||
|
data = f.read(chunk_size)
|
||||||
|
if not data:
|
||||||
|
break
|
||||||
|
sha256.update(data)
|
||||||
|
return sha256.hexdigest()
|
||||||
Reference in New Issue
Block a user