Spaces:
Runtime error
Runtime error
Download api_chat.py from sdudeja/agentic-extractor: direct link, hf CLI and curl.
- Browser
- Download file 8.98 kB
-
https://huggingface.co/spaces/sdudeja/agentic-extractor/resolve/main/api_chat.py
- Command line
-
hf download hf://spaces/sdudeja/agentic-extractor/api_chat.py
-
curl -L -o api_chat.py https://huggingface.co/spaces/sdudeja/agentic-extractor/resolve/main/api_chat.py
8.98 kB
| """ | |
| DocuLens - Chat API endpoint. | |
| Provides /api/v1/chat for the AI Assistant panel. | |
| Uses the same Qwen2.5-VL models via HF Inference API for document Q&A. | |
| """ | |
| import os | |
| import io | |
| import base64 | |
| import logging | |
| from typing import Optional | |
| from fastapi import APIRouter, UploadFile, File, Form, HTTPException, Depends | |
| from fastapi.responses import StreamingResponse | |
| from PIL import Image | |
| from huggingface_hub import InferenceClient | |
| from hf_client import create_inference_client, robust_chat_completion | |
| from middleware import check_rate_limit, validate_file_upload, sanitize_filename | |
| logger = logging.getLogger(__name__) | |
| chat_router = APIRouter() | |
| HF_TOKEN = os.environ.get("HF_TOKEN", "") | |
| # Model used for chat — the 72B is best for conversational Q&A | |
| CHAT_MODEL = "Qwen/Qwen2.5-VL-72B-Instruct" | |
| # Fallback models if primary is unavailable | |
| CHAT_FALLBACK_MODELS = [ | |
| "Qwen/Qwen2.5-VL-7B-Instruct", | |
| "Qwen/Qwen2.5-VL-3B-Instruct", | |
| ] | |
| SYSTEM_PROMPT = """You are DocuLens Assistant, an expert in document analysis and data extraction. | |
| You help users understand their documents, extract information, and answer questions about document content. | |
| When analyzing a document image: | |
| - Identify the document type (invoice, receipt, tax form, etc.) | |
| - Describe key fields and their values accurately | |
| - Point out any issues or anomalies you notice | |
| - Be precise with numbers, dates, and amounts | |
| When answering questions about extracted data: | |
| - Reference specific fields and values | |
| - Perform calculations if asked (totals, tax rates, etc.) | |
| - Compare values if multiple documents are discussed | |
| Keep responses concise and helpful. Use markdown formatting for clarity when appropriate.""" | |
| def _image_to_data_url(image_bytes: bytes) -> str: | |
| """Convert image bytes to a data URL for the VLM.""" | |
| img = Image.open(io.BytesIO(image_bytes)) | |
| # Convert to RGB if needed | |
| if img.mode in ("RGBA", "P", "LA"): | |
| img = img.convert("RGB") | |
| buf = io.BytesIO() | |
| img.save(buf, format="JPEG", quality=85) | |
| b64 = base64.b64encode(buf.getvalue()).decode() | |
| return f"data:image/jpeg;base64,{b64}" | |
| def _build_messages( | |
| conversation_json: str, | |
| image_data_url: Optional[str] = None, | |
| extraction_context: Optional[str] = None, | |
| ) -> list[dict]: | |
| """Build the messages array for the VLM from conversation history.""" | |
| import json | |
| messages = [{"role": "system", "content": SYSTEM_PROMPT}] | |
| try: | |
| conversation = json.loads(conversation_json) | |
| except (json.JSONDecodeError, TypeError): | |
| conversation = [] | |
| for msg in conversation: | |
| role = msg.get("role", "user") | |
| text = msg.get("content", "") | |
| msg_image = msg.get("image") | |
| # Build content blocks | |
| if role == "user": | |
| content_parts: list[dict] = [] | |
| # Add extraction context if this is the first user message with it | |
| if extraction_context and msg == conversation[-1]: | |
| content_parts.append({ | |
| "type": "text", | |
| "text": f"[Current extraction context]\n{extraction_context}", | |
| }) | |
| content_parts.append({"type": "text", "text": text}) | |
| # Attach image if present (either from this message or the provided image) | |
| img_url = None | |
| if msg == conversation[-1] and image_data_url: | |
| img_url = image_data_url | |
| elif msg_image: | |
| img_url = msg_image | |
| if img_url: | |
| content_parts.append({ | |
| "type": "image_url", | |
| "image_url": {"url": img_url}, | |
| }) | |
| messages.append({"role": "user", "content": content_parts}) | |
| else: | |
| messages.append({"role": "assistant", "content": text}) | |
| return messages | |
| async def api_chat( | |
| messages: str = Form(..., description="JSON array of conversation messages"), | |
| file: Optional[UploadFile] = File(None, description="Optional image/PDF to discuss"), | |
| extraction_context: Optional[str] = Form(None, description="Current extraction result JSON for context"), | |
| model: Optional[str] = Form(None, description="Model override"), | |
| stream: bool = Form(False, description="Whether to stream the response"), | |
| api_key: str = Depends(check_rate_limit), | |
| ): | |
| """ | |
| Chat with the AI assistant about documents. | |
| Accepts a conversation history (messages) and optionally an image/PDF. | |
| Returns the assistant's response. | |
| """ | |
| if not HF_TOKEN: | |
| raise HTTPException(status_code=500, detail="HF_TOKEN not configured") | |
| # Process uploaded file if present | |
| image_data_url = None | |
| if file: | |
| validate_file_upload(file.filename, file.size or 0, file.content_type) | |
| try: | |
| file_bytes = await file.read() | |
| safe_name = sanitize_filename(file.filename) | |
| file_ext = os.path.splitext(safe_name)[1].lower().lstrip(".") | |
| if file_ext == "pdf": | |
| # Convert first page of PDF to image | |
| try: | |
| import fitz # PyMuPDF | |
| doc = fitz.open(stream=file_bytes, filetype="pdf") | |
| page = doc[0] | |
| pix = page.get_pixmap(dpi=200) | |
| img_bytes = pix.tobytes("jpeg") | |
| image_data_url = f"data:image/jpeg;base64,{base64.b64encode(img_bytes).decode()}" | |
| doc.close() | |
| except ImportError: | |
| # Fallback: use pdf2image | |
| from pdf2image import convert_from_bytes | |
| images = convert_from_bytes(file_bytes, first_page=1, last_page=1, dpi=200) | |
| if images: | |
| buf = io.BytesIO() | |
| images[0].save(buf, format="JPEG", quality=85) | |
| image_data_url = f"data:image/jpeg;base64,{base64.b64encode(buf.getvalue()).decode()}" | |
| else: | |
| image_data_url = _image_to_data_url(file_bytes) | |
| except Exception as e: | |
| logger.warning(f"Failed to process uploaded file: {e}") | |
| raise HTTPException(status_code=400, detail=f"Failed to process file: {str(e)}") | |
| # Build VLM messages | |
| vlm_messages = _build_messages(messages, image_data_url, extraction_context) | |
| # Try models in order | |
| use_model = model or CHAT_MODEL | |
| models_to_try = [use_model] + [m for m in CHAT_FALLBACK_MODELS if m != use_model] | |
| if stream: | |
| return StreamingResponse( | |
| _stream_chat(models_to_try, vlm_messages), | |
| media_type="text/event-stream", | |
| headers={ | |
| "Cache-Control": "no-cache", | |
| "Connection": "keep-alive", | |
| "X-Accel-Buffering": "no", | |
| }, | |
| ) | |
| # Non-streaming response | |
| errors = [] | |
| for model_id in models_to_try: | |
| try: | |
| client = create_inference_client(api_key=HF_TOKEN) | |
| response = robust_chat_completion( | |
| client, | |
| model=model_id, | |
| messages=vlm_messages, | |
| max_tokens=2048, | |
| temperature=0.3, | |
| ) | |
| content = response.choices[0].message.content | |
| return { | |
| "role": "assistant", | |
| "content": content, | |
| "model": model_id, | |
| } | |
| except Exception as e: | |
| short = model_id.split("/")[-1] | |
| errors.append(f"{short}: {type(e).__name__}") | |
| logger.warning(f"Chat failed with {model_id}: {e}") | |
| raise HTTPException( | |
| status_code=502, | |
| detail=f"All models failed: {'; '.join(errors)}", | |
| ) | |
| async def _stream_chat(models: list[str], vlm_messages: list[dict]): | |
| """Generator for SSE streaming.""" | |
| import json | |
| errors = [] | |
| for model_id in models: | |
| try: | |
| client = create_inference_client(api_key=HF_TOKEN) | |
| stream = client.chat_completion( | |
| model=model_id, | |
| messages=vlm_messages, | |
| max_tokens=2048, | |
| temperature=0.3, | |
| stream=True, | |
| ) | |
| # Send model info first | |
| yield f"data: {json.dumps({'type': 'meta', 'model': model_id})}\n\n" | |
| for chunk in stream: | |
| if chunk.choices and chunk.choices[0].delta.content: | |
| token = chunk.choices[0].delta.content | |
| yield f"data: {json.dumps({'type': 'token', 'content': token})}\n\n" | |
| yield f"data: {json.dumps({'type': 'done'})}\n\n" | |
| return | |
| except Exception as e: | |
| short = model_id.split("/")[-1] | |
| errors.append(f"{short}: {type(e).__name__}") | |
| logger.warning(f"Stream chat failed with {model_id}: {e}") | |
| error_detail = "All models failed: " + "; ".join(errors) | |
| yield f"data: {json.dumps({'type': 'error', 'detail': error_detail})}\n\n" | |