adudeja commited on
Commit
8a8efdb
·
1 Parent(s): 5ede76e

AI assistant panel

Browse files
Files changed (2) hide show
  1. api_chat.py +241 -0
  2. app.py +2 -0
api_chat.py ADDED
@@ -0,0 +1,241 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ DocuMint AI - Chat API endpoint.
3
+
4
+ Provides /api/v1/chat for the AI Assistant panel.
5
+ Uses the same Qwen2.5-VL models via HF Inference API for document Q&A.
6
+ """
7
+
8
+ import os
9
+ import io
10
+ import base64
11
+ import logging
12
+ from typing import Optional
13
+
14
+ from fastapi import APIRouter, UploadFile, File, Form, HTTPException
15
+ from fastapi.responses import StreamingResponse
16
+ from PIL import Image
17
+ from huggingface_hub import InferenceClient
18
+
19
+ logger = logging.getLogger(__name__)
20
+
21
+ chat_router = APIRouter()
22
+
23
+ HF_TOKEN = os.environ.get("HF_TOKEN", "")
24
+
25
+ # Model used for chat — the 72B is best for conversational Q&A
26
+ CHAT_MODEL = "Qwen/Qwen2.5-VL-72B-Instruct"
27
+
28
+ # Fallback models if primary is unavailable
29
+ CHAT_FALLBACK_MODELS = [
30
+ "Qwen/Qwen2.5-VL-7B-Instruct",
31
+ "Qwen/Qwen2.5-VL-3B-Instruct",
32
+ ]
33
+
34
+ SYSTEM_PROMPT = """You are DocuMint AI Assistant, an expert in document analysis and data extraction.
35
+ You help users understand their documents, extract information, and answer questions about document content.
36
+
37
+ When analyzing a document image:
38
+ - Identify the document type (invoice, receipt, tax form, etc.)
39
+ - Describe key fields and their values accurately
40
+ - Point out any issues or anomalies you notice
41
+ - Be precise with numbers, dates, and amounts
42
+
43
+ When answering questions about extracted data:
44
+ - Reference specific fields and values
45
+ - Perform calculations if asked (totals, tax rates, etc.)
46
+ - Compare values if multiple documents are discussed
47
+
48
+ Keep responses concise and helpful. Use markdown formatting for clarity when appropriate."""
49
+
50
+
51
+ def _image_to_data_url(image_bytes: bytes) -> str:
52
+ """Convert image bytes to a data URL for the VLM."""
53
+ img = Image.open(io.BytesIO(image_bytes))
54
+ # Convert to RGB if needed
55
+ if img.mode in ("RGBA", "P", "LA"):
56
+ img = img.convert("RGB")
57
+ buf = io.BytesIO()
58
+ img.save(buf, format="JPEG", quality=85)
59
+ b64 = base64.b64encode(buf.getvalue()).decode()
60
+ return f"data:image/jpeg;base64,{b64}"
61
+
62
+
63
+ def _build_messages(
64
+ conversation_json: str,
65
+ image_data_url: Optional[str] = None,
66
+ extraction_context: Optional[str] = None,
67
+ ) -> list[dict]:
68
+ """Build the messages array for the VLM from conversation history."""
69
+ import json
70
+
71
+ messages = [{"role": "system", "content": SYSTEM_PROMPT}]
72
+
73
+ try:
74
+ conversation = json.loads(conversation_json)
75
+ except (json.JSONDecodeError, TypeError):
76
+ conversation = []
77
+
78
+ for msg in conversation:
79
+ role = msg.get("role", "user")
80
+ text = msg.get("content", "")
81
+ msg_image = msg.get("image")
82
+
83
+ # Build content blocks
84
+ if role == "user":
85
+ content_parts: list[dict] = []
86
+
87
+ # Add extraction context if this is the first user message with it
88
+ if extraction_context and msg == conversation[-1]:
89
+ content_parts.append({
90
+ "type": "text",
91
+ "text": f"[Current extraction context]\n{extraction_context}",
92
+ })
93
+
94
+ content_parts.append({"type": "text", "text": text})
95
+
96
+ # Attach image if present (either from this message or the provided image)
97
+ img_url = None
98
+ if msg == conversation[-1] and image_data_url:
99
+ img_url = image_data_url
100
+ elif msg_image:
101
+ img_url = msg_image
102
+
103
+ if img_url:
104
+ content_parts.append({
105
+ "type": "image_url",
106
+ "image_url": {"url": img_url},
107
+ })
108
+
109
+ messages.append({"role": "user", "content": content_parts})
110
+ else:
111
+ messages.append({"role": "assistant", "content": text})
112
+
113
+ return messages
114
+
115
+
116
+ @chat_router.post("/api/v1/chat")
117
+ async def api_chat(
118
+ messages: str = Form(..., description="JSON array of conversation messages"),
119
+ file: Optional[UploadFile] = File(None, description="Optional image/PDF to discuss"),
120
+ extraction_context: Optional[str] = Form(None, description="Current extraction result JSON for context"),
121
+ model: Optional[str] = Form(None, description="Model override"),
122
+ stream: bool = Form(False, description="Whether to stream the response"),
123
+ ):
124
+ """
125
+ Chat with the AI assistant about documents.
126
+
127
+ Accepts a conversation history (messages) and optionally an image/PDF.
128
+ Returns the assistant's response.
129
+ """
130
+ if not HF_TOKEN:
131
+ raise HTTPException(status_code=500, detail="HF_TOKEN not configured")
132
+
133
+ # Process uploaded file if present
134
+ image_data_url = None
135
+ if file:
136
+ try:
137
+ file_bytes = await file.read()
138
+ file_ext = (file.filename or "").lower().split(".")[-1]
139
+
140
+ if file_ext == "pdf":
141
+ # Convert first page of PDF to image
142
+ try:
143
+ import fitz # PyMuPDF
144
+ doc = fitz.open(stream=file_bytes, filetype="pdf")
145
+ page = doc[0]
146
+ pix = page.get_pixmap(dpi=200)
147
+ img_bytes = pix.tobytes("jpeg")
148
+ image_data_url = f"data:image/jpeg;base64,{base64.b64encode(img_bytes).decode()}"
149
+ doc.close()
150
+ except ImportError:
151
+ # Fallback: use pdf2image
152
+ from pdf2image import convert_from_bytes
153
+ images = convert_from_bytes(file_bytes, first_page=1, last_page=1, dpi=200)
154
+ if images:
155
+ buf = io.BytesIO()
156
+ images[0].save(buf, format="JPEG", quality=85)
157
+ image_data_url = f"data:image/jpeg;base64,{base64.b64encode(buf.getvalue()).decode()}"
158
+ else:
159
+ image_data_url = _image_to_data_url(file_bytes)
160
+ except Exception as e:
161
+ logger.warning(f"Failed to process uploaded file: {e}")
162
+ raise HTTPException(status_code=400, detail=f"Failed to process file: {str(e)}")
163
+
164
+ # Build VLM messages
165
+ vlm_messages = _build_messages(messages, image_data_url, extraction_context)
166
+
167
+ # Try models in order
168
+ use_model = model or CHAT_MODEL
169
+ models_to_try = [use_model] + [m for m in CHAT_FALLBACK_MODELS if m != use_model]
170
+
171
+ if stream:
172
+ return StreamingResponse(
173
+ _stream_chat(models_to_try, vlm_messages),
174
+ media_type="text/event-stream",
175
+ headers={
176
+ "Cache-Control": "no-cache",
177
+ "Connection": "keep-alive",
178
+ "X-Accel-Buffering": "no",
179
+ },
180
+ )
181
+
182
+ # Non-streaming response
183
+ errors = []
184
+ for model_id in models_to_try:
185
+ try:
186
+ client = InferenceClient(api_key=HF_TOKEN)
187
+ response = client.chat_completion(
188
+ model=model_id,
189
+ messages=vlm_messages,
190
+ max_tokens=2048,
191
+ temperature=0.3,
192
+ )
193
+ content = response.choices[0].message.content
194
+ return {
195
+ "role": "assistant",
196
+ "content": content,
197
+ "model": model_id,
198
+ }
199
+ except Exception as e:
200
+ short = model_id.split("/")[-1]
201
+ errors.append(f"{short}: {type(e).__name__}")
202
+ logger.warning(f"Chat failed with {model_id}: {e}")
203
+
204
+ raise HTTPException(
205
+ status_code=502,
206
+ detail=f"All models failed: {'; '.join(errors)}",
207
+ )
208
+
209
+
210
+ async def _stream_chat(models: list[str], vlm_messages: list[dict]):
211
+ """Generator for SSE streaming."""
212
+ import json
213
+
214
+ errors = []
215
+ for model_id in models:
216
+ try:
217
+ client = InferenceClient(api_key=HF_TOKEN)
218
+ stream = client.chat_completion(
219
+ model=model_id,
220
+ messages=vlm_messages,
221
+ max_tokens=2048,
222
+ temperature=0.3,
223
+ stream=True,
224
+ )
225
+ # Send model info first
226
+ yield f"data: {json.dumps({'type': 'meta', 'model': model_id})}\n\n"
227
+
228
+ for chunk in stream:
229
+ if chunk.choices and chunk.choices[0].delta.content:
230
+ token = chunk.choices[0].delta.content
231
+ yield f"data: {json.dumps({'type': 'token', 'content': token})}\n\n"
232
+
233
+ yield f"data: {json.dumps({'type': 'done'})}\n\n"
234
+ return
235
+
236
+ except Exception as e:
237
+ short = model_id.split("/")[-1]
238
+ errors.append(f"{short}: {type(e).__name__}")
239
+ logger.warning(f"Stream chat failed with {model_id}: {e}")
240
+
241
+ yield f"data: {json.dumps({'type': 'error', 'detail': f'All models failed: {'; '.join(errors)}'})}\n\n"
app.py CHANGED
@@ -334,9 +334,11 @@ if __name__ == "__main__":
334
  from api import api_router
335
  from api_runs import runs_router
336
  from api_prompts import prompts_router
 
337
  app.include_router(api_router)
338
  app.include_router(runs_router)
339
  app.include_router(prompts_router)
 
340
 
341
  print(f"✅ REST API routes mounted at {local_url}api/v1/")
342
 
 
334
  from api import api_router
335
  from api_runs import runs_router
336
  from api_prompts import prompts_router
337
+ from api_chat import chat_router
338
  app.include_router(api_router)
339
  app.include_router(runs_router)
340
  app.include_router(prompts_router)
341
+ app.include_router(chat_router)
342
 
343
  print(f"✅ REST API routes mounted at {local_url}api/v1/")
344