Merge branch 'main' into auth-for-ui

This commit is contained in:
Christian Krakau-Louis
2025-03-25 14:49:19 +01:00
committed by GitHub
5 changed files with 171 additions and 23 deletions
+52 -4
View File
@@ -1,15 +1,63 @@
#!/usr/bin/env python3
# app/database.py
from sqlalchemy import create_engine, Column, String, Integer
import os
import logging
from sqlalchemy import create_engine, exc
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import sessionmaker
from .config import settings
from sqlalchemy.engine.url import make_url
from app.config import settings
logger = logging.getLogger(__name__)
Base = declarative_base()
engine = create_engine(settings.database_url, connect_args={"check_same_thread": False})
# Parse the DATABASE_URL
DB_URL = settings.database_url
engine = create_engine(DB_URL, connect_args={"check_same_thread": False})
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
def init_db():
"""
Ensures the SQLite database file and its parent directory exist (if using sqlite).
Then runs Base.metadata.create_all(bind=engine) to initialize tables.
Logs a message if a new SQLite DB file is created.
"""
# 1. Parse the DB URL to see if it's sqlite
url = make_url(DB_URL)
if url.get_backend_name() == "sqlite":
# 2. Extract the database path from the URL
database_path = url.database # e.g. "/workdir/db/database.db" or ":memory:"
if database_path != ":memory:":
# 3. Ensure directory exists
db_dir = os.path.dirname(database_path)
if db_dir and not os.path.exists(db_dir):
logger.info(f"Creating directory for SQLite DB: {db_dir}")
os.makedirs(db_dir, exist_ok=True)
# 4. If the file does not exist, create an empty one
if not os.path.exists(database_path):
logger.info(f"Creating new SQLite database file at {database_path}")
open(database_path, "a").close()
# 5. Now create tables if they don't exist yet
try:
Base.metadata.create_all(bind=engine)
logger.info("Database initialization complete (tables created if not exist).")
except exc.SQLAlchemyError as e:
logger.error(f"Error initializing database: {e}")
raise
def get_db():
"""
Dependency for FastAPI routes or general DB usage.
Yields a SQLAlchemy session, and closes it upon exit.
"""
db = SessionLocal()
try:
yield db
+6 -1
View File
@@ -6,7 +6,7 @@ from starlette.middleware.sessions import SessionMiddleware
from starlette.config import Config
from starlette.middleware.trustedhost import TrustedHostMiddleware
from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware
from app.database import init_db
from app.config import settings
from app.tasks.upload_to_s3 import upload_to_s3
from app.tasks.upload_to_dropbox import upload_to_dropbox
@@ -25,6 +25,7 @@ SESSION_SECRET = config(
app = FastAPI(title="Document Processing API")
# 1) Session Middleware (for request.session to work)
app.add_middleware(SessionMiddleware, secret_key=SESSION_SECRET)
@@ -39,6 +40,10 @@ app.add_middleware(TrustedHostMiddleware, allowed_hosts=[
"127.0.0.1"
])
@app.on_event("startup")
def on_startup():
init_db() # Create tables if they don't exist
@app.get("/")
def root():
return {"message": "Document Processing API"}
+36 -2
View File
@@ -1,7 +1,9 @@
# app/models.py
#!/usr/bin/env python3
from .database import Base
from sqlalchemy import Column, String, Integer
from sqlalchemy import Column, String, Integer, DateTime, func, ForeignKey
from sqlalchemy.ext.declarative import declarative_base
from app.database import Base
class DocumentMetadata(Base):
__tablename__ = "documents"
@@ -12,3 +14,35 @@ class DocumentMetadata(Base):
recipient = Column(String)
tags = 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())
+61 -16
View File
@@ -4,15 +4,20 @@ import os
import uuid
import boto3
import shutil
import mimetypes
import fitz # PyMuPDF for checking embedded text
from app.config import settings
from app.tasks.retry_config import BaseTaskWithRetry
from app.tasks.process_with_textract import process_with_textract
from app.tasks.extract_metadata_with_gpt import extract_metadata_with_gpt
# Import the shared Celery instance
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
s3_client = boto3.client(
"s3",
@@ -21,13 +26,20 @@ s3_client = boto3.client(
region_name=settings.aws_region,
)
@celery.task(base=BaseTaskWithRetry)
def upload_to_s3(original_local_file: str):
"""
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
if not bucket_name:
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.")
return {"error": "File not found"}
# Generate UUID and create a new filename
file_ext = os.path.splitext(original_local_file)[1] # Preserve original file extension
file_uuid = str(uuid.uuid4())
new_filename = f"{file_uuid}{file_ext}"
# 0. Compute the file hash and check for duplicates
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"
# Construct the new local path using settings.workdir and a 'tmp' subdirectory
tmp_dir = os.path.join(settings.workdir, "tmp")
new_local_path = os.path.join(tmp_dir, new_filename)
# 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."
}
# Ensure the target tmp directory exists
os.makedirs(tmp_dir, exist_ok=True)
# 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)
# Copy the file instead of moving it
shutil.copy(original_local_file, new_local_path)
# 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())
new_filename = f"{file_uuid}{file_ext}"
# Check for embedded text
tmp_dir = os.path.join(settings.workdir, "tmp")
os.makedirs(tmp_dir, exist_ok=True)
new_local_path = os.path.join(tmp_dir, new_filename)
# Copy the file instead of moving it
shutil.copy(original_local_file, new_local_path)
# 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)
has_text = any(page.get_text() for page in pdf_doc)
pdf_doc.close()
@@ -72,6 +116,7 @@ def upload_to_s3(original_local_file: str):
return {"file": new_local_path, "status": "Text extracted locally"}
# 3. If no embedded text, upload to S3 and queue Textract processing
try:
print(f"[INFO] Uploading {new_local_path} to s3://{bucket_name}/{new_filename}...")
s3_client.upload_file(new_local_path, bucket_name, new_filename)
+16
View File
@@ -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()