from __future__ import annotations import math import os from typing import List, Optional, Sequence, Any, Tuple from concurrent.futures import ThreadPoolExecutor, as_completed from kbdebugger.types.ui import ProgressCallback from rich.progress import track from kbdebugger.compat.langchain import Document from kbdebugger.utils import batched from .sentence_to_qualities import build_sentence_decomposer from .chunk_to_qualities import build_chunk_decomposer, build_chunk_batch_decomposer from .types import Qualities, TextDecomposer, BatchTextDecomposer, DecomposeMode from .logging import save_qualities_json # --------------------------------------------------------------------------- # Module-level decomposer singletons # --------------------------------------------------------------------------- # These are initialized once at import time to avoid re-loading prompt resources # and few-shot examples repeatedly inside tight loops. _sentence_to_qualities_decomposer: TextDecomposer = build_sentence_decomposer() _chunk_to_qualities_decomposer: TextDecomposer = build_chunk_decomposer() _chunk_batch_to_qualities_decomposer: BatchTextDecomposer = build_chunk_batch_decomposer() def decompose( text: str, *, mode: DecomposeMode ) -> Qualities: """ Decompose a single input text into "qualities" under the selected mode. Parameters ---------- text: Input text: either a single sentence or a larger chunk. mode: - DecomposeMode.SENTENCES: Use when `text` is already sentence-like but may contain multiple atomic statements that should be split. Example: "The cat sat on the mat and looked at the dog." → ["The cat sat on the mat.", "The cat looked at the dog."] - DecomposeMode.CHUNKS: Use when `text` is a larger paragraph or chunk and you want to extract key qualities / statements. Example: "Cats are great pets. They are independent and curious animals..." → ["Cats are great pets.", "Cats are independent animals.", "Cats are curious animals."] Returns ------- list[str] A list of short, atomic sentences (qualities). """ match mode: case DecomposeMode.SENTENCES: return _sentence_to_qualities_decomposer(text) case DecomposeMode.CHUNKS: return _chunk_to_qualities_decomposer(text) case _: pass # Defensive: this should never happen with the Enum, but keeps mypy happy raise ValueError(f"Unsupported DecomposeMode: {mode}") def _decompose_one_batch( batch_id: int, group: List[str], ) -> Tuple[int, List[Qualities]]: """ Worker wrapper for parallel batch decomposition. Returns ------- (batch_id, batch_results) batch_id is used to optionally re-order results deterministically. """ return batch_id, _chunk_batch_to_qualities_decomposer(group) def _safe_chunk_batch_to_qualities_decomposer(group: List[str]) -> List[Qualities]: """ Safe wrapper around the batched decomposer. Why this exists --------------- `ThreadPoolExecutor.map()` will propagate exceptions and stop iteration on the first failure. For a pipeline stage, it's usually better to be *best-effort* and preserve output alignment. Contract -------- Returns `List[Qualities]` aligned with `group` length: - one Qualities list per input chunk text - on failure: returns `[[], [], ...]` (same length as group) """ try: return _chunk_batch_to_qualities_decomposer(group) except Exception as e: # noqa: BLE001 (intentionally broad in pipeline boundary) print(f"[decompose_documents] Batch failed (size={len(group)}): {e}") return [[] for _ in range(len(group))] # --------------------------------------------------------------------------- # Public API # --------------------------------------------------------------------------- def decompose_documents( docs: Sequence[Document], *, mode: DecomposeMode, batch_size: int = 5, use_batch_decomposer: bool = True, parallel: bool = False, max_workers: Optional[int] = 2, progress: Optional[ProgressCallback] = None ) -> Tuple[Qualities, dict]: """ Decompose a list of LangChain Documents into a flat list of qualities. Returns ------- (qualities, log_payload) - qualities: flat list of atomic qualities - log_payload: the same payload that was written to disk """ all_qualities: Qualities = [] if not docs: log_payload = save_qualities_json( qualities=all_qualities, mode=mode, num_input_docs=0, use_batch_decomposer=use_batch_decomposer, batch_size=batch_size if use_batch_decomposer else None, num_batches=0 if use_batch_decomposer else None, parallel=parallel, max_workers=max_workers if parallel else None, ) return all_qualities, log_payload texts: List[str] = [getattr(doc, "page_content", "") for doc in docs] # --- Fast path: batched chunk decomposition --- if mode == DecomposeMode.CHUNKS and use_batch_decomposer: num_batches = math.ceil(len(texts) / batch_size) if parallel: with ThreadPoolExecutor(max_workers=max_workers) as pool: results_iter = pool.map( _safe_chunk_batch_to_qualities_decomposer, batched(texts, batch_size=batch_size), ) for batch_idx, group_results in track( enumerate(results_iter), total=num_batches, description=( f"🧷 LLM Decomposer (parallel): paragraphs → qualities " f"(num_batches={num_batches}, batch size={batch_size})" ), ): if progress: progress( batch_idx + 1, # nicer: 1-based progress num_batches, f"🧷 LLM Decomposer (parallel): Processing batch ({batch_idx+1}/{num_batches}) ..." ) for qualities in group_results: all_qualities.extend(qualities) else: for batch_idx, group in track( enumerate(batched(texts, batch_size=batch_size)), total=num_batches, description=( f"🧷 LLM Decomposer: paragraphs → qualities " f"(num_batches={num_batches}, batch size={batch_size})" ), ): if progress: progress( batch_idx + 1, num_batches, f"🧷 LLM Decomposer: Processing batch ({batch_idx+1}/{num_batches}) ..." ) group_results: List[Qualities] = _chunk_batch_to_qualities_decomposer(group) for qualities in group_results: all_qualities.extend(qualities) log_payload = save_qualities_json( qualities=all_qualities, mode=mode, num_input_docs=len(docs), use_batch_decomposer=True, batch_size=batch_size, num_batches=num_batches, parallel=parallel, max_workers=max_workers if parallel else None, ) return all_qualities, log_payload # --- Default path: one document -> one decompose() call --- for text in texts: qualities = decompose(text, mode=mode) all_qualities.extend(qualities) log_payload = save_qualities_json( qualities=all_qualities, mode=mode, num_input_docs=len(docs), use_batch_decomposer=False, batch_size=None, num_batches=None, parallel=False, max_workers=None, ) return all_qualities, log_payload