Doc-OCR / app.py
kd13's picture
Update app.py
95bd0c2 verified
Raw History Blame Contribute Delete
9.22 kB
import html
import logging
import os
import threading
from dataclasses import dataclass, field
from functools import lru_cache
from pathlib import Path
import spaces
import gradio as gr
from lexquery.config import GEMINI_MODEL, LANGUAGES
from lexquery.models import LocalModels
from lexquery.pipeline import Engine
from lexquery.schema import Result
logging.basicConfig(level=logging.INFO, format='%(levelname)s %(name)s: %(message)s')
@lru_cache(maxsize=1)
def load_models():
return LocalModels()
@dataclass
class Session:
engine: Engine = field(default_factory=lambda: Engine(load_models()))
chat: list = field(default_factory=list)
lock: object = field(default_factory=threading.RLock)
_sessions = {}
_sessions_lock = threading.RLock()
def session(request):
if not request.session_hash:
raise gr.Error('Open the app in a browser session first.')
with _sessions_lock:
if request.session_hash not in _sessions:
_sessions[request.session_hash] = Session()
return _sessions[request.session_hash]
def cleanup_id(session_id):
with _sessions_lock:
_sessions.pop(session_id, None)
def unload(request: gr.Request):
cleanup_id(request.session_hash)
def initialize(request: gr.Request):
session(request)
return request.session_hash
def safe_error(exc):
text = str(exc)
for name in ('GOOGLE_API_KEY', 'GEMINI_API_KEY'):
key = os.getenv(name)
if key:
text = text.replace(key, '[key hidden]')
return text
def read_uploads(files, request: gr.Request, progress=gr.Progress()):
item = session(request)
with item.lock:
keep, messages = set(), []
old_ids = set(item.engine.documents)
for i, file in enumerate(files or []):
path = Path(file)
progress(i / max(1, len(files)), desc='Reading ' + path.name)
try:
doc = item.engine.ingest(path.name, path.read_bytes())
keep.add(doc.id)
messages.append(f'**{html.escape(doc.name)}** · {len(doc.pages)} pages · ready')
uncertain = [str(p.number) for p in doc.pages if p.warnings]
if uncertain:
messages.append('Extraction notes on pages ' + ', '.join(uncertain) + '; see Extracted text.')
except Exception as exc:
logging.exception('PDF extraction failed')
messages.append(f'**{html.escape(path.name)}** · {html.escape(safe_error(exc))}')
for doc_id in set(item.engine.documents) - keep:
del item.engine.documents[doc_id]
if old_ids != keep:
item.engine.invalidate()
item.engine.history.clear()
item.chat.clear()
choices = [(d.name, d.id) for d in item.engine.documents.values()]
ids = [i for _, i in choices]
progress(1, desc='Finished')
return ('\n\n'.join(messages) or 'Upload a PDF to begin.',
gr.Dropdown(choices=choices, value=ids, multiselect=True),
gr.Dropdown(choices=choices, value=ids[0] if ids else None),
[], '', '')
def extracted_text(doc_id, request: gr.Request):
item = session(request)
with item.lock:
doc = item.engine.documents.get(doc_id)
if not doc:
return ''
return '\n\n'.join(f'PAGE {p.number} · {p.language}\n{p.original}\n' +
('\nNotes: ' + '; '.join(p.warnings) if p.warnings else '') for p in doc.pages)
def result_markdown(result):
paragraphs = [html.escape(text) for text in result.displayed] if result.status in (
'checked', 'conversational') else []
for note in result.notes:
paragraphs.append(html.escape(safe_error(note)))
return '\n\n'.join(paragraphs) or 'I could not find a supported answer in the selected documents.'
def ask(question, ids, language, request: gr.Request):
item = session(request)
with item.lock:
if not question.strip():
return item.chat, ''
result = item.engine.present(item.engine.ask(question, ids or []), language)
item.chat.extend([{'role': 'user', 'content': question},
{'role': 'assistant', 'content': result_markdown(result)}])
return item.chat, ''
@spaces.GPU
def ask(question, ids, language, request: gr.Request):
item = session(request)
with item.lock:
if not question.strip():
return item.chat, ''
result = item.engine.present(item.engine.ask(question, ids or []), language)
item.chat.extend([{'role': 'user', 'content': question},
{'role': 'assistant', 'content': result_markdown(result)}])
return item.chat, ''
@spaces.GPU
def summarize(ids, language, request: gr.Request, progress=gr.Progress()):
item = session(request)
with item.lock:
if not ids:
return 'Upload and select a PDF first.'
sections = []
def summarize(ids, language, request: gr.Request, progress=gr.Progress()):
item = session(request)
with item.lock:
if not ids:
return 'Upload and select a PDF first.'
sections = []
for i, doc_id in enumerate(ids):
doc = item.engine.documents.get(doc_id)
if doc is None:
continue
progress(i / len(ids), desc='Summarizing ' + doc.name)
try:
report = item.engine.report(doc_id)
result = item.engine.present(Result.model_validate(report['summary']), language)
section = '### ' + html.escape(doc.name) + '\n\n' + result_markdown(result)
fields = []
for name, facts in report['fields'].items():
if facts:
fields.append('**' + name.replace('_', ' ').title() + ':** ' +
'; '.join(html.escape(f['value']) for f in facts))
if fields:
section += '\n\n' + '\n\n'.join(fields)
for note in report.get('field_notes', []):
section += '\n\n' + html.escape(safe_error(note))
sections.append(section)
except Exception as exc:
logging.exception('Report failed')
sections.append('### ' + html.escape(doc.name) + '\n\n' + html.escape(safe_error(exc)))
progress(1, desc='Finished')
return '\n\n---\n\n'.join(sections)
def clear_chat(request: gr.Request):
item = session(request)
with item.lock:
item.engine.history.clear()
item.chat.clear()
return []
def build_app():
with gr.Blocks(title='LexQuery', delete_cache=(3600, 3600)) as demo:
lifecycle = gr.State(None, delete_callback=cleanup_id)
gr.Markdown('# LexQuery\nUpload your PDFs. Ask questions. Get clear summaries.')
with gr.Row():
with gr.Column(scale=1, min_width=280):
files = gr.File(label='Documents', file_types=['.pdf'], file_count='multiple', type='filepath')
gr.Markdown('Scanned and text PDFs · English and 22 Indian languages')
status = gr.Markdown('Upload a PDF to begin.')
selected = gr.Dropdown(choices=[], multiselect=True, label='Use these documents')
language = gr.Dropdown(LANGUAGES, value='English', label='Answer language')
with gr.Column(scale=3):
with gr.Tab('Questions'):
chat = gr.Chatbot(label='Conversation', height=460)
question = gr.Textbox(label='Your question', placeholder='What does this document say about…?')
with gr.Row():
ask_button = gr.Button('Ask', variant='primary')
reset = gr.Button('New conversation')
with gr.Tab('Summaries'):
gr.Markdown('Summaries cover each selected document independently of your conversation.')
summary_button = gr.Button('Summarize documents', variant='primary')
summary = gr.Markdown()
with gr.Accordion('Extracted text', open=False):
text_doc = gr.Dropdown(choices=[], label='Document')
transcript = gr.Textbox(label='Original text', lines=18, interactive=False)
gr.Markdown('Check important findings against the original PDF.')
files.change(read_uploads, files, [status, selected, text_doc, chat, summary, transcript])
text_doc.change(extracted_text, text_doc, transcript)
for trigger in (ask_button.click, question.submit):
trigger(ask, [question, selected, language], [chat, question])
reset.click(clear_chat, outputs=chat)
summary_button.click(summarize, [selected, language], summary)
demo.load(initialize, outputs=lifecycle)
demo.unload(unload)
return demo.queue(default_concurrency_limit=1)
if __name__ == "__main__":
logging.info("LexQuery Gemini model: %s", GEMINI_MODEL)
build_app().launch(
server_name="0.0.0.0",
server_port=7860,
share=False,
)