Team Ai
Apppublic

MasoomChoudhury/processor

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
batch_manager.py141 linesDownload Raw Back to root
1import asyncio2import logging3from datetime import datetime, timezone4from supabase import AsyncClient5from api_clients import get_gemini_ocr, get_gemini_analysis, get_claude_analysis6 7# Configure logging8logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')9 10class BatchManager:11    def __init__(self, supabase: AsyncClient, app_status: dict):12        self.supabase = supabase13        self.app_status = app_status14        self.current_batch = []15        self.timer = None16        self.timer_lock = asyncio.Lock()17 18    async def add_job_to_batch(self, job):19        """Adds a new job to the current batch and resets the processing timer."""20        async with self.timer_lock:21            self.current_batch.append(job)22            23            # Update application status24            self.app_status["last_job_received_at"] = datetime.now(timezone.utc).isoformat()25            self.app_status["jobs_in_current_batch"] = len(self.current_batch)26            self.app_status["current_batch_status"] = "collecting"27            28            logging.info(f"Job {job['id']} added to batch. Current batch size: {len(self.current_batch)}")29 30            # Reset the 7-second timer31            if self.timer and not self.timer.done():32                self.timer.cancel()33            34            self.timer = asyncio.create_task(self._schedule_batch_processing())35 36    async def _schedule_batch_processing(self):37        """Waits for the timer to expire, then processes the batch."""38        try:39            await asyncio.sleep(7)40            logging.info("Batch timer expired. Processing batch...")41            await self.process_batch()42        except asyncio.CancelledError:43            logging.info("Batch timer reset.")44 45    async def process_batch(self):46        """Processes all jobs currently in the batch."""47        async with self.timer_lock:48            if not self.current_batch:49                self.app_status["current_batch_status"] = "idle"50                return51 52            batch_to_process = self.current_batch53            self.current_batch = []54            55            self.app_status["current_batch_status"] = f"processing {len(batch_to_process)} jobs"56            self.app_status["jobs_in_current_batch"] = 057            logging.info(f"Processing batch of {len(batch_to_process)} jobs.")58 59        try:60            # Create a record for the batch in the database61            batch_insert_result = await self.supabase.table("batches").insert({62                "image_count": len(batch_to_process),63                "status": "processing"64            }).execute()65            batch_id = batch_insert_result.data[0]['id']66 67            # Link jobs to the new batch68            for job in batch_to_process:69                await self.supabase.table("jobs").update({"batch_id": batch_id, "status": "ocr_processing"}).eq("id", job['id']).execute()70 71        except Exception as e:72            logging.error(f"Error creating batch record: {e}")73            self.app_status["current_batch_status"] = f"failed: {e}"74            return75 76        # --- Perform OCR and AI Analysis ---77        ocr_tasks = [self.perform_ocr(job) for job in batch_to_process]78        ocr_results = await asyncio.gather(*ocr_tasks, return_exceptions=True)79 80        analysis_tasks = []81        for job, ocr_text in zip(batch_to_process, ocr_results):82            if isinstance(ocr_text, str) and ocr_text:83                job_with_ocr = {**job, "ocr_text": ocr_text}84                analysis_tasks.append(self._perform_ai_analysis(job_with_ocr))85 86        await asyncio.gather(*analysis_tasks, return_exceptions=True)87 88        # --- Finalize Batch ---89        try:90            await self.supabase.table("batches").update({91                "status": "completed",92                "processed_at": datetime.now(timezone.utc).isoformat()93            }).eq("id", batch_id).execute()94            self.app_status["total_jobs_processed"] += len(batch_to_process)95            logging.info(f"Batch {batch_id} completed successfully.")96        except Exception as e:97            logging.error(f"Error marking batch as completed: {e}")98        99        self.app_status["current_batch_status"] = "idle"100 101 102    async def perform_ocr(self, job):103        """Downloads an image, performs OCR, and updates the job status."""104        try:105            logging.info(f"Starting OCR for job {job['id']}...")106            image_response = await self.supabase.storage.from_('images-bucket').download(job['image_path'])107            108            ocr_text = await get_gemini_ocr(image_response)109 110            await self.supabase.table("jobs").update({"ocr_text": ocr_text, "status": "ai_processing"}).eq("id", job['id']).execute()111            logging.info(f"OCR successful for job {job['id']}.")112            return ocr_text113        except Exception as e:114            logging.error(f"Error during OCR for job {job['id']}: {e}")115            await self.supabase.table("jobs").update({"status": "failed"}).eq("id", job['id']).execute()116            return e117 118 119    async def _perform_ai_analysis(self, job):120        """Performs AI analysis from Gemini and Claude concurrently."""121        try:122            logging.info(f"Starting AI analysis for job {job['id']}...")123            gemini_task = get_gemini_analysis(job['ocr_text'])124            claude_task = get_claude_analysis(job['ocr_text'])125 126            gemini_response, claude_response = await asyncio.gather(gemini_task, claude_task, return_exceptions=True)127 128            update_payload = { "status": "completed" }129            if not isinstance(gemini_response, Exception):130                update_payload["gemini_response"] = gemini_response131            if not isinstance(claude_response, Exception):132                update_payload["claude_response"] = claude_response133 134            await self.supabase.table("jobs").update(update_payload).eq("id", job['id']).execute()135            logging.info(f"AI analysis successful for job {job['id']}.")136 137        except Exception as e:138            logging.error(f"Error during AI analysis for job {job['id']}: {e}")139            await self.supabase.table("jobs").update({"status": "failed"}).eq("id", job['id']).execute()140            return e141