rag-wright 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (184) hide show
  1. rag_wright/__init__.py +13 -0
  2. rag_wright/api/__init__.py +33 -0
  3. rag_wright/api/config.py +59 -0
  4. rag_wright/api/discover.py +70 -0
  5. rag_wright/api/documents.py +39 -0
  6. rag_wright/api/ids.py +31 -0
  7. rag_wright/api/invoke.py +99 -0
  8. rag_wright/api/kg.py +61 -0
  9. rag_wright/api/mcp.py +94 -0
  10. rag_wright/api/usage.py +30 -0
  11. rag_wright/api/workspace.py +85 -0
  12. rag_wright/capabilities/__init__.py +8 -0
  13. rag_wright/capabilities/answer_generator.py +427 -0
  14. rag_wright/capabilities/ard.py +286 -0
  15. rag_wright/capabilities/assertion_extraction.py +79 -0
  16. rag_wright/capabilities/chunk_read.py +58 -0
  17. rag_wright/capabilities/chunk_write.py +163 -0
  18. rag_wright/capabilities/claim_extraction.py +153 -0
  19. rag_wright/capabilities/clause_exception_linking.py +117 -0
  20. rag_wright/capabilities/compliance_judgment.py +322 -0
  21. rag_wright/capabilities/compliance_store.py +87 -0
  22. rag_wright/capabilities/contract_kg_serve.py +156 -0
  23. rag_wright/capabilities/contract_kg_store.py +251 -0
  24. rag_wright/capabilities/dg_extraction.py +585 -0
  25. rag_wright/capabilities/disambiguation.py +163 -0
  26. rag_wright/capabilities/document_parse.py +87 -0
  27. rag_wright/capabilities/document_scope.py +49 -0
  28. rag_wright/capabilities/embedding.py +164 -0
  29. rag_wright/capabilities/embedding_profiles.py +43 -0
  30. rag_wright/capabilities/entity_resolution.py +154 -0
  31. rag_wright/capabilities/fusion.py +64 -0
  32. rag_wright/capabilities/graph_extraction.py +243 -0
  33. rag_wright/capabilities/graph_query.py +73 -0
  34. rag_wright/capabilities/graph_storage.py +111 -0
  35. rag_wright/capabilities/highlight_serve.py +142 -0
  36. rag_wright/capabilities/hybrid_search.py +65 -0
  37. rag_wright/capabilities/invoke.py +31 -0
  38. rag_wright/capabilities/jev_decision.py +38 -0
  39. rag_wright/capabilities/manifests.py +872 -0
  40. rag_wright/capabilities/okf_navigate.py +456 -0
  41. rag_wright/capabilities/parsing.py +286 -0
  42. rag_wright/capabilities/property_boosted_retrieval.py +125 -0
  43. rag_wright/capabilities/query_function_classifier.py +94 -0
  44. rag_wright/capabilities/query_understanding.py +109 -0
  45. rag_wright/capabilities/registry.py +262 -0
  46. rag_wright/capabilities/remote_encoders.py +94 -0
  47. rag_wright/capabilities/requirement_extraction.py +247 -0
  48. rag_wright/capabilities/reranking.py +123 -0
  49. rag_wright/capabilities/retrieval_core.py +126 -0
  50. rag_wright/capabilities/rlm_chunking.py +808 -0
  51. rag_wright/capabilities/rlm_synthesis.py +316 -0
  52. rag_wright/capabilities/scan_quality.py +136 -0
  53. rag_wright/capabilities/span_relevance_judgment.py +191 -0
  54. rag_wright/capabilities/vision_to_text.py +85 -0
  55. rag_wright/capabilities/vlm_ocr.py +85 -0
  56. rag_wright/contracts/__init__.py +6 -0
  57. rag_wright/contracts/chunk.py +79 -0
  58. rag_wright/contracts/compliance.py +303 -0
  59. rag_wright/contracts/contract_meta.py +27 -0
  60. rag_wright/contracts/extraction.py +130 -0
  61. rag_wright/contracts/function.py +167 -0
  62. rag_wright/contracts/function_routing.py +91 -0
  63. rag_wright/contracts/highlight.py +74 -0
  64. rag_wright/contracts/identifiers.py +153 -0
  65. rag_wright/contracts/jurisdiction.py +96 -0
  66. rag_wright/contracts/ontology.py +142 -0
  67. rag_wright/contracts/property.py +201 -0
  68. rag_wright/contracts/provenance.py +78 -0
  69. rag_wright/contracts/query_intent.py +53 -0
  70. rag_wright/contracts/span.py +76 -0
  71. rag_wright/contracts/value_match.py +84 -0
  72. rag_wright/corpus/__init__.py +0 -0
  73. rag_wright/corpus/canonicalize.py +116 -0
  74. rag_wright/corpus/cuad.py +153 -0
  75. rag_wright/corpus/cuad_ingestion.py +72 -0
  76. rag_wright/corpus/document_parser.py +299 -0
  77. rag_wright/corpus/edgar.py +231 -0
  78. rag_wright/corpus/gcs_ingestion.py +120 -0
  79. rag_wright/corpus/http.py +110 -0
  80. rag_wright/corpus/selection.py +152 -0
  81. rag_wright/mcp/__init__.py +11 -0
  82. rag_wright/mcp/compliance_server.py +299 -0
  83. rag_wright/mcp/intra_document_qa_server.py +170 -0
  84. rag_wright/mcp/relational_qa_server.py +171 -0
  85. rag_wright/mcp/session_store.py +64 -0
  86. rag_wright/mcp/typed_property_retrieval_server.py +191 -0
  87. rag_wright/models/__init__.py +8 -0
  88. rag_wright/models/profiles.py +331 -0
  89. rag_wright/models/seam.py +497 -0
  90. rag_wright/models/tag_structured.py +285 -0
  91. rag_wright/models/tracing.py +179 -0
  92. rag_wright/models/usage.py +102 -0
  93. rag_wright/okf/__init__.py +11 -0
  94. rag_wright/okf/compile.py +292 -0
  95. rag_wright/okf/document.py +47 -0
  96. rag_wright/okf/enrich.py +176 -0
  97. rag_wright/okf/links.py +190 -0
  98. rag_wright/okf/lint.py +105 -0
  99. rag_wright/ontology/__init__.py +6 -0
  100. rag_wright/ontology/_generated_template_meta.py +60 -0
  101. rag_wright/ontology/_generated_vocab.py +52 -0
  102. rag_wright/ontology/clause_template.py +964 -0
  103. rag_wright/ontology/codegen.py +84 -0
  104. rag_wright/ontology/compliance_bridge.ttl +186 -0
  105. rag_wright/ontology/contract_bridge.ttl +2685 -0
  106. rag_wright/ontology/contract_taxonomy.py +24 -0
  107. rag_wright/ontology/derive.py +58 -0
  108. rag_wright/ontology/loader.py +435 -0
  109. rag_wright/ontology/packs/ftc_16cfr255.ttl +29 -0
  110. rag_wright/ontology/registry.py +87 -0
  111. rag_wright/ontology/template_introspect.py +100 -0
  112. rag_wright/py.typed +0 -0
  113. rag_wright/reference/__init__.py +2 -0
  114. rag_wright/reference/compliance.py +41 -0
  115. rag_wright/reference/contract_seam.py +123 -0
  116. rag_wright/skills/__init__.py +7 -0
  117. rag_wright/skills/claim_extraction/SKILL.md +47 -0
  118. rag_wright/skills/claim_extraction/__init__.py +1 -0
  119. rag_wright/skills/claim_extraction/template.py +50 -0
  120. rag_wright/skills/compliance_judgment/SKILL.md +59 -0
  121. rag_wright/skills/corpus_ingest/SKILL.md +106 -0
  122. rag_wright/skills/extraction_semantic_judge/SKILL.md +51 -0
  123. rag_wright/skills/extraction_semantic_judge/__init__.py +1 -0
  124. rag_wright/skills/generation/SKILL.md +64 -0
  125. rag_wright/skills/generation/__init__.py +1 -0
  126. rag_wright/skills/generic_compliance_judgment/SKILL.md +58 -0
  127. rag_wright/skills/okf_navigate/SKILL.md +137 -0
  128. rag_wright/skills/requirement_extraction/SKILL.md +47 -0
  129. rag_wright/skills/requirement_extraction/__init__.py +1 -0
  130. rag_wright/skills/requirement_extraction/template.py +50 -0
  131. rag_wright/skills/rlm/SKILL.md +186 -0
  132. rag_wright/skills/rlm/__init__.py +31 -0
  133. rag_wright/skills/rlm/agent.py +292 -0
  134. rag_wright/skills/span_relevance_judgment/SKILL.md +67 -0
  135. rag_wright/skills/vision_to_text/SKILL.md +36 -0
  136. rag_wright/skills/vision_to_text/__init__.py +1 -0
  137. rag_wright/spans/__init__.py +1 -0
  138. rag_wright/spans/boundary.py +78 -0
  139. rag_wright/spans/clause_function_classifier.py +490 -0
  140. rag_wright/spans/clause_kg_extractor.py +337 -0
  141. rag_wright/spans/cuad_labels.py +81 -0
  142. rag_wright/spans/dim_classifier.py +158 -0
  143. rag_wright/spans/dim_fleet.json +411 -0
  144. rag_wright/spans/function_classifier.py +77 -0
  145. rag_wright/spans/function_families.py +62 -0
  146. rag_wright/spans/hybrid_classifier.py +103 -0
  147. rag_wright/spans/legalbert_classifier.py +83 -0
  148. rag_wright/spans/model_capabilities.py +107 -0
  149. rag_wright/spans/new_function_labels.py +111 -0
  150. rag_wright/spans/page_map.py +68 -0
  151. rag_wright/spans/property_extractor.py +365 -0
  152. rag_wright/spans/property_grounding.py +182 -0
  153. rag_wright/spans/reclassify.py +77 -0
  154. rag_wright/spans/scarce_function_labels.py +105 -0
  155. rag_wright/spans/segment.py +341 -0
  156. rag_wright/spans/semantic_judge.py +197 -0
  157. rag_wright/spans/symbolic_validation.py +131 -0
  158. rag_wright/spans/tag_clause_extractor.py +182 -0
  159. rag_wright/store/__init__.py +6 -0
  160. rag_wright/store/arcadedb.py +1135 -0
  161. rag_wright/store/chunk_text.py +66 -0
  162. rag_wright/store/seam.py +213 -0
  163. rag_wright/subgraphs/__init__.py +0 -0
  164. rag_wright/subgraphs/async_ingestion.py +204 -0
  165. rag_wright/subgraphs/compliance_check.py +1042 -0
  166. rag_wright/subgraphs/compliance_ingestion.py +306 -0
  167. rag_wright/subgraphs/contract_ingestion_pipeline.py +999 -0
  168. rag_wright/subgraphs/graph_extraction.py +102 -0
  169. rag_wright/subgraphs/intra_document_qa.py +328 -0
  170. rag_wright/subgraphs/observability.py +140 -0
  171. rag_wright/subgraphs/query_constraint_extraction.py +73 -0
  172. rag_wright/subgraphs/relational_qa.py +165 -0
  173. rag_wright/subgraphs/requirement_extraction.py +137 -0
  174. rag_wright/subgraphs/scaffold.py +65 -0
  175. rag_wright/subgraphs/semantic_chunking.py +183 -0
  176. rag_wright/subgraphs/typed_clause_extraction.py +172 -0
  177. rag_wright/subgraphs/typed_property_retrieval.py +278 -0
  178. rag_wright/util/__init__.py +1 -0
  179. rag_wright/util/concurrent.py +153 -0
  180. rag_wright/util/spacy_model.py +45 -0
  181. rag_wright-0.1.0.dist-info/METADATA +168 -0
  182. rag_wright-0.1.0.dist-info/RECORD +184 -0
  183. rag_wright-0.1.0.dist-info/WHEEL +4 -0
  184. rag_wright-0.1.0.dist-info/licenses/LICENSE +21 -0
@@ -0,0 +1,490 @@
1
+ """INGEST-LLM-CLASSIFIER (ADR-0048): the clause-level function-classification seam.
2
+
3
+ Replaces the span-level, single-label LegalBERT call at ingestion with a FULL-CLAUSE, MULTI-LABEL, confidence-
4
+ scored classifier by the graph-building LLM. Two implementations behind one `classify(clause_text) -> [FunctionScore]`
5
+ seam: `LlmClauseClassifier` (the default going forward, via the model-profile structured seam) and
6
+ `LegalBertClauseAdapter` (wraps the existing single-label classifier for back-compat). Injected into
7
+ `production_document_ingest` (a later step); the injection point query-time already had.
8
+
9
+ Robustness: the post-processor canonicalizes the LLM's labels, drops off-taxonomy / below-floor labels, and caps
10
+ the count -- so LLM slop never crashes ingestion, and a runnable failure degrades to NO function (the clause
11
+ becomes NONE, exactly as an off-taxonomy span does today)."""
12
+
13
+ from __future__ import annotations
14
+
15
+ import asyncio
16
+ import re
17
+ from typing import Any, Protocol, runtime_checkable
18
+
19
+ from pydantic import BaseModel, Field
20
+
21
+ from rag_wright.contracts.function import (
22
+ FUNCTION_LABELS,
23
+ FunctionConfidence,
24
+ FunctionScore,
25
+ canonical_function,
26
+ )
27
+
28
+ _FLOOR: frozenset[FunctionConfidence] = frozenset({FunctionConfidence.HIGH, FunctionConfidence.MEDIUM})
29
+ _MAX_FUNCTIONS = 3
30
+
31
+ # ADR-0048 Phase A: the structured `function` field advertises the closed label set (the 52 taxonomy labels +
32
+ # "OTHER" for a real clause type we lack a label for). Emitted into the JSON schema so guided decoding HARD-
33
+ # constrains the model to a valid label (killing the granite-8B failure mode: inventing free-form names like
34
+ # "Exclusive Source of Supply" that then fall to NONE), and strongly guides it under function-calling. Kept as a
35
+ # `str` field (json_schema_extra is schema-only, not pydantic-enforced), so a stray value never crashes a whole
36
+ # sub-batch and the taxonomy-gap channel (`other_label`) still works. "OTHER" is not a FUNCTION_LABELS entry.
37
+ _FUNCTION_ENUM: list[str] = [*FUNCTION_LABELS, "OTHER"]
38
+
39
+
40
+ @runtime_checkable
41
+ class ClauseFunctionClassifier(Protocol):
42
+ """The ingestion seam (option B): classify a chunk's spans with the chunk as SHARED context, returning one
43
+ ranked `FunctionScore` list per span (aligned to the input order). One LLM call per chunk, not per span."""
44
+
45
+ def classify_spans(self, chunk_text: str, span_texts: list[str]) -> list[list[FunctionScore]]: ...
46
+
47
+
48
+ class RawScore(BaseModel):
49
+ """The LLM's raw (pre-validation) score: a function label (one of the 52 taxonomy labels or "OTHER") + coarse
50
+ confidence. When `function` is "OTHER" (a real clause type not in our taxonomy), `other_label` names it -- the
51
+ taxonomy-gap signal (ADR-0048 option 2). Filtered/categorized downstream. `function` advertises the closed
52
+ label enum in the JSON schema (guided decoding), but stays a `str` so a stray never crashes the sub-batch."""
53
+
54
+ function: str = Field(json_schema_extra={"enum": _FUNCTION_ENUM})
55
+ confidence: str
56
+ other_label: str = ""
57
+
58
+
59
+ class ClauseFunctionClassification(BaseModel):
60
+ """One span/clause's structured output: applicable functions RANKED primary-first, each with a coarse
61
+ `high|medium|low` confidence. Kept raw (str fields) so post-processing can drop LLM slop rather than fail the
62
+ whole call."""
63
+
64
+ functions: list[RawScore]
65
+
66
+
67
+ class SpanFunctions(BaseModel):
68
+ """One span's classification, keyed by its EXPLICIT `span_index` (the `[n]` in the prompt) -- alignment is by
69
+ index, NOT list position, so a dropped/reordered span can't silently shift every later span's label. The
70
+ classifier may OMIT spans it assigns no function."""
71
+
72
+ span_index: int
73
+ functions: list[RawScore]
74
+
75
+
76
+ class BatchSpanClassification(BaseModel):
77
+ """The batched (option B) output: per-span classifications keyed by `span_index` (sparse -- no-function spans
78
+ may be omitted)."""
79
+
80
+ spans: list[SpanFunctions]
81
+
82
+
83
+ def _to_scores(raw_scores: list[RawScore]) -> list[FunctionScore]:
84
+ """Canonicalize + confidence-floor (>= medium) + cap (<=3), preserving rank. Drops off-taxonomy / NONE /
85
+ below-floor / unparseable-confidence entries (robust to LLM slop)."""
86
+ out: list[FunctionScore] = []
87
+ for r in raw_scores:
88
+ canon = canonical_function(r.function)
89
+ if canon is None: # off-taxonomy or NONE
90
+ continue
91
+ try:
92
+ conf = FunctionConfidence(r.confidence.strip().lower())
93
+ except ValueError:
94
+ continue
95
+ if conf in _FLOOR:
96
+ out.append(FunctionScore(function=canon, confidence=conf))
97
+ if len(out) >= _MAX_FUNCTIONS:
98
+ break
99
+ return out
100
+
101
+
102
+ def categorize_raw(raw_scores: list[RawScore]) -> tuple[list[FunctionScore], list[str]]:
103
+ """ADR-0048 option 2: split a span's raw scores into (in-taxonomy FunctionScores [canonical+floored],
104
+ out-of-taxonomy labels). An entry whose `function` is off-taxonomy / "OTHER" is captured by its real clause
105
+ type (`other_label`, else the raw `function`) -- distinguishing "no function" from "a function we lack a
106
+ label for" (the taxonomy-gap signal), instead of silently collapsing both to NONE."""
107
+ in_tax = _to_scores(raw_scores)
108
+ others: list[str] = []
109
+ for r in raw_scores:
110
+ if canonical_function(r.function) is None: # off-taxonomy
111
+ # the LLM may follow the convention (function="OTHER", real type in other_label) OR put the real
112
+ # type directly in `function` (with other_label empty/"None"). Prefer other_label ONLY when function
113
+ # is the literal "OTHER"; otherwise the off-taxonomy `function` string IS the real type.
114
+ if r.function.strip().upper() == "OTHER":
115
+ label = r.other_label.strip()
116
+ else:
117
+ label = r.function.strip()
118
+ if label and label.upper() not in ("NONE", "OTHER"):
119
+ others.append(label)
120
+ return in_tax, others
121
+
122
+
123
+ _PROMPT = (
124
+ "You are classifying a single contract CLAUSE by its legal function. Read the whole clause (not a fragment) "
125
+ "and list EVERY function it genuinely serves, ranked most-relevant first. A clause may serve more than one "
126
+ "function; most serve exactly one. For each, give a coarse confidence: high, medium, or low. Use ONLY these "
127
+ "function types; if none apply, return an empty list.\n\nFunction types:\n{labels}\n\nClause:\n{text}"
128
+ )
129
+
130
+
131
+ class LlmClauseClassifier:
132
+ """Classify a full clause via the graph-building LLM (structured). `runnable` is the injected structured seam
133
+ (`.invoke(prompt) -> ClauseFunctionClassification`); production wires it through the model-profile seam."""
134
+
135
+ def __init__(self, runnable: Any) -> None:
136
+ self._runnable = runnable
137
+
138
+ def classify(self, clause_text: str) -> list[FunctionScore]:
139
+ from rag_wright.contracts.function import FUNCTION_LABELS
140
+
141
+ prompt = _PROMPT.format(labels="\n".join(FUNCTION_LABELS), text=clause_text)
142
+ try:
143
+ raw = self._runnable.invoke(prompt)
144
+ except Exception: # noqa: BLE001 - a classify failure degrades to NO function (never crash ingestion)
145
+ return []
146
+ if not isinstance(raw, ClauseFunctionClassification):
147
+ return []
148
+ return _to_scores(raw.functions)
149
+
150
+
151
+ _BATCH_CAP = 10 # max spans per LLM call -- a big multi-provision chunk is split into sub-batches (bounded prompt)
152
+
153
+ _BATCH_PROMPT = (
154
+ "You are classifying the operative provisions of ONE contract section by their legal function. Read the WHOLE "
155
+ "section for context, then classify each numbered span below by the function(s) it serves -- ranked "
156
+ "most-relevant first, each with a coarse confidence (high, medium, or low). For `function`, use EXACTLY one of "
157
+ "the function types listed below. If a span's real function is NOT in the list, set `function` to \"OTHER\" and "
158
+ "put the actual clause type in `other_label` -- do NOT invent a name in `function`. For EACH span you classify, "
159
+ "return its `span_index` (the [n] number) and its functions; OMIT any span that serves no function. Classify "
160
+ "only indices 0..{max_index}.\n\nFunction types:\n{labels}\n\nSECTION (context):\n{context}\n\nSPANS:\n{spans}"
161
+ )
162
+
163
+ _SUBBATCH_CONCURRENCY = 8 # max concurrent sub-batch LLM calls per chunk (async), bounded against provider limits
164
+
165
+
166
+ def _subs(span_texts: list[str]) -> list[tuple[int, list[str]]]:
167
+ """Split spans into sub-batches of at most `_BATCH_CAP`, each a `(start_index, spans)` pair."""
168
+ return [(start, span_texts[start:start + _BATCH_CAP]) for start in range(0, len(span_texts), _BATCH_CAP)]
169
+
170
+
171
+ def _sub_prompt(chunk_text: str, sub: list[str]) -> str:
172
+ """The batch-classify prompt for one sub-batch (shared by the sync and async paths)."""
173
+ from rag_wright.contracts.function import FUNCTION_LABELS
174
+
175
+ numbered = "\n".join(f"[{i}] {t}" for i, t in enumerate(sub))
176
+ return _BATCH_PROMPT.format(
177
+ labels="\n".join(FUNCTION_LABELS), context=chunk_text, spans=numbered, max_index=len(sub) - 1)
178
+
179
+
180
+ def _merge_subbatches(results: list[tuple[int, int, Any]], span_texts: list[str]) -> list[list[RawScore]]:
181
+ """Merge sub-batch results into index-aligned per-span scores, aligning by the returned `span_index` (relative
182
+ to each sub-batch, offset by its start) so a dropped/reordered span can't shift later labels."""
183
+ out: list[list[RawScore]] = [[] for _ in span_texts]
184
+ for start, sublen, raw in results:
185
+ if raw is None:
186
+ continue
187
+ for sf in raw.spans:
188
+ idx = start + sf.span_index
189
+ if start <= idx < start + sublen:
190
+ out[idx] = list(sf.functions)
191
+ return out
192
+
193
+
194
+ class LlmBatchClauseClassifier:
195
+ """Option B: classify a chunk's spans with the chunk as shared context. A big chunk is split into sub-batches
196
+ of `_BATCH_CAP` spans (bounded prompt), each ONE LLM call. Output is aligned by the returned `span_index` (not
197
+ list position), so a dropped/reordered span can't shift later labels. `runnable` is the injected structured
198
+ seam (`.invoke(prompt) -> BatchSpanClassification`). Robust: a failed sub-batch leaves its spans empty."""
199
+
200
+ def __init__(self, runnable: Any) -> None:
201
+ self._runnable = runnable
202
+
203
+ def _classify_raw(self, chunk_text: str, span_texts: list[str]) -> list[list[RawScore]]:
204
+ """Sub-batched LLM calls, run CONCURRENTLY within the chunk (a big multi-provision chunk's sub-batches are
205
+ independent, so a 29-span chunk costs ~1 call's latency, not 3x). Index-aligned RAW per-span scores (the
206
+ LLM's function+confidence strings, NO canonicalize/floor/cap). A failed sub-batch leaves its spans empty."""
207
+ if not span_texts:
208
+ return []
209
+ from concurrent.futures import ThreadPoolExecutor
210
+
211
+ subs = _subs(span_texts)
212
+
213
+ def _call(item):
214
+ start, sub = item
215
+ try:
216
+ raw = self._runnable.invoke(_sub_prompt(chunk_text, sub))
217
+ except Exception: # noqa: BLE001 - a failed sub-batch leaves its spans empty (never crash)
218
+ return (start, len(sub), None)
219
+ return (start, len(sub), raw if isinstance(raw, BatchSpanClassification) else None)
220
+
221
+ if len(subs) <= 1: # single sub-batch -> no thread pool
222
+ results = [_call(subs[0])]
223
+ else: # concurrent sub-batches (independent calls, merged by span_index)
224
+ with ThreadPoolExecutor(max_workers=len(subs)) as ex:
225
+ results = list(ex.map(_call, subs))
226
+ return _merge_subbatches(results, span_texts)
227
+
228
+ async def _aclassify_raw(self, chunk_text: str, span_texts: list[str],
229
+ *, sem: asyncio.Semaphore | None = None) -> list[list[RawScore]]:
230
+ """ASYNC-B1 (ADR-0057): the async twin of `_classify_raw`. Sub-batches run concurrently via
231
+ `asyncio.gather` bounded by a `Semaphore` (native form of the sync thread pool); each `.ainvoke` carries
232
+ the true wall-clock deadline. A failed sub-batch -- including a `ModelCallTimeout` (an Exception, so it is
233
+ caught here) -- leaves its spans empty: the degrade path stays reachable and the node never raises.
234
+
235
+ CLASSIFY-CONCURRENCY-1: a caller classifying MANY chunks concurrently passes ONE shared `sem`, so the total
236
+ in-flight sub-batch calls across all chunks are bounded by a single deliberate knob (else each chunk got
237
+ its own `_SUBBATCH_CONCURRENCY` budget). The sem is acquired only at the leaf `.ainvoke` -- never held
238
+ across another acquire -- so nesting the chunk gather over it cannot deadlock."""
239
+ if not span_texts:
240
+ return []
241
+ subs = _subs(span_texts)
242
+ sem = sem if sem is not None else asyncio.Semaphore(_SUBBATCH_CONCURRENCY)
243
+
244
+ async def _acall(item):
245
+ start, sub = item
246
+ async with sem:
247
+ try:
248
+ raw = await self._runnable.ainvoke(_sub_prompt(chunk_text, sub))
249
+ except Exception: # noqa: BLE001 - failed sub-batch -> empty spans (never crash the document)
250
+ return (start, len(sub), None)
251
+ return (start, len(sub), raw if isinstance(raw, BatchSpanClassification) else None)
252
+
253
+ results = await asyncio.gather(*(_acall(s) for s in subs))
254
+ return _merge_subbatches(results, span_texts)
255
+
256
+ def classify_spans(self, chunk_text: str, span_texts: list[str]) -> list[list[FunctionScore]]:
257
+ return [_to_scores(raws) for raws in self._classify_raw(chunk_text, span_texts)]
258
+
259
+ async def aclassify_spans(self, chunk_text: str, span_texts: list[str],
260
+ *, sem: asyncio.Semaphore | None = None) -> list[list[FunctionScore]]:
261
+ return [_to_scores(raws) for raws in await self._aclassify_raw(chunk_text, span_texts, sem=sem)]
262
+
263
+ def classify_spans_raw(self, chunk_text: str, span_texts: list[str]) -> list[list[RawScore]]:
264
+ """Debug: the RAW per-span scores (pre-floor), to see what the confidence floor drops."""
265
+ return self._classify_raw(chunk_text, span_texts)
266
+
267
+
268
+ class LegalBertClauseAdapter:
269
+ """Back-compat: wrap the span-level single-label LegalBERT classifier. Its one label becomes the PRIMARY at
270
+ `high` confidence; NONE / off-taxonomy -> [] (no function). It is span-level, so `classify_spans` ignores the
271
+ chunk context (each span classified independently, exactly as today)."""
272
+
273
+ def __init__(self, classifier: Any) -> None:
274
+ self._classifier = classifier
275
+
276
+ def classify(self, clause_text: str) -> list[FunctionScore]:
277
+ return self.classify_spans("", [clause_text])[0]
278
+
279
+ def classify_spans(self, chunk_text: str, span_texts: list[str]) -> list[list[FunctionScore]]: # noqa: ARG002
280
+ out: list[list[FunctionScore]] = []
281
+ for label in self._classifier.classify(span_texts):
282
+ canon = canonical_function(label)
283
+ out.append([FunctionScore(function=canon, confidence=FunctionConfidence.HIGH)] if canon else [])
284
+ return out
285
+
286
+
287
+ # --- issue 0005 / route (b): CLIENT-SIDE free-text tag classification (no server-side guided decoding) ---------
288
+ #
289
+ # The classifier schemas are nested list[BaseModel] (BatchSpanClassification.spans -> SpanFunctions.functions ->
290
+ # RawScore), so `build_tag_structured` (FLAT only) can't serve them. Server-side guided decoding
291
+ # (`build_structured`) runs away on self-hosted Granite (profiles.py: no `client_side_structured` escape) -- the
292
+ # ADR-0058 boundary-call failure at a new call site. So, like `answer_generator.parse_tagged_answer`, the
293
+ # classifier uses its OWN compact free-text tag format that FLATTENS the nesting into the tag BODY, parsed
294
+ # CLIENT-SIDE into the same Pydantic contract. `function` stays raw -- the downstream categorize/floor/cap cleans
295
+ # a stray, so a malformed token never crashes a sub-batch.
296
+
297
+ _BATCH_TAG_INSTRUCTIONS = (
298
+ "\n\nReturn ONLY the classification as tagged lines -- ONE per span you assign a function to (OMIT any span "
299
+ "with no function). Between the tags put the applicable function labels PRIMARY-FIRST as `Label:confidence` "
300
+ "(confidence = high, medium, or low), comma-separated. For a real clause type NOT in the list above, use "
301
+ "`OTHER:confidence:short_name`. Use ONLY the [n] span indices shown. Example:\n"
302
+ '<span index="0">Cap On Liability:high, Indemnification:medium</span>\n'
303
+ '<span index="3">Governing Law:high</span>'
304
+ )
305
+ _CLAUSE_TAG_INSTRUCTIONS = (
306
+ "\n\nReturn ONLY the applicable function labels PRIMARY-FIRST as `Label:confidence` (confidence = high, "
307
+ "medium, or low), comma-separated, between <functions></functions> tags; for a real clause type NOT in the "
308
+ "list above, use `OTHER:confidence:short_name`. If none apply, return `<functions></functions>`. Example:\n"
309
+ "<functions>Cap On Liability:high, Indemnification:low</functions>"
310
+ )
311
+
312
+ _SPAN_TAG_RE = re.compile(r'<span\s+index="?(\d+)"?\s*>(.*?)</span>', re.DOTALL | re.IGNORECASE)
313
+ _FUNCTIONS_TAG_RE = re.compile(r"<functions>(.*?)</functions>", re.DOTALL | re.IGNORECASE)
314
+
315
+
316
+ def _parse_score_item(item: str) -> RawScore | None:
317
+ """One `Label:confidence` (or `OTHER:confidence:other_label`) token -> a RawScore, kept RAW."""
318
+ parts = [p.strip() for p in item.split(":")]
319
+ if not parts or not parts[0]:
320
+ return None
321
+ if parts[0].upper() == "OTHER":
322
+ return RawScore(function="OTHER", confidence=parts[1] if len(parts) > 1 else "low",
323
+ other_label=parts[2] if len(parts) > 2 else "")
324
+ if len(parts) == 1:
325
+ return RawScore(function=parts[0], confidence="low") # no confidence given -> lowest (dropped by floor)
326
+ return RawScore(function=":".join(parts[:-1]), confidence=parts[-1]) # rejoin a label that itself had a colon
327
+
328
+
329
+ def _parse_scores(body: str) -> list[RawScore]:
330
+ return [rs for item in re.split(r"[,\n]", body) if (rs := _parse_score_item(item.strip())) is not None]
331
+
332
+
333
+ def parse_batch_span_tags(text: str) -> BatchSpanClassification:
334
+ """CLIENT-SIDE parse of the batched free-text tags -> BatchSpanClassification. Sparse (a no-function span is
335
+ absent) and aligned by the EXPLICIT index in each tag, matching the schema's span_index contract."""
336
+ spans = [SpanFunctions(span_index=int(m.group(1)), functions=scores)
337
+ for m in _SPAN_TAG_RE.finditer(text) if (scores := _parse_scores(m.group(2)))]
338
+ return BatchSpanClassification(spans=spans)
339
+
340
+
341
+ def parse_clause_function_tags(text: str) -> ClauseFunctionClassification:
342
+ """CLIENT-SIDE parse of the single-clause free-text tags -> ClauseFunctionClassification."""
343
+ m = _FUNCTIONS_TAG_RE.search(text)
344
+ return ClauseFunctionClassification(functions=_parse_scores(m.group(1)) if m else [])
345
+
346
+
347
+ class _TagClassifierRunnable:
348
+ """A `build_structured`-shaped runnable (`.invoke`/`.ainvoke(prompt) -> the Pydantic contract`) that drives the
349
+ classifier CLIENT-SIDE: a plain free-text call (NO server guided decoding) whose tagged output is parsed by
350
+ `parse`. `.ainvoke` streams via `astream_text` (true wall-clock deadline + the stage `label`, issue 0005)."""
351
+
352
+ def __init__(self, model_id: str, *, instructions: str, parse: Any, label: str, max_tokens: int = 2048) -> None:
353
+ self._model_id = model_id
354
+ self._instructions = instructions
355
+ self._parse = parse
356
+ self._label = label
357
+ self._max_tokens = max_tokens
358
+
359
+ def invoke(self, prompt: Any, config: Any = None) -> Any: # config accepted for runnable-compat, unused
360
+ from rag_wright.models.seam import build_model
361
+
362
+ text = build_model(self._model_id, max_tokens=self._max_tokens).invoke(prompt + self._instructions).content
363
+ return self._parse(str(text))
364
+
365
+ async def ainvoke(self, prompt: Any, config: Any = None) -> Any:
366
+ from rag_wright.models.seam import astream_text
367
+
368
+ text = await astream_text(self._model_id, prompt + self._instructions,
369
+ max_tokens=self._max_tokens, label=self._label)
370
+ return self._parse(text)
371
+
372
+
373
+ def production_llm_clause_classifier(model_id: str) -> LlmClauseClassifier:
374
+ """Wire the single-clause classifier over CLIENT-SIDE free-text tag parse (issue 0005: no server-side guided
375
+ decoding, which runs away on self-hosted Granite). Same `ClauseFunctionClassification` contract."""
376
+ return LlmClauseClassifier(_TagClassifierRunnable(
377
+ model_id, instructions=_CLAUSE_TAG_INSTRUCTIONS, parse=parse_clause_function_tags,
378
+ label="clause_function_classifier.classify"))
379
+
380
+
381
+ def production_batch_clause_classifier(model_id: str) -> LlmBatchClauseClassifier:
382
+ """Wire the BATCHED (option B) classifier -- the ingestion default -- over CLIENT-SIDE free-text tag parse
383
+ (issue 0005). Same `BatchSpanClassification` contract; parsed by `parse_batch_span_tags`."""
384
+ return LlmBatchClauseClassifier(_TagClassifierRunnable(
385
+ model_id, instructions=_BATCH_TAG_INSTRUCTIONS, parse=parse_batch_span_tags,
386
+ label="clause_function_classifier.classify_spans"))
387
+
388
+
389
+ # --- T55 / SETFIT-SEG-1: trained SetFit soft-tagger (replaces the LLM classifier for ingestion latency) ----------
390
+ # Behind the SAME ClauseFunctionClassifier seam as LlmBatchClauseClassifier / LegalBertClauseAdapter -- a new
391
+ # IMPLEMENTATION of the already-registered `clause_function_classification` capability, no contract/API change.
392
+ # Function is a SOFT tag (ADR-0047), so emitting multiple tags per span is intended.
393
+ class SetFitClauseAdapter:
394
+ """In-process ENSEMBLE of trained SetFit models. Each model on disk is a Sentence-Transformer body
395
+ (`model.safetensors` + configs) + a joblib-pickled sklearn head (`model_head.pkl`), so inference needs only
396
+ sentence-transformers + scikit-learn + joblib -- NO `setfit` dependency. `classify_spans` encodes every span
397
+ with each body, AVERAGES the per-class probabilities across the ensemble, and emits the top-k labels above
398
+ `threshold` as `FunctionScore` soft tags (primary-first). Span-level (ignores `chunk_text`, like
399
+ `LegalBertClauseAdapter`). Load-don't-retrain from local checkpoints; milliseconds/span, no network hop."""
400
+
401
+ def __init__(self, model_dirs, *, top_k: int = 3, threshold: float = 0.0,
402
+ hi: float = 0.6, mid: float = 0.3, device: str | None = None, batch_size: int = 32) -> None:
403
+ # DEFAULT = the finalized operating point: pure avg-prob TOP-3 (threshold 0.0) -> ~3.0 tags/span, the
404
+ # validated 50/52 classes >0.65 recall. A higher threshold trims tags but drops recall (e.g. 0.10 -> ~2.0
405
+ # tags, ~47/52); tune via RAG_SETFIT_THRESHOLD / RAG_SETFIT_TOPK only with a re-measured operating point.
406
+ import json
407
+ from pathlib import Path
408
+
409
+ import joblib
410
+ import numpy as np
411
+ from sentence_transformers import SentenceTransformer
412
+
413
+ self._np = np
414
+ self._top_k, self._threshold, self._hi, self._mid, self._batch = top_k, threshold, hi, mid, batch_size
415
+ self._models: list[tuple] = []
416
+ for d in model_dirs:
417
+ d = Path(d)
418
+ cfg = json.loads((d / "config_setfit.json").read_text())
419
+ body = SentenceTransformer(str(d), device=device)
420
+ head = joblib.load(d / "model_head.pkl")
421
+ self._models.append((body, head, bool(cfg.get("normalize_embeddings", False)),
422
+ [str(c) for c in head.classes_]))
423
+ if not self._models:
424
+ raise ValueError("SetFitClauseAdapter needs at least one model dir")
425
+ self._labels = sorted({lab for *_, cols in self._models for lab in cols}) # union label space (robust)
426
+ self._lab_idx = {lab: i for i, lab in enumerate(self._labels)}
427
+
428
+ def classify(self, clause_text: str) -> list[FunctionScore]:
429
+ return self.classify_spans("", [clause_text])[0]
430
+
431
+ async def aclassify_spans(self, chunk_text: str, span_texts: list[str],
432
+ *, sem: asyncio.Semaphore | None = None) -> list[list[FunctionScore]]:
433
+ # the ingestion pipeline calls aclassify_spans under a shared concurrency bound. SetFit is CPU-bound and
434
+ # in-process, so run the sync encode+predict off the event loop in a thread; the sem bounds in-flight work.
435
+ if sem is None:
436
+ return await asyncio.to_thread(self.classify_spans, chunk_text, span_texts)
437
+ async with sem:
438
+ return await asyncio.to_thread(self.classify_spans, chunk_text, span_texts)
439
+
440
+ def classify_spans(self, chunk_text: str, span_texts: list[str]) -> list[list[FunctionScore]]: # noqa: ARG002
441
+ if not span_texts:
442
+ return []
443
+ np = self._np
444
+ agg = np.zeros((len(span_texts), len(self._labels)))
445
+ for body, head, norm, cols in self._models:
446
+ proba = np.asarray(head.predict_proba(
447
+ body.encode(list(span_texts), normalize_embeddings=norm, batch_size=self._batch)), dtype=float)
448
+ for j, c in enumerate(cols):
449
+ idx = self._lab_idx.get(c)
450
+ if idx is not None:
451
+ agg[:, idx] += proba[:, j]
452
+ agg /= len(self._models)
453
+ out: list[list[FunctionScore]] = []
454
+ for row in agg:
455
+ scores: list[FunctionScore] = []
456
+ for j in np.argsort(-row)[: self._top_k]:
457
+ p = float(row[j])
458
+ if p < self._threshold:
459
+ break
460
+ canon = canonical_function(self._labels[int(j)])
461
+ if not canon:
462
+ continue
463
+ conf = (FunctionConfidence.HIGH if p >= self._hi
464
+ else FunctionConfidence.MEDIUM if p >= self._mid else FunctionConfidence.LOW)
465
+ scores.append(FunctionScore(function=canon, confidence=conf))
466
+ out.append(scores)
467
+ return out
468
+
469
+
470
+ def production_setfit_clause_classifier(model_root: str | None = None, **kwargs) -> SetFitClauseAdapter:
471
+ """Wire the ENSEMBLE SetFit soft-tagger -- the ingestion default (replaces the LLM classifier for latency,
472
+ T55/SETFIT-SEG-1). Loads the 3 finalized checkpoints (LegalBERT + BGE-large + MPNet) from `model_root`
473
+ (env RAG_SETFIT_CLAUSE_DIR; default data/models/setfit_clause). In-process, no network hop, no `setfit` dep.
474
+ Tunables: RAG_SETFIT_TOPK, RAG_SETFIT_THRESHOLD, RAG_SETFIT_DEVICE."""
475
+ import os
476
+ from pathlib import Path
477
+
478
+ root = Path(model_root or os.getenv("RAG_SETFIT_CLAUSE_DIR", "data/models/setfit_clause"))
479
+ subdirs = [root / n for n in ("cap128b_legalbert", "cap128b_bge", "cap128b_mpnet")]
480
+ present = [d for d in subdirs if (d / "model_head.pkl").exists()]
481
+ if not present:
482
+ raise FileNotFoundError(
483
+ f"No SetFit clause checkpoints under {root} (expected cap128b_legalbert/bge/mpnet). "
484
+ "Download them from the model store, or set RAG_FUNCTION_CLASSIFIER=llm to use the LLM classifier.")
485
+ return SetFitClauseAdapter(
486
+ present,
487
+ top_k=int(os.getenv("RAG_SETFIT_TOPK", str(kwargs.pop("top_k", 3)))),
488
+ threshold=float(os.getenv("RAG_SETFIT_THRESHOLD", str(kwargs.pop("threshold", 0.0)))),
489
+ device=os.getenv("RAG_SETFIT_DEVICE") or kwargs.pop("device", None),
490
+ **kwargs)