File size: 9,661 Bytes
04c4194
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
"""
DocuLens β€” Batch processing endpoints.

Accepts multiple files in a single request, processes them sequentially,
and returns per-file results with a batch summary.
"""

import os
import json
import time
import tempfile
import logging
from typing import Optional
from uuid import uuid4

from fastapi import APIRouter, UploadFile, File, Form, HTTPException, Depends
from PIL import Image

from extraction import (
    extract_invoice, extract_from_pdf, ExtractionResult,
)
from db.supabase import save_document, create_run, complete_run, save_stage_result
from middleware import check_rate_limit, validate_file_upload, sanitize_filename
from api_webhooks import deliver_webhook

logger = logging.getLogger(__name__)

batch_router = APIRouter()

HF_TOKEN = os.environ.get("HF_TOKEN", "")

# Maximum files per batch (free tier)
MAX_BATCH_SIZE = int(os.environ.get("MAX_BATCH_SIZE", "20"))


# ---------------------------------------------------------------------------
# Batch processing endpoint
# ---------------------------------------------------------------------------

@batch_router.post("/api/v1/batch/extract")
async def batch_extract(
    files: list[UploadFile] = File(...),
    model: str = Form("Qwen/Qwen2.5-VL-72B-Instruct"),
    force_ocr: str = Form("false"),
    persist: str = Form("true"),
    use_case: str = Form("invoice"),
    org_id: str = Form(""),
    user_id: str = Form(""),
    api_key: str = Depends(check_rate_limit),
):
    """
    Extract structured data from multiple uploaded files.

    Returns a batch result with per-file status, extraction data,
    and a summary of passed/failed counts.
    """
    if not files:
        raise HTTPException(status_code=400, detail="No files provided")

    # ── Tier-based batch limits ──
    effective_max = MAX_BATCH_SIZE
    if user_id:
        try:
            from api_usage import check_feature, check_quota, TIER_FEATURES, get_user_tier
            if not check_feature(user_id, "batch_upload"):
                raise HTTPException(
                    status_code=403,
                    detail="Batch upload is not available on the free plan. Upgrade to Starter or above.",
                )
            tier = get_user_tier(user_id)
            tier_max = TIER_FEATURES.get(tier, {}).get("max_batch_files", MAX_BATCH_SIZE)
            effective_max = min(tier_max, MAX_BATCH_SIZE) if tier_max else MAX_BATCH_SIZE

            # Check quota for all files in the batch
            quota = check_quota(user_id, pages_requested=len(files))
            if not quota["allowed"]:
                raise HTTPException(status_code=429, detail=quota["message"])
        except HTTPException:
            raise
        except Exception as e:
            logger.warning("Batch tier check failed (non-blocking): %s", e)

    if len(files) > effective_max:
        raise HTTPException(
            status_code=400,
            detail=f"Batch size exceeds limit. Maximum {effective_max} files per batch.",
        )

    batch_id = str(uuid4())
    batch_start = time.time()
    use_ocr_only = force_ocr.lower() in ("true", "1", "yes") or model == "ocr-only"
    should_persist = persist.lower() in ("true", "1", "yes")
    token = "" if use_ocr_only else HF_TOKEN

    results = []
    passed = 0
    failed = 0

    for i, file in enumerate(files):
        file_start = time.time()
        safe_filename = sanitize_filename(file.filename)
        file_result = {
            "index": i,
            "filename": safe_filename,
            "status": "pending",
            "run_id": None,
            "data": None,
            "error": None,
            "processing_time_ms": 0,
        }

        try:
            # Validate each file
            file_size = file.size or 0
            validate_file_upload(file.filename, file_size, file.content_type)

            # Persistence
            doc_id = None
            run_id = None
            if should_persist:
                ext = os.path.splitext(safe_filename)[1].lower()
                file_type = "pdf" if ext == ".pdf" else "image"
                doc_id = save_document(
                    filename=safe_filename,
                    file_type=file_type,
                    file_size=file_size,
                    org_id=org_id or None,
                    user_id=user_id or None,
                )
                run_id = create_run(
                    document_id=doc_id,
                    use_case=use_case,
                    org_id=org_id or None,
                    user_id=user_id or None,
                )
                file_result["run_id"] = run_id

            # Read file contents
            contents = await file.read()
            ext = os.path.splitext(safe_filename)[1].lower()

            # Extract
            if ext == ".pdf":
                page_results = extract_from_pdf(
                    contents, hf_token=token, force_ocr=use_ocr_only, use_case=use_case, model_id=model,
                )
                if not page_results:
                    raise ValueError("Could not extract data from PDF")
                result = page_results[0]
                # Merge multi-page line items
                base_li = len(result.line_items)
                for pi, pr in enumerate(page_results[1:], start=1):
                    for bb in pr.bounding_boxes:
                        if bb.field.startswith("line_item_"):
                            li_idx = int(bb.field.split("_")[-1])
                            bb.field = f"line_item_{base_li + li_idx}"
                            bb.page = pi
                            result.bounding_boxes.append(bb)
                    result.line_items.extend(pr.line_items)
                    base_li += len(pr.line_items)
            else:
                with tempfile.NamedTemporaryFile(suffix=ext or ".jpg", delete=False) as tmp:
                    tmp.write(contents)
                    tmp_path = tmp.name
                try:
                    image = Image.open(tmp_path)
                    result = extract_invoice(
                        image, hf_token=token, force_ocr=use_ocr_only, use_case=use_case, model_id=model,
                    )
                finally:
                    os.unlink(tmp_path)

            file_ms = int((time.time() - file_start) * 1000)
            result.processing_time_ms = file_ms

            is_failure = result.extraction_method.startswith("failed")

            if is_failure:
                file_result["status"] = "failed"
                file_result["error"] = _describe_failure(result.extraction_method)
                failed += 1
            else:
                file_result["status"] = "passed"
                file_result["data"] = result.to_dict()
                # Remove large base64 images from batch response to keep payload small
                file_result["data"].pop("page_image", None)
                passed += 1

                # Increment usage for successful extraction
                if user_id:
                    try:
                        from api_usage import increment_usage
                        increment_usage(user_id, pages=1)
                    except Exception:
                        pass

            file_result["processing_time_ms"] = file_ms

            # Persist run result
            if run_id:
                save_stage_result(
                    run_id=run_id,
                    stage_type="extraction",
                    stage_name="Data Extraction",
                    status="failed" if is_failure else "passed",
                    sort_order=1,
                    output=result.to_dict() if not is_failure else None,
                    error_message=file_result.get("error"),
                    duration_ms=file_ms,
                )
                complete_run(
                    run_id=run_id,
                    status="completed",
                    overall_result="failed" if is_failure else "passed",
                    extraction_data=result.to_dict() if not is_failure else {},
                    processing_time_ms=file_ms,
                    error_message=file_result.get("error"),
                )

        except HTTPException as he:
            file_result["status"] = "failed"
            file_result["error"] = he.detail
            file_result["processing_time_ms"] = int((time.time() - file_start) * 1000)
            failed += 1
        except Exception as e:
            logger.warning(f"Batch file {i} ({safe_filename}) failed: {e}")
            file_result["status"] = "failed"
            file_result["error"] = str(e)
            file_result["processing_time_ms"] = int((time.time() - file_start) * 1000)
            failed += 1

        results.append(file_result)

    batch_ms = int((time.time() - batch_start) * 1000)

    batch_response = {
        "batch_id": batch_id,
        "total": len(files),
        "passed": passed,
        "failed": failed,
        "processing_time_ms": batch_ms,
        "results": results,
    }

    # Fire webhook
    try:
        deliver_webhook("batch.completed", batch_response)
    except Exception:
        pass

    return batch_response


# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------

def _describe_failure(method: str) -> str:
    if "no_token" in method:
        return "HF_TOKEN not set β€” AI Vision unavailable."
    if "vlm" in method:
        return "AI Vision models failed. Inference API may be loading or rate-limited."
    return "Could not extract readable text. Image may be blurry or handwritten."