Merge pull request #169 from ggozad/fix/docling-fixes

Fix docling OCR options.
This commit is contained in:
Yiorgis Gozadinos 2025-11-27 15:24:02 +02:00 committed by GitHub
commit 5d00bbf6dd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 78 additions and 48 deletions

View file

@ -1,6 +1,14 @@
# Changelog # Changelog
## [Unreleased] ## [Unreleased]
### Changed
- **Async Chunker**: `DoclingServeChunker` now uses `httpx.AsyncClient` instead of sync `requests`
### Fixed
- **OCR Options**: Fixed `DoclingLocalConverter` using base `OcrOptions` class which docling's OCR factory doesn't recognize. Now uses `OcrAutoOptions` for automatic OCR engine selection.
## [0.19.2] - 2025-11-27 ## [0.19.2] - 2025-11-27
### Changed ### Changed

View file

@ -1,7 +1,7 @@
from io import BytesIO from io import BytesIO
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
import requests import httpx
from haiku.rag.chunkers.base import DocumentChunker from haiku.rag.chunkers.base import DocumentChunker
from haiku.rag.config import AppConfig, Config from haiku.rag.config import AppConfig, Config
@ -75,33 +75,31 @@ class DoclingServeChunker(DocumentChunker):
if self.api_key: if self.api_key:
headers["X-Api-Key"] = self.api_key headers["X-Api-Key"] = self.api_key
response = requests.post( async with httpx.AsyncClient(timeout=self.timeout) as client:
url, response = await client.post(
files=files, url,
data=data, files=files,
headers=headers, data=data,
timeout=self.timeout, headers=headers,
) )
response.raise_for_status()
response.raise_for_status() result = response.json()
result = response.json()
# Extract text from chunks # Extract text from chunks
chunks = result.get("chunks", []) chunks = result.get("chunks", [])
return [chunk["text"] for chunk in chunks] return [chunk["text"] for chunk in chunks]
except requests.exceptions.ConnectionError as e: except httpx.ConnectError as e:
raise ValueError( raise ValueError(
f"Could not connect to docling-serve at {self.base_url}. " f"Could not connect to docling-serve at {self.base_url}. "
f"Ensure the service is running and accessible. Error: {e}" f"Ensure the service is running and accessible. Error: {e}"
) )
except requests.exceptions.Timeout as e: except httpx.TimeoutException as e:
raise ValueError( raise ValueError(
f"Request to docling-serve timed out after {self.timeout}s. " f"Request to docling-serve timed out after {self.timeout}s. "
f"Consider increasing the timeout in configuration. Error: {e}" f"Consider increasing the timeout in configuration. Error: {e}"
) )
except requests.exceptions.HTTPError as e: except httpx.HTTPStatusError as e:
if e.response.status_code == 401: if e.response.status_code == 401:
raise ValueError( raise ValueError(
"Authentication failed. Check your API key configuration." "Authentication failed. Check your API key configuration."

View file

@ -59,7 +59,7 @@ class DoclingLocalConverter(DocumentConverter):
from docling.backend.docling_parse_backend import DoclingParseDocumentBackend from docling.backend.docling_parse_backend import DoclingParseDocumentBackend
from docling.datamodel.base_models import InputFormat from docling.datamodel.base_models import InputFormat
from docling.datamodel.pipeline_options import ( from docling.datamodel.pipeline_options import (
OcrOptions, OcrAutoOptions,
PdfPipelineOptions, PdfPipelineOptions,
TableFormerMode, TableFormerMode,
TableStructureOptions, TableStructureOptions,
@ -86,7 +86,7 @@ class DoclingLocalConverter(DocumentConverter):
else TableFormerMode.ACCURATE else TableFormerMode.ACCURATE
), ),
), ),
ocr_options=OcrOptions( ocr_options=OcrAutoOptions(
force_full_page_ocr=opts.force_ocr, force_full_page_ocr=opts.force_ocr,
lang=opts.ocr_lang if opts.ocr_lang else [], lang=opts.ocr_lang if opts.ocr_lang else [],
), ),

View file

