85 lines
3.0 KiB
Python
85 lines
3.0 KiB
Python
import time
|
|
|
|
from django.conf import settings
|
|
from django.utils import timezone
|
|
|
|
from .algorithms import ProcessingError, decode_image, histogram, process_image
|
|
from .models import ImageSession, ProcessingJob
|
|
from .storage import load_image_array, payload_for_image, save_image_array
|
|
from .tasks import run_batch_job
|
|
|
|
|
|
def image_session_create(*, uploaded_file=None, image_base64=None):
|
|
if uploaded_file and uploaded_file.size > settings.MAX_UPLOAD_MB * 1024 * 1024:
|
|
raise ProcessingError(f"Upload exceeds {settings.MAX_UPLOAD_MB} MB.")
|
|
|
|
image = decode_image(uploaded_file=uploaded_file, base64_image=image_base64)
|
|
relative_path = save_image_array(image, "original")
|
|
hist = histogram(image)
|
|
session = ImageSession.objects.create(
|
|
original_image=relative_path,
|
|
width=image.shape[1],
|
|
height=image.shape[0],
|
|
channels=image.shape[2] if image.ndim == 3 else 1,
|
|
color_mode="RGB" if image.ndim == 3 else "L",
|
|
original_histogram=hist,
|
|
expires_at=timezone.now() + timezone.timedelta(hours=settings.IMAGE_SESSION_TTL_HOURS),
|
|
)
|
|
payload = {
|
|
"session_id": str(session.id),
|
|
"width": session.width,
|
|
"height": session.height,
|
|
"channels": session.channels,
|
|
"color_mode": session.color_mode,
|
|
"original_histogram": hist,
|
|
"expires_at": session.expires_at.isoformat(),
|
|
}
|
|
payload.update(payload_for_image(image, relative_path))
|
|
return payload
|
|
|
|
|
|
def image_session_process(*, session, operation, params):
|
|
if session.expired:
|
|
raise ProcessingError("Image session has expired.")
|
|
|
|
started_at = time.perf_counter()
|
|
source = load_image_array(session.original_image)
|
|
result = process_image(source, operation, params)
|
|
relative_path = save_image_array(result, f"processed-{operation}")
|
|
hist = histogram(result)
|
|
|
|
session.processed_image = relative_path
|
|
session.processed_histogram = hist
|
|
session.save(update_fields=["processed_image", "processed_histogram"])
|
|
|
|
payload = {
|
|
"session_id": str(session.id),
|
|
"operation": operation,
|
|
"params": params,
|
|
"processed_histogram": hist,
|
|
"elapsed_ms": round((time.perf_counter() - started_at) * 1000, 2),
|
|
}
|
|
payload.update(payload_for_image(result, relative_path))
|
|
return payload
|
|
|
|
|
|
def batch_job_create(*, operation, session_ids, params=None):
|
|
job = ProcessingJob.objects.create(operation=operation, params=params or {})
|
|
run_batch_job.delay(str(job.id), operation, [str(session_id) for session_id in session_ids])
|
|
return job
|
|
|
|
|
|
def processing_job_payload(*, job):
|
|
payload = {
|
|
"job_id": str(job.id),
|
|
"operation": job.operation,
|
|
"status": job.status,
|
|
"progress": job.progress,
|
|
"error": job.error,
|
|
"result_histogram": job.result_histogram,
|
|
}
|
|
if job.result_image:
|
|
image = load_image_array(job.result_image)
|
|
payload.update(payload_for_image(image, job.result_image))
|
|
return payload
|