diff --git a/app/api/openai.py b/app/api/openai.py index 07859e50..273d09cb 100644 --- a/app/api/openai.py +++ b/app/api/openai.py @@ -1,14 +1,20 @@ """ AI provider and OpenAI API endpoints. -Exposes two endpoints: -- GET /api/ai/test – tests the currently configured AI provider (generic, provider-agnostic) -- GET /api/openai/test – backward-compatible alias that tests the OpenAI API specifically +Exposes three endpoints: +- GET /api/ai/test – tests the currently configured AI provider (generic, provider-agnostic) +- GET /api/openai/test – backward-compatible alias that tests the OpenAI API specifically +- POST /api/ai/test-extraction – runs the metadata-extraction prompt against the configured AI provider + with caller-supplied plaintext and returns the raw response, parsed JSON, + and extracted tags so operators can evaluate model quality. """ +import json import logging +import re from fastapi import APIRouter, Request +from pydantic import BaseModel, Field from app.auth import require_login from app.config import settings @@ -18,6 +24,10 @@ logger = logging.getLogger(__name__) router = APIRouter() +# Maximum number of characters accepted for a test-extraction request. +# Keeps individual requests reasonable without blocking any real-world document. +_MAX_EXTRACTION_TEXT_LEN = 50_000 + def _get_exception_chain_detail(exc: Exception) -> str: """ @@ -202,3 +212,123 @@ async def test_ai_provider_connection(request: Request): "message": f"Connection failed: {detail}", "provider": provider_name, } + + +class ExtractionTestRequest(BaseModel): + """Request body for the AI extraction test endpoint.""" + + text: str = Field(..., min_length=1, max_length=_MAX_EXTRACTION_TEXT_LEN, description="Plain-text document content") + + +def _build_extraction_prompt(text: str) -> str: + """Return the metadata-extraction prompt used in the standard processing pipeline.""" + return ( + "You are a specialized document analyzer trained to extract structured metadata from documents.\n" + "Your task is to analyze the given text and return a well-structured JSON object.\n\n" + "Extract and return the following fields:\n" + "1. **filename**: Machine-readable filename " + "(YYYY-MM-DD_DescriptiveTitle, use only letters, numbers, periods, and underscores).\n" + '2. **empfaenger**: The recipient, or "Unknown" if not found.\n' + '3. **absender**: The sender, or "Unknown" if not found.\n' + "4. **correspondent**: The entity or company that issued the document " + '(shortest possible name, e.g., "Amazon" instead of "Amazon EU SARL, German branch").\n' + "5. **kommunikationsart**: One of [Behoerdlicher_Brief, Rechnung, Kontoauszug, Vertrag, " + "Quittung, Privater_Brief, Einladung, Gewerbliche_Korrespondenz, Newsletter, Werbung, Sonstiges].\n" + "6. **kommunikationskategorie**: One of [Amtliche_Postbehoerdliche_Dokumente, " + "Finanz_und_Vertragsdokumente, Geschaeftliche_Kommunikation, " + "Private_Korrespondenz, Sonstige_Informationen].\n" + "7. **document_type**: Precise classification (e.g., Invoice, Contract, Information, Unknown).\n" + "8. **tags**: A list of up to 4 relevant thematic keywords.\n" + '9. **language**: Detected document language (ISO 639-1 code, e.g., "de" or "en").\n' + "10. **title**: A human-readable title summarizing the document content.\n" + "11. **confidence_score**: A numeric value (0-100) indicating the confidence level " + "of the extracted metadata.\n" + "12. **reference_number**: Extracted invoice/order/reference number if available.\n" + "13. **monetary_amounts**: A list of key monetary values detected in the document.\n\n" + "### Important Rules:\n" + "- **OCR Correction**: Assume the text has been corrected for OCR errors.\n" + "- **Tagging**: Max 4 tags, avoiding generic or overly specific terms.\n" + "- **Title**: Concise, no addresses, and contains key identifying features.\n" + "- **Date Selection**: Use the most relevant date if multiple are found.\n" + "- **Output Language**: Maintain the document's original language.\n\n" + f"Extracted text:\n{text}\n\n" + "Return only valid JSON with no additional commentary.\n" + ) + + +def _extract_json_from_text(text: str): + """Try to extract a JSON object from the LLM response text.""" + pattern = r"```(?:json)?\s*(\{.*?\})\s*```" + match = re.search(pattern, text, re.DOTALL) + if match: + return match.group(1) + start = text.find("{") + end = text.rfind("}") + if start != -1 and end != -1 and end > start: + return text[start : end + 1] + return None + + +@router.post("/ai/test-extraction") +@require_login +async def test_ai_extraction(body: ExtractionTestRequest, request: Request): + """ + Run the metadata-extraction prompt against the configured AI provider. + + Accepts plain-text document content, sends it through the same prompt used + by the background processing pipeline, and returns: + - ``raw_response``: verbatim LLM output + - ``parsed_json``: the extracted JSON object (null when parsing fails) + - ``tags``: the ``tags`` list from the parsed JSON (empty list on failure) + - ``provider`` / ``model``: which provider / model was used + """ + from app.utils.ai_provider import get_ai_provider + + provider_name = settings.ai_provider + model = settings.ai_model or settings.openai_model + + logger.info(f"AI extraction test requested: provider={provider_name}, model={model}") + + try: + provider = get_ai_provider() + prompt = _build_extraction_prompt(body.text) + raw_response = provider.chat_completion( + messages=[ + {"role": "system", "content": "You are an intelligent document classifier."}, + {"role": "user", "content": prompt}, + ], + model=model, + temperature=0, + ) + except ValueError as e: + logger.warning(f"AI extraction test – configuration error: {e}") + return {"status": "error", "message": str(e), "provider": provider_name} + except Exception as e: + detail = _get_exception_chain_detail(e) + logger.error(f"AI extraction test failed for provider '{provider_name}': {detail}", exc_info=True) + return {"status": "error", "message": f"AI call failed: {detail}", "provider": provider_name} + + # Attempt to parse JSON from the response + parsed_json = None + tags: list = [] + parse_error = None + json_text = _extract_json_from_text(raw_response) + if json_text: + try: + parsed_json = json.loads(json_text) + tags = parsed_json.get("tags", []) + except json.JSONDecodeError as exc: + parse_error = str(exc) + logger.warning(f"AI extraction test: JSON parse error: {exc}") + else: + parse_error = "No JSON object found in response" + + return { + "status": "success", + "provider": provider_name, + "model": model, + "raw_response": raw_response, + "parsed_json": parsed_json, + "tags": tags, + "parse_error": parse_error, + } diff --git a/frontend/templates/status_dashboard.html b/frontend/templates/status_dashboard.html index 259a2afa..1e4e95d2 100644 --- a/frontend/templates/status_dashboard.html +++ b/frontend/templates/status_dashboard.html @@ -181,6 +181,12 @@ data-provider="ai_provider"> Test Connection + {% endif %} {% elif name == "Azure AI" %} {% if provider.configured %} @@ -273,6 +279,85 @@ + + +
{% endblock %} @@ -573,6 +658,152 @@ document.addEventListener('DOMContentLoaded', function() { }); }); }); + + // ── AI Extraction Test Modal ────────────────────────────────────────────── + + const aiExtractionModal = document.getElementById('aiExtractionModal'); + const aiExtractionText = document.getElementById('aiExtractionText'); + const aiExtractionInput = document.getElementById('aiExtractionInput'); + const aiExtractionResults = document.getElementById('aiExtractionResults'); + const aiExtractionMeta = document.getElementById('aiExtractionMeta'); + const aiRawResponse = document.getElementById('aiRawResponse'); + const aiParsedJson = document.getElementById('aiParsedJson'); + const aiParsedJsonSection = document.getElementById('aiParsedJsonSection'); + const aiTags = document.getElementById('aiTags'); + const aiTagsSection = document.getElementById('aiTagsSection'); + const aiParseWarning = document.getElementById('aiParseWarning'); + const testAiExtractionBtn = document.getElementById('testAiExtractionBtn'); + const runAiExtractionBtn = document.getElementById('runAiExtractionBtn'); + const cancelAiExtractionBtn= document.getElementById('cancelAiExtractionBtn'); + const aiExtractionBackBtn = document.getElementById('aiExtractionBackBtn'); + const aiExtractionCloseBtn = document.getElementById('aiExtractionCloseBtn'); + const closeAiExtractionModal = document.getElementById('closeAiExtractionModal'); + + function openAiExtractionModal() { + // Reset to input view + aiExtractionText.value = ''; + aiExtractionInput.classList.remove('hidden'); + aiExtractionResults.classList.add('hidden'); + aiExtractionModal.classList.remove('hidden'); + } + + function closeAiExtractionModalFn() { + aiExtractionModal.classList.add('hidden'); + } + + function showAiExtractionResults(data) { + // Provider / model meta + aiExtractionMeta.textContent = `Provider: ${data.provider || '—'} | Model: ${data.model || '—'}`; + + // Raw response + aiRawResponse.textContent = data.raw_response || '(empty)'; + + // Parse warning + if (data.parse_error) { + aiParseWarning.textContent = `JSON parse issue: ${data.parse_error}`; + aiParseWarning.classList.remove('hidden'); + } else { + aiParseWarning.classList.add('hidden'); + } + + // Parsed JSON + if (data.parsed_json) { + aiParsedJson.textContent = JSON.stringify(data.parsed_json, null, 2); + aiParsedJsonSection.classList.remove('hidden'); + } else { + aiParsedJsonSection.classList.add('hidden'); + } + + // Tags + aiTags.innerHTML = ''; + const tags = Array.isArray(data.tags) ? data.tags : []; + if (tags.length > 0) { + tags.forEach(tag => { + const span = document.createElement('span'); + span.className = 'inline-flex items-center px-2.5 py-0.5 rounded-full text-xs font-medium bg-indigo-100 text-indigo-800'; + span.textContent = tag; + aiTags.appendChild(span); + }); + aiTagsSection.classList.remove('hidden'); + } else { + aiTagsSection.classList.add('hidden'); + } + + // Switch view + aiExtractionInput.classList.add('hidden'); + aiExtractionResults.classList.remove('hidden'); + } + + if (testAiExtractionBtn) { + testAiExtractionBtn.addEventListener('click', openAiExtractionModal); + } + + if (closeAiExtractionModal) { + closeAiExtractionModal.addEventListener('click', closeAiExtractionModalFn); + } + + if (cancelAiExtractionBtn) { + cancelAiExtractionBtn.addEventListener('click', closeAiExtractionModalFn); + } + + if (aiExtractionCloseBtn) { + aiExtractionCloseBtn.addEventListener('click', closeAiExtractionModalFn); + } + + if (aiExtractionBackBtn) { + aiExtractionBackBtn.addEventListener('click', function() { + aiExtractionResults.classList.add('hidden'); + aiExtractionInput.classList.remove('hidden'); + }); + } + + // Close when clicking the backdrop + if (aiExtractionModal) { + aiExtractionModal.addEventListener('click', function(e) { + if (e.target === aiExtractionModal) { + closeAiExtractionModalFn(); + } + }); + } + + if (runAiExtractionBtn) { + runAiExtractionBtn.addEventListener('click', function() { + const text = aiExtractionText.value.trim(); + if (!text) { + aiExtractionText.focus(); + return; + } + + const originalHTML = runAiExtractionBtn.innerHTML; + runAiExtractionBtn.innerHTML = 'Running…'; + runAiExtractionBtn.disabled = true; + cancelAiExtractionBtn.disabled = true; + + fetch('/api/ai/test-extraction', { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ text: text }), + }) + .then(response => response.json()) + .then(data => { + if (data.status === 'success') { + showAiExtractionResults(data); + } else { + closeAiExtractionModalFn(); + showModal('error', 'AI Extraction Failed', data.message || 'Unknown error'); + } + }) + .catch(error => { + closeAiExtractionModalFn(); + showModal('error', 'Connection Error', 'Error running extraction: ' + error.message); + }) + .finally(() => { + runAiExtractionBtn.innerHTML = originalHTML; + runAiExtractionBtn.disabled = false; + cancelAiExtractionBtn.disabled = false; + }); + }); + } }); {% endblock %} diff --git a/tests/test_api_openai.py b/tests/test_api_openai.py index 9d2c5d7b..71bc3968 100644 --- a/tests/test_api_openai.py +++ b/tests/test_api_openai.py @@ -309,3 +309,136 @@ class TestOpenAIConnectionErrors: assert data["status"] == "error" assert data.get("is_auth_error") is False + + +@pytest.mark.unit +class TestAiExtractionEndpoint: + """Tests for POST /api/ai/test-extraction endpoint.""" + + @patch("app.api.openai.settings") + def test_extraction_missing_text_returns_422(self, mock_settings, client): + """Test that missing text body returns 422 validation error.""" + response = client.post("/api/ai/test-extraction", json={}) + assert response.status_code == 422 + + @patch("app.api.openai.settings") + def test_extraction_empty_text_returns_422(self, mock_settings, client): + """Test that empty string text returns 422 validation error.""" + response = client.post("/api/ai/test-extraction", json={"text": ""}) + assert response.status_code == 422 + + @patch("app.utils.ai_provider.get_ai_provider") + @patch("app.api.openai.settings") + def test_extraction_returns_raw_response_and_parsed_json(self, mock_settings, mock_get_provider, client): + """Test successful extraction returns raw response, parsed JSON, and tags.""" + import json as json_module + + mock_settings.ai_provider = "openai" + mock_settings.ai_model = "gpt-4o-mini" + mock_settings.openai_model = "gpt-4o-mini" + + expected_json = { + "filename": "2024-01-01_Invoice", + "tags": ["invoice", "payment"], + "title": "January Invoice", + "document_type": "Invoice", + } + raw = "```json\n" + json_module.dumps(expected_json) + "\n```" + + mock_provider = MagicMock() + mock_provider.chat_completion.return_value = raw + mock_get_provider.return_value = mock_provider + + response = client.post("/api/ai/test-extraction", json={"text": "Invoice content here"}) + assert response.status_code == 200 + data = response.json() + + assert data["status"] == "success" + assert data["raw_response"] == raw + assert data["parsed_json"]["filename"] == "2024-01-01_Invoice" + assert data["tags"] == ["invoice", "payment"] + assert data["parse_error"] is None + assert data["provider"] == "openai" + assert data["model"] == "gpt-4o-mini" + + @patch("app.utils.ai_provider.get_ai_provider") + @patch("app.api.openai.settings") + def test_extraction_handles_invalid_json_in_response(self, mock_settings, mock_get_provider, client): + """Test that invalid JSON in LLM response is reported via parse_error.""" + mock_settings.ai_provider = "openai" + mock_settings.ai_model = "gpt-4o-mini" + mock_settings.openai_model = "gpt-4o-mini" + + mock_provider = MagicMock() + mock_provider.chat_completion.return_value = "Sorry, I cannot help with that." + mock_get_provider.return_value = mock_provider + + response = client.post("/api/ai/test-extraction", json={"text": "Some document text"}) + assert response.status_code == 200 + data = response.json() + + assert data["status"] == "success" + assert data["parsed_json"] is None + assert data["tags"] == [] + assert data["parse_error"] is not None + + @patch("app.utils.ai_provider.get_ai_provider") + @patch("app.api.openai.settings") + def test_extraction_provider_config_error(self, mock_settings, mock_get_provider, client): + """Test that configuration errors (missing keys) are returned as error status.""" + mock_settings.ai_provider = "anthropic" + mock_settings.ai_model = "claude-3" + mock_settings.openai_model = "gpt-4o-mini" + + mock_get_provider.side_effect = ValueError("ANTHROPIC_API_KEY must be set") + + response = client.post("/api/ai/test-extraction", json={"text": "Some document text"}) + assert response.status_code == 200 + data = response.json() + + assert data["status"] == "error" + assert "ANTHROPIC_API_KEY" in data["message"] + + @patch("app.utils.ai_provider.get_ai_provider") + @patch("app.api.openai.settings") + def test_extraction_provider_runtime_error(self, mock_settings, mock_get_provider, client): + """Test that runtime errors during AI call are returned as error status.""" + mock_settings.ai_provider = "openai" + mock_settings.ai_model = "gpt-4o-mini" + mock_settings.openai_model = "gpt-4o-mini" + + mock_provider = MagicMock() + mock_provider.chat_completion.side_effect = Exception("Connection refused") + mock_get_provider.return_value = mock_provider + + response = client.post("/api/ai/test-extraction", json={"text": "Some document text"}) + assert response.status_code == 200 + data = response.json() + + assert data["status"] == "error" + assert "Connection refused" in data["message"] + + @patch("app.utils.ai_provider.get_ai_provider") + @patch("app.api.openai.settings") + def test_extraction_plain_json_without_code_fences(self, mock_settings, mock_get_provider, client): + """Test that JSON returned without code fences is still parsed correctly.""" + import json as json_module + + mock_settings.ai_provider = "openai" + mock_settings.ai_model = "gpt-4o-mini" + mock_settings.openai_model = "gpt-4o-mini" + + expected_json = {"tags": ["contract"], "title": "Service Agreement"} + raw = json_module.dumps(expected_json) + + mock_provider = MagicMock() + mock_provider.chat_completion.return_value = raw + mock_get_provider.return_value = mock_provider + + response = client.post("/api/ai/test-extraction", json={"text": "Contract content"}) + assert response.status_code == 200 + data = response.json() + + assert data["status"] == "success" + assert data["tags"] == ["contract"] + assert data["parse_error"] is None