@ -1,4 +1,4 @@
from unittest.mock import Mock, patch from unittest.mock import AsyncMock, Mock, patch
import pytest import pytest
from datasets import Dataset from datasets import Dataset
@ -170,8 +170,8 @@ class TestDoclingServeChunker:
return DoclingServeChunker(config) return DoclingServeChunker(config)
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("haiku.rag.chunkers.docling_serve.requests.post") @patch("haiku.rag.chunkers.docling_serve.httpx.AsyncClient")
async def test_chunk_success(self, mock_post, chunker): async def test_chunk_success(self, mock_client_class, chunker):
"""Test successful chunking via docling-serve.""" """Test successful chunking via docling-serve."""
mock_response = Mock() mock_response = Mock()
mock_response.status_code = 200 mock_response.status_code = 200
@ -181,7 +181,11 @@ class TestDoclingServeChunker:
{"text": "Chunk 2", "chunk_index": 1}, {"text": "Chunk 2", "chunk_index": 1},
] ]
} }
mock_post.return_value = mock_response mock_response.raise_for_status = Mock()
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
mock_client_class.return_value.__aenter__.return_value = mock_client
# Create a simple document # Create a simple document
converter = get_converter(Config) converter = get_converter(Config)
@ -191,11 +195,11 @@ class TestDoclingServeChunker:
assert len(chunks) == 2 assert len(chunks) == 2
assert chunks[0] == "Chunk 1" assert chunks[0] == "Chunk 1"
assert chunks[1] == "Chunk 2" assert chunks[1] == "Chunk 2"
mock_post.assert_called_once() mock_client.post.assert_called_once()
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("haiku.rag.chunkers.docling_serve.requests.post") @patch("haiku.rag.chunkers.docling_serve.httpx.AsyncClient")
async def test_chunk_with_api_key(self, mock_post, config): async def test_chunk_with_api_key(self, mock_client_class, config):
"""Test that API key is included in request headers.""" """Test that API key is included in request headers."""
config.providers.docling_serve.api_key = "test-key" config.providers.docling_serve.api_key = "test-key"
chunker = DoclingServeChunker(config) chunker = DoclingServeChunker(config)
@ -205,19 +209,23 @@ class TestDoclingServeChunker:
mock_response.json.return_value = { mock_response.json.return_value = {
"chunks": [{"text": "Chunk 1", "chunk_index": 0}] "chunks": [{"text": "Chunk 1", "chunk_index": 0}]
} }
mock_post.return_value = mock_response mock_response.raise_for_status = Mock()
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
mock_client_class.return_value.__aenter__.return_value = mock_client
converter = get_converter(Config) converter = get_converter(Config)
doc = await converter.convert_text("# Test", name="test.md") doc = await converter.convert_text("# Test", name="test.md")
await chunker.chunk(doc) await chunker.chunk(doc)
call_kwargs = mock_post.call_args.kwargs call_kwargs = mock_client.post.call_args.kwargs
assert "headers" in call_kwargs assert "headers" in call_kwargs
assert call_kwargs["headers"]["X-Api-Key"] == "test-key" assert call_kwargs["headers"]["X-Api-Key"] == "test-key"
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("haiku.rag.chunkers.docling_serve.requests.post") @patch("haiku.rag.chunkers.docling_serve.httpx.AsyncClient")
async def test_chunk_hierarchical_endpoint(self, mock_post, config): async def test_chunk_hierarchical_endpoint(self, mock_client_class, config):
"""Test that hierarchical chunker uses correct endpoint.""" """Test that hierarchical chunker uses correct endpoint."""
config.processing.chunker_type = "hierarchical" config.processing.chunker_type = "hierarchical"
chunker = DoclingServeChunker(config) chunker = DoclingServeChunker(config)
@ -227,18 +235,22 @@ class TestDoclingServeChunker:
mock_response.json.return_value = { mock_response.json.return_value = {
"chunks": [{"text": "Chunk 1", "chunk_index": 0}] "chunks": [{"text": "Chunk 1", "chunk_index": 0}]
} }
mock_post.return_value = mock_response mock_response.raise_for_status = Mock()
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
mock_client_class.return_value.__aenter__.return_value = mock_client
converter = get_converter(Config) converter = get_converter(Config)
doc = await converter.convert_text("# Test", name="test.md") doc = await converter.convert_text("# Test", name="test.md")
await chunker.chunk(doc) await chunker.chunk(doc)
call_args = mock_post.call_args call_args = mock_client.post.call_args
assert "/v1/chunk/hierarchical/file" in call_args[0][0] assert "/v1/chunk/hierarchical/file" in call_args[0][0]
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("haiku.rag.chunkers.docling_serve.requests.post") @patch("haiku.rag.chunkers.docling_serve.httpx.AsyncClient")
async def test_chunk_passes_config_parameters(self, mock_post, config): async def test_chunk_passes_config_parameters(self, mock_client_class, config):
"""Test that all config parameters are passed to API.""" """Test that all config parameters are passed to API."""
config.processing.chunk_size = 512 config.processing.chunk_size = 512
config.processing.chunking_merge_peers = False config.processing.chunking_merge_peers = False
@ -250,25 +262,31 @@ class TestDoclingServeChunker:
mock_response.json.return_value = { mock_response.json.return_value = {
"chunks": [{"text": "Chunk 1", "chunk_index": 0}] "chunks": [{"text": "Chunk 1", "chunk_index": 0}]
} }
mock_post.return_value = mock_response mock_response.raise_for_status = Mock()
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
mock_client_class.return_value.__aenter__.return_value = mock_client
converter = get_converter(Config) converter = get_converter(Config)
doc = await converter.convert_text("# Test", name="test.md") doc = await converter.convert_text("# Test", name="test.md")
await chunker.chunk(doc) await chunker.chunk(doc)
call_kwargs = mock_post.call_args.kwargs call_kwargs = mock_client.post.call_args.kwargs
data = call_kwargs["data"] data = call_kwargs["data"]
assert data["chunking_max_tokens"] == "512" assert data["chunking_max_tokens"] == "512"
assert data["chunking_merge_peers"] == "false" assert data["chunking_merge_peers"] == "false"
assert data["chunking_use_markdown_tables"] == "true" assert data["chunking_use_markdown_tables"] == "true"
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("haiku.rag.chunkers.docling_serve.requests.post") @patch("haiku.rag.chunkers.docling_serve.httpx.AsyncClient")
async def test_chunk_connection_error(self, mock_post, chunker): async def test_chunk_connection_error(self, mock_client_class, chunker):
"""Test handling of connection errors.""" """Test handling of connection errors."""
import requests import httpx
mock_post.side_effect = requests.exceptions.ConnectionError("Connection failed") mock_client = AsyncMock()
mock_client.post.side_effect = httpx.ConnectError("Connection failed")
mock_client_class.return_value.__aenter__.return_value = mock_client
converter = get_converter(Config) converter = get_converter(Config)
doc = await converter.convert_text("# Test", name="test.md") doc = await converter.convert_text("# Test", name="test.md")
@ -277,12 +295,14 @@ class TestDoclingServeChunker:
await chunker.chunk(doc) await chunker.chunk(doc)
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("haiku.rag.chunkers.docling_serve.requests.post") @patch("haiku.rag.chunkers.docling_serve.httpx.AsyncClient")
async def test_chunk_timeout_error(self, mock_post, chunker): async def test_chunk_timeout_error(self, mock_client_class, chunker):
"""Test handling of timeout errors.""" """Test handling of timeout errors."""
import requests import httpx
mock_post.side_effect = requests.exceptions.Timeout("Timeout") mock_client = AsyncMock()
mock_client.post.side_effect = httpx.TimeoutException("Timeout")
mock_client_class.return_value.__aenter__.return_value = mock_client
converter = get_converter(Config) converter = get_converter(Config)
doc = await converter.convert_text("# Test", name="test.md") doc = await converter.convert_text("# Test", name="test.md")
@ -291,17 +311,21 @@ class TestDoclingServeChunker:
await chunker.chunk(doc) await chunker.chunk(doc)
@pytest.mark.asyncio @pytest.mark.asyncio
@patch("haiku.rag.chunkers.docling_serve.requests.post") @patch("haiku.rag.chunkers.docling_serve.httpx.AsyncClient")
async def test_chunk_auth_error(self, mock_post, chunker): async def test_chunk_auth_error(self, mock_client_class, chunker):
"""Test handling of authentication errors.""" """Test handling of authentication errors."""
import requests import httpx
mock_request = Mock()
mock_response = Mock() mock_response = Mock()
mock_response.status_code = 401 mock_response.status_code = 401
mock_response.raise_for_status.side_effect = requests.exceptions.HTTPError( mock_response.raise_for_status.side_effect = httpx.HTTPStatusError(
response=mock_response "401", request=mock_request, response=mock_response
) )
mock_post.return_value = mock_response
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
mock_client_class.return_value.__aenter__.return_value = mock_client
converter = get_converter(Config) converter = get_converter(Config)
doc = await converter.convert_text("# Test", name="test.md") doc = await converter.convert_text("# Test", name="test.md")