DocDoeAI / app /services /chunking.py
asnannp's picture
Deploy backend cd4237ff: support routes + rate limit + exam_date nullable + upload 413 fix
7c6ffa6
Raw
History Blame Contribute Delete
6.12 kB
from __future__ import annotations
import asyncio
import logging
import re
from sqlalchemy import delete
from sqlalchemy.orm import Session
from app.models.document import Document
from app.models.document_chunk import DocumentChunk
logger = logging.getLogger(__name__)
def clean_extracted_text(text: str) -> str:
normalized = remove_invalid_unicode(text).replace("\x00", " ")
normalized = re.sub(r"[ \t]+", " ", normalized)
normalized = re.sub(r"\n{3,}", "\n\n", normalized)
normalized = "\n".join(line.strip() for line in normalized.splitlines())
return normalized.strip()
def remove_invalid_unicode(text: str) -> str:
"""Drop code points Postgres/UTF-8 cannot store, such as lone surrogates."""
return text.encode("utf-8", errors="ignore").decode("utf-8")
def estimate_tokens(text: str) -> int:
# A practical rough estimate for English-heavy study material.
return max(1, round(len(text) / 4))
def split_text_into_chunks(
text: str,
max_chars: int = 1200,
overlap_chars: int = 150,
) -> list[str]:
cleaned = clean_extracted_text(text)
if not cleaned:
return []
paragraphs = [paragraph.strip() for paragraph in cleaned.split("\n\n") if paragraph.strip()]
units = _paragraphs_to_units(paragraphs, max_chars)
chunks: list[str] = []
current = ""
for unit in units:
candidate = _join_chunk_parts(current, unit)
if len(candidate) <= max_chars:
current = candidate
continue
if current:
chunks.append(current)
current = _join_chunk_parts(_overlap_tail(current, overlap_chars), unit)
else:
chunks.append(unit[:max_chars].strip())
current = unit[max_chars - overlap_chars :].strip()
while len(current) > max_chars:
chunks.append(current[:max_chars].strip())
current = current[max_chars - overlap_chars :].strip()
if current:
chunks.append(current.strip())
return [chunk for chunk in chunks if chunk]
def infer_chunk_page_number(chunk_text: str) -> int | None:
match = re.search(r"\bPage\s+(\d+)\b", chunk_text[:80], flags=re.IGNORECASE)
if match:
return int(match.group(1))
return None
def infer_chunk_heading(chunk_text: str) -> str | None:
first_line = next((line.strip() for line in chunk_text.splitlines() if line.strip()), "")
if not first_line:
return None
if first_line.lower().startswith("page ") and len(chunk_text.splitlines()) > 1:
first_line = chunk_text.splitlines()[1].strip()
return first_line[:120]
def replace_document_chunks(
db: Session,
document: Document,
max_chars: int = 1200,
overlap_chars: int = 150,
) -> list[DocumentChunk]:
db.execute(delete(DocumentChunk).where(DocumentChunk.document_id == document.id))
chunks = split_text_into_chunks(
document.extracted_text or "",
max_chars=max_chars,
overlap_chars=overlap_chars,
)
chunk_records = [
DocumentChunk(
document_id=document.id,
chunk_index=index,
chunk_text=chunk_text,
token_estimate=estimate_tokens(chunk_text),
page_number=infer_chunk_page_number(chunk_text),
heading=infer_chunk_heading(chunk_text),
)
for index, chunk_text in enumerate(chunks)
]
db.add_all(chunk_records)
# Generate embeddings for all chunks (non-blocking, failure-safe)
try:
from app.services.embedding_service import generate_embeddings, embedding_to_str
texts = [c.chunk_text for c in chunk_records]
embeddings = asyncio.run(generate_embeddings(texts))
for chunk_rec, emb in zip(chunk_records, embeddings):
chunk_rec.embedding = embedding_to_str(emb)
except Exception as exc:
from app.core.config import get_settings
settings = get_settings()
if settings.environment == "production" or not settings.ai_fallback_to_mock:
raise exc
logger.warning("Embedding generation failed (chunks saved without vectors): %s", exc)
document.chunk_count = len(chunk_records)
db.add(document)
db.flush()
return chunk_records
def _paragraphs_to_units(paragraphs: list[str], max_chars: int) -> list[str]:
units: list[str] = []
for paragraph in paragraphs:
if len(paragraph) <= max_chars:
units.append(paragraph)
continue
units.extend(_split_long_paragraph(paragraph, max_chars))
return units
def _split_long_paragraph(paragraph: str, max_chars: int) -> list[str]:
sentences = re.split(r"(?<=[.!?])\s+", paragraph)
units: list[str] = []
current = ""
for sentence in sentences:
if len(sentence) > max_chars:
if current:
units.append(current)
current = ""
units.extend(_hard_split(sentence, max_chars))
continue
candidate = _join_chunk_parts(current, sentence)
if len(candidate) <= max_chars:
current = candidate
else:
if current:
units.append(current)
current = sentence
if current:
units.append(current)
return units
def _hard_split(text: str, max_chars: int) -> list[str]:
parts: list[str] = []
start = 0
while start < len(text):
parts.append(text[start : start + max_chars].strip())
start += max_chars
return [part for part in parts if part]
def _join_chunk_parts(left: str, right: str) -> str:
if not left:
return right.strip()
if not right:
return left.strip()
return f"{left.strip()}\n\n{right.strip()}"
def _overlap_tail(text: str, overlap_chars: int) -> str:
if overlap_chars <= 0:
return ""
if len(text) <= overlap_chars:
return text
tail = text[-overlap_chars:]
sentence_boundary = max(tail.find(". "), tail.find("? "), tail.find("! "))
if sentence_boundary > 0:
return tail[sentence_boundary + 2 :].strip()
return tail.strip()