File size: 9,153 Bytes
399944f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
import os
from random import random
import re
from time import time
from typing import Any, Callable, List, Optional, TypeVar

from kbdebugger.novelty.types import NoveltyDecision
from kbdebugger.types import ExtractionResult, TripletSubjectObjectPredicate
from kbdebugger.utils.json import write_json
from kbdebugger.utils.time import now_utc_compact
from .types import Qualities
from typing import Any, Dict


def coerce_triplets(item: Dict[str, Any], fallback_sentence: str) -> ExtractionResult:
    """
    Coerce a single item dict to ExtractionResult:
    { "sentence": str, "triplets": list[TripletSubjectObjectPredicate] }
    """
    sentence = item.get("sentence", fallback_sentence)
    raw_triplets = item.get("triplets", [])
    triplets: list[TripletSubjectObjectPredicate] = []

    if isinstance(raw_triplets, list):
        for t in raw_triplets:
            if isinstance(t, (list, tuple)) and len(t) == 3:
                subj, obj, rel = t
                if all(isinstance(x, str) for x in (subj, obj, rel)):
                    triplets.append((subj.strip(), obj.strip(), rel.strip()))

    return {"sentence": str(sentence), "triplets": triplets}


def coerce_triplets_batch(obj: Dict[str, Any], sentences: List[str]) -> List[ExtractionResult]:
    """
    Coerce the LLM batch output of shape:
    {
      "triplets_batch": [
        {"id": 0, "sentence": "...", "triplets": [...]},
        ...
      ]
    }
    into a list[ExtractionResult], aligned by input index.
    """
    # Default: one empty result per input sentence
    empty: ExtractionResult = {"sentence": "", "triplets": []}
    results: List[ExtractionResult] = [empty for _ in sentences]

    batch = obj.get("triplets_batch", [])
    if not isinstance(batch, list):
        return results

    # Map by id, but also be robust
    for item in batch:
        if not isinstance(item, dict):
            continue

        idx = item.get("id")
        if isinstance(idx, int) and 0 <= idx < len(sentences):
            results[idx] = coerce_triplets(item, sentences[idx])

    # Fill any missing entries with fallback (no triplets)
    for i, res in enumerate(results):
        if res["sentence"] == "":
            results[i] = {"sentence": sentences[i], "triplets": []}

    return results


def coerce_qualities(obj: Dict) -> Qualities:
    if not isinstance(obj, dict):
        return []
    qualities = obj.get("qualities")
    if not isinstance(qualities, list):
        return []
    out: Qualities = []
    for q in qualities:
        if isinstance(q, str):
            s = q.strip()
            if s:
                out.append(s)
    return out


def save_results_json(results: List[ExtractionResult]) -> None:
    """
    Write extraction results to a JSON file.
    """
    created_at = now_utc_compact()
    data = {
        "results": results,
    }
    path = f"logs/05_triplet_extraction_results_{created_at}.json"
    write_json(path, data)
    print(f"\n[INFO] Wrote JSON results to {path}")


# ---------------------------------------------------------------------------
# Helpers for `build_chunk_batch_decomposer`
# ---------------------------------------------------------------------------
_WS_RE = re.compile(r"\s+") # this matches all whitespace sequences i.e. newlines, tabs, multiple spaces, etc.

def sanitize_chunk(text: str) -> str:
    """
    Normalize a chunk into a single-line string.

    We intentionally avoid aggressive cleaning here: the upstream PDF cleaning
    stage already handles boilerplate/DOI stripping etc. Our goal is only to
    prevent formatting artifacts from confusing the LLM.
    """
    # replace all whitespace sequences (newlines, tabs, multiple spaces) with single space " "
    return _WS_RE.sub(" ", text or "").strip()


def coerce_batch_qualities(
    obj: Any,
    *,
    expected_n: int,
) -> Dict[int, Qualities]:
    """
    Parse the JSON object returned by the batch prompt into an id->qualities map.

    Expected schema (strict, by prompt contract):
        {
          "results": [
            {"id": 0, "qualities": ["...", "..."]},
            {"id": 1, "qualities": []}
          ]
        }

    This parser is defensive:
    - Accepts "id" as int or numeric string.
    - Accepts "qualities" as list[str] or other coercible structures.
    - Ignores unknown items; only keeps ids within range.
    - Returns a possibly sparse mapping; caller fills missing ids with [].
    """
    if not isinstance(obj, dict):
        return {}

    results = obj.get("results")
    if not isinstance(results, list):
        return {}

    out: Dict[int, Qualities] = {}

    for item in results:
        if not isinstance(item, dict):
            continue

        raw_id = item.get("id")
        if raw_id is None:
            continue

        # Coerce id -> int if possible
        chunk_id: Optional[int] = None
        if isinstance(raw_id, int):
            chunk_id = raw_id
        elif isinstance(raw_id, str) and raw_id.strip().isdigit():
            chunk_id = int(raw_id.strip())

        if chunk_id is None:
            continue
        if chunk_id < 0 or chunk_id >= expected_n:
            continue

        raw_qualities = item.get("qualities", [])
        # Try to coerce qualities robustly.
        # - If it's already a list, keep string-like entries.
        # - If it's a dict (rare), attempt coerce_qualities on it.
        qualities: Qualities = []

        if isinstance(raw_qualities, list):
            qualities = [str(x).strip() for x in raw_qualities if str(x).strip()]
        else:
            # Some models might accidentally return {"qualities": [...]} per item.
            # coerce_qualities can often salvage this.
            try:
                qualities = coerce_qualities(raw_qualities)  # type: ignore[arg-type]
            except Exception:
                qualities = []

        out[chunk_id] = qualities

    return out

def load_triplet_qualifying_decisions() -> set[NoveltyDecision]:
    """
    Load which novelty decisions qualify a quality for triplet extraction.

    Environment variable:
        KB_TRIPLET_QUALIFY_DECISIONS=PARTIALLY_NEW,NEW

    Defaults to:
        {"PARTIALLY_NEW", "NEW"}
    """
    raw = os.getenv("KB_TRIPLET_QUALIFY_DECISIONS", "").strip()

    fallback = {
        NoveltyDecision.PARTIALLY_NEW,
        NoveltyDecision.NEW,
    }

    if not raw:
        return fallback
    
    decisions: set[NoveltyDecision] = set()
    for token in raw.split(","):
        token = token.strip().upper()
        if not token:
            continue
        try:
            decisions.add(NoveltyDecision(token))
        except ValueError:
            # Ignore unknown tokens silently
            continue

    # Safety fallback
    if not decisions:
        decisions = fallback

    return decisions

# ---------------------------------------------------------------------------
# Parallelism helpers
# ---------------------------------------------------------------------------
T = TypeVar("T")

_RETRY_AFTER_RE = re.compile(r"try again in\s+([0-9]*\.?[0-9]+)s", re.IGNORECASE)


def _extract_retry_after_seconds(error_text: str) -> Optional[float]:
    """
    Extract a retry delay (in seconds) from Groq-style 429 error messages.

    Example message fragment:
        "Please try again in 13.45s."

    Returns
    -------
    Optional[float]
        The parsed delay in seconds if present, otherwise None.
    """
    match = _RETRY_AFTER_RE.search(error_text or "")
    if not match:
        return None

    try:
        return float(match.group(1))
    except ValueError:
        return None
    

def _call_with_rate_limit_retries(
    fn: Callable[[], T],
    *,
    max_retries: int = 8,
    default_backoff_s: float = 2.0,
    max_sleep_s: float = 30.0,
) -> T:
    """
    Call `fn()` with rate-limit-aware retries.

    Strategy
    --------
    - If the exception message contains "try again in Xs", sleep for X seconds
      (+ small jitter) and retry.
    - Otherwise, use a conservative exponential backoff.

    Why this exists
    ---------------
    Groq on-demand has strict TPM (Token-per-Minute) limits. When we batch or parallelize,
    occasional 429s are expected. Dropping a batch silently corrupts results.

    Raises
    ------
    RuntimeError
        If all retry attempts fail.
    """
    last_err: Exception | None = None

    for attempt in range(1, max_retries + 1):
        try:
            return fn()
        except Exception as e:  # SDKs often raise generic exceptions
            last_err = e
            msg = str(e)

            retry_after = _extract_retry_after_seconds(msg)
            if retry_after is not None:
                # Add a tiny jitter to avoid synchronizing retries across threads.
                sleep_s = retry_after + random.uniform(0.1, 0.4)
            else:
                # Exponential backoff for unknown transient failures
                sleep_s = default_backoff_s * (2 ** (attempt - 1))

            sleep_s = min(sleep_s, max_sleep_s)
            time.sleep(sleep_s)

    raise RuntimeError(f"LLM call failed after {max_retries} retries") from last_err