neuralos 3.0.2__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.
needle/__init__.py ADDED
@@ -0,0 +1,645 @@
1
+ from __future__ import annotations
2
+
3
+ import ctypes
4
+ import decimal
5
+ import datetime
6
+ import json
7
+ import os
8
+ import re
9
+ import warnings
10
+
11
+ from .agent.tools import Field, build_schema, pydantic_schema, tool, _is_pydantic_model
12
+ from ._telemetry import track as _track
13
+ from ._worker import FineTuneWorker
14
+
15
+ __version__ = "3.0.2"
16
+ __all__ = ["Needle", "ExtractionValidationError", "tool", "Field", "extract",
17
+ "__version__"]
18
+
19
+
20
+ class ExtractionValidationError(ValueError):
21
+ """The engine produced structured values that are not grounded in the input."""
22
+
23
+
24
+ UNRESET_TURNS = 4
25
+
26
+ _CACT_GENERATIONS = {
27
+ 0x05E12A83: 2,
28
+ 0x05E12A84: 3,
29
+ }
30
+
31
+
32
+ def _weight_generation(path):
33
+ with open(path, "rb") as handle:
34
+ tag_bytes = handle.read(4)
35
+ if len(tag_bytes) != 4:
36
+ raise RuntimeError(f"{path} is not a complete .cact archive")
37
+ tag = int.from_bytes(tag_bytes, "little")
38
+ try:
39
+ return _CACT_GENERATIONS[tag]
40
+ except KeyError as exc:
41
+ raise RuntimeError(
42
+ f"{path} has unknown .cact format tag 0x{tag:08x}; "
43
+ "cannot choose a compatible Needle engine") from exc
44
+
45
+
46
+ _CACT_HEADER = "<48If"
47
+ _CACT_RECORD = "<BBHIIIIQQII"
48
+ _CACT_FP16, _CACT_RAW = 1, 4
49
+ _CONFIDENCE_HEAD_CODE = 2
50
+
51
+
52
+ def _confidence_head_present(path):
53
+ """Whether a Needle 3 archive carries a confidence head, read from its head manifest."""
54
+ import struct
55
+
56
+ with open(path, "rb") as handle:
57
+ header = handle.read(struct.calcsize(_CACT_HEADER))
58
+ if len(header) < struct.calcsize(_CACT_HEADER):
59
+ return False
60
+ fields = struct.unpack(_CACT_HEADER, header)
61
+ num_tensors, codebook = fields[1], fields[2]
62
+ handle.seek(codebook * 4, 1)
63
+ size = struct.calcsize(_CACT_RECORD)
64
+ records = [struct.unpack(_CACT_RECORD, handle.read(size)) for _ in range(num_tensors)]
65
+ if not records or records[-1][0] != _CACT_RAW:
66
+ return False
67
+ for heads in (1, 2, 3):
68
+ index = num_tensors - 2 - 6 * heads
69
+ if index < 0:
70
+ continue
71
+ dtype, ndim, _, length = records[index][:4]
72
+ if dtype != _CACT_FP16 or ndim != 1 or length != heads:
73
+ continue
74
+ handle.seek(records[index][7])
75
+ codes = struct.unpack(f"<{heads}e", handle.read(2 * heads))
76
+ return _CONFIDENCE_HEAD_CODE in {int(round(c)) for c in codes}
77
+ return False
78
+
79
+
80
+ def _library_path(generation=2):
81
+ from .agent import fetch
82
+
83
+ generation = int(generation)
84
+ override = os.environ.get(f"NEEDLE{generation}_LIB_PATH")
85
+ if generation == 2 and not override:
86
+ # NEEDLE_LIB_PATH predates multi-generation dispatch and therefore
87
+ # names the Needle 2 engine. Never route a v3 archive through it.
88
+ override = os.environ.get("NEEDLE_LIB_PATH")
89
+ if override:
90
+ return override
91
+ here = os.path.dirname(os.path.abspath(__file__))
92
+ lib_name = fetch._lib_name()
93
+ stem, suffix = os.path.splitext(lib_name)
94
+ local_names = [f"{stem}{generation}{suffix}"]
95
+ if generation == 2:
96
+ # Wheels published before the split shipped Needle 2 as libneedle.*.
97
+ local_names.append(lib_name)
98
+ for name in local_names:
99
+ local = os.path.join(here, name)
100
+ if os.path.exists(local):
101
+ return local
102
+ version = fetch.engine_version(generation)
103
+ cache = os.path.join(os.path.expanduser("~"), ".cache", "cactus-needle",
104
+ f"v{generation}", version)
105
+ cached = os.path.join(cache, fetch._lib_name())
106
+ if os.path.exists(cached):
107
+ return cached
108
+ os.makedirs(cache, exist_ok=True)
109
+ return fetch.fetch_library(version, cache, generation=generation)
110
+
111
+
112
+ def _base_weights_path(generation):
113
+ from .agent import fetch
114
+
115
+ name = fetch.base_weights(generation)
116
+ # Engine-bundled wheels ship the weights next to the package code.
117
+ bundled = os.path.join(os.path.dirname(os.path.abspath(__file__)), name)
118
+ if os.path.exists(bundled):
119
+ return bundled
120
+ local = os.path.join(fetch.cache_dir(generation), name)
121
+ if os.path.exists(local):
122
+ return local
123
+ return fetch.fetch_weights(generation)
124
+
125
+
126
+ _lib_handles = {}
127
+ _active = {}
128
+ _loaded_base = {}
129
+
130
+
131
+ def _load_base(lib, generation):
132
+ if generation < 3 or _loaded_base.get(generation):
133
+ return
134
+ path = _base_weights_path(generation)
135
+ with open(path, "rb") as handle:
136
+ data = handle.read()
137
+ if lib.needle_load(data, len(data)) < 0:
138
+ raise RuntimeError(f"needle_load failed for {path}")
139
+ _loaded_base[generation] = path
140
+ _active.pop(generation, None)
141
+
142
+
143
+ def _load_cdll(generation):
144
+ """Load the engine, retrying with the other libc build if the first is wrong.
145
+
146
+ A musl system that reports itself as glibc (or the reverse) otherwise fails
147
+ at load with a missing symbol such as strtoll_l.
148
+ """
149
+ path = _library_path(generation)
150
+ try:
151
+ return ctypes.CDLL(path)
152
+ except OSError as first:
153
+ from .agent import fetch
154
+
155
+ other = fetch.other_libc_tag()
156
+ if other is None:
157
+ raise
158
+ try:
159
+ alt = fetch.fetch_library(dest_dir=os.path.dirname(path) or ".",
160
+ tag=other, generation=generation)
161
+ handle = ctypes.CDLL(alt)
162
+ except Exception:
163
+ raise first from None
164
+ warnings.warn(f"the {fetch._platform_tag()} engine did not load ({first}); "
165
+ f"using the {other} build instead", stacklevel=3)
166
+ return handle
167
+
168
+
169
+ def _lib(generation=2):
170
+ generation = int(generation)
171
+ if generation not in _lib_handles:
172
+ lib = _load_cdll(generation)
173
+ lib.needle_init.argtypes = [ctypes.c_char_p, ctypes.c_char_p, ctypes.c_char_p]
174
+ lib.needle_init.restype = ctypes.c_int
175
+ lib.needle_complete.argtypes = [
176
+ ctypes.c_char_p, ctypes.c_int, ctypes.c_char_p, ctypes.c_int]
177
+ lib.needle_complete.restype = ctypes.c_int
178
+ if generation >= 3:
179
+ lib.needle_embed.argtypes = [
180
+ ctypes.c_char_p, ctypes.POINTER(ctypes.c_float), ctypes.c_int]
181
+ lib.needle_embed.restype = ctypes.c_int
182
+ lib.needle_reset.argtypes = []
183
+ lib.needle_reset.restype = None
184
+ lib.needle_load.argtypes = [ctypes.c_char_p, ctypes.c_uint64]
185
+ lib.needle_load.restype = ctypes.c_int
186
+ _lib_handles[generation] = lib
187
+ return _lib_handles[generation]
188
+
189
+
190
+ class Needle:
191
+ def __init__(self, tools=None, system=None, weights=None, tool_index_path=None, buffer_size=65536,
192
+ auto_date=True, generation=None, stateless=False):
193
+ self._functions = {}
194
+ self._stateless = bool(stateless)
195
+ self._turns = 0
196
+ self._weights = os.fspath(weights) if weights is not None else None
197
+ self._tuned = self._weights is not None
198
+ if self._tuned:
199
+ self._generation = _weight_generation(self._weights)
200
+ else:
201
+ self._generation = int(generation or 3)
202
+ self._worker = None
203
+ self._closed = False
204
+ self._calibrated = (not self._tuned
205
+ or (self._generation >= 3 and _confidence_head_present(self._weights)))
206
+ if not self._calibrated:
207
+ warnings.warn("these weights carry no confidence head trained for them, so this "
208
+ "agent reports confidence as None; platform fine-tunes keep the head",
209
+ stacklevel=2)
210
+ self._system_text = _with_date_fact(system or "") if auto_date else (system or "")
211
+ self._system = self._system_text.encode("utf-8")
212
+ tools_json = tools if isinstance(tools, str) else json.dumps(self._resolve(tools))
213
+ self._tools_json = tools_json.encode("utf-8")
214
+ try:
215
+ parsed_tools = json.loads(tools_json)
216
+ self._n_tools = len(parsed_tools)
217
+ except (json.JSONDecodeError, TypeError):
218
+ parsed_tools, self._n_tools = [], None
219
+ self._tool_schemas = [entry for entry in (parsed_tools or [])
220
+ if isinstance(entry, dict)]
221
+ self._seen_years = set()
222
+ self._tool_index_path = (os.fspath(tool_index_path).encode("utf-8")
223
+ if tool_index_path else None)
224
+ self._buffer = ctypes.create_string_buffer(buffer_size)
225
+ if self._weights:
226
+ self._worker = FineTuneWorker(
227
+ _library_path(self._generation), self._weights,
228
+ self._system.decode("utf-8"), self._tools_json.decode("utf-8"),
229
+ os.fspath(tool_index_path) if tool_index_path else None,
230
+ buffer_size, generation=self._generation)
231
+ else:
232
+ self._bind()
233
+
234
+ def _bind(self):
235
+ if self._closed:
236
+ raise RuntimeError("Needle instance is closed")
237
+ if self._worker is not None:
238
+ return
239
+ generation = self._generation
240
+ lib = _lib(generation)
241
+ _load_base(lib, generation)
242
+ if _active.get(generation) is self:
243
+ return
244
+ if lib.needle_init(self._system, self._tools_json, self._tool_index_path) < 0:
245
+ _active.pop(generation, None)
246
+ raise RuntimeError("needle_init failed")
247
+ _active[generation] = self
248
+ self._seen_years = set()
249
+
250
+ def _resolve(self, tools):
251
+ schemas = []
252
+ for entry in tools or []:
253
+ if _is_pydantic_model(entry):
254
+ schema = pydantic_schema(entry)
255
+ self._functions[schema["name"]] = entry
256
+ schemas.append(schema)
257
+ elif callable(entry):
258
+ schema = getattr(entry, "_needle_tool", None) or build_schema(entry)
259
+ self._functions[schema["name"]] = entry
260
+ schemas.append(schema)
261
+ elif isinstance(entry, dict):
262
+ schemas.append(entry)
263
+ return schemas
264
+
265
+ def _track_props(self):
266
+ return {"n_tools": self._n_tools, "tuned": self._tuned,
267
+ "generation": self._generation}
268
+
269
+ def complete(self, text: str = "", max_new_tokens: int = 512) -> dict:
270
+ _track("complete", self._track_props())
271
+ if self._stateless:
272
+ self.reset()
273
+ self._count_query()
274
+ return self._complete(text, max_new_tokens)
275
+
276
+ def _count_query(self):
277
+ self._turns += 1
278
+ if self._turns == UNRESET_TURNS + 1:
279
+ warnings.warn(
280
+ f"{UNRESET_TURNS} queries on this agent without reset(): every turn stays in "
281
+ "the conversation, so unrelated queries lose accuracy and confidence. "
282
+ "Call reset() between independent queries, or construct with stateless=True",
283
+ stacklevel=3)
284
+
285
+ def _complete(self, text: str, max_new_tokens: int = 512,
286
+ ground: bool = True) -> dict:
287
+ self._bind()
288
+ self._seen_years |= _source_years(text or "")
289
+ if self._worker is not None:
290
+ raw = self._worker.complete(text, max_new_tokens)
291
+ else:
292
+ lib = _lib(self._generation)
293
+ rc = lib.needle_complete(
294
+ text.encode("utf-8"), int(max_new_tokens), self._buffer,
295
+ len(self._buffer))
296
+ if rc < 0:
297
+ detail = self._buffer.value.decode("utf-8", "replace")
298
+ raise RuntimeError(detail or f"needle_complete failed (code {rc})")
299
+ raw = self._buffer.value.decode("utf-8")
300
+ try:
301
+ response = json.loads(raw)
302
+ except json.JSONDecodeError as err:
303
+ raise RuntimeError(
304
+ f"engine returned an unparseable envelope ({err}); this is an "
305
+ f"engine bug - please report it with the prompt and schema") from err
306
+ if not self._calibrated:
307
+ response["confidence"] = None
308
+ if ground:
309
+ _annotate_ungrounded(response, self._tool_schemas, self._seen_years,
310
+ self._system_text, text)
311
+ return response
312
+
313
+ def embed(self, text: str = "") -> list[float]:
314
+ if self._generation < 3:
315
+ raise ValueError("embeddings require a Needle 3 model")
316
+ self._bind()
317
+ if self._worker is not None:
318
+ return self._worker.embed(text)
319
+ lib = _lib(self._generation)
320
+ dim = lib.needle_embed(text.encode("utf-8"), None, 0)
321
+ if dim <= 0:
322
+ raise RuntimeError(f"needle_embed failed (code {dim})")
323
+ output = (ctypes.c_float * dim)()
324
+ rc = lib.needle_embed(text.encode("utf-8"), output, dim)
325
+ if rc != dim:
326
+ raise RuntimeError(f"needle_embed failed (code {rc})")
327
+ return list(output)
328
+
329
+ def run(self, query: str = "", max_steps: int = 8,
330
+ max_new_tokens: int = 512, strict: bool = True) -> dict:
331
+ _track("run", self._track_props())
332
+ if self._stateless:
333
+ self.reset()
334
+ self._count_query()
335
+ response = self._complete(query, max_new_tokens)
336
+ executed = []
337
+ for _ in range(max_steps):
338
+ calls = response.get("function_calls") or []
339
+ if response.get("type") != "call" or not calls:
340
+ break
341
+ ungrounded = _ungrounded_paths(response)
342
+ results = []
343
+ for call in calls:
344
+ name = str(call.get("name"))
345
+ fabricated = sorted(ungrounded.get(name, ()))
346
+ if strict and fabricated:
347
+ grounded = _grounded_number_paths(
348
+ call.get("arguments") or {}, query, self._system_text)
349
+ fabricated = [path for path in fabricated
350
+ if path not in grounded]
351
+ if strict and fabricated:
352
+ results.append({"error": "ungrounded " + ", ".join(fabricated)})
353
+ continue
354
+ fn = self._functions.get(call.get("name"))
355
+ if fn is None:
356
+ results.append({"error": "unknown tool: " + name})
357
+ continue
358
+ try:
359
+ results.append(fn(**(call.get("arguments") or {})))
360
+ except Exception as exc:
361
+ results.append({"error": str(exc)})
362
+ executed.extend(results)
363
+ response = self._complete(json.dumps(results, default=_jsonable),
364
+ max_new_tokens, ground=False)
365
+ response["results"] = executed
366
+ return response
367
+
368
+ def extract(self, text: str, schema: type | dict, max_new_tokens: int = 512,
369
+ strict: bool = True) -> object:
370
+ return extract(text, schema, max_new_tokens=max_new_tokens,
371
+ weights=self._weights, strict=strict)
372
+
373
+ def reset(self):
374
+ self._bind()
375
+ self._turns = 0
376
+ if self._worker is not None:
377
+ self._worker.reset()
378
+ else:
379
+ _lib(self._generation).needle_reset()
380
+ self._seen_years = set()
381
+
382
+ def close(self):
383
+ if self._worker is not None:
384
+ try:
385
+ self._worker.close()
386
+ finally:
387
+ self._worker = None
388
+ self._closed = True
389
+
390
+ def __enter__(self):
391
+ return self
392
+
393
+ def __exit__(self, exc_type, exc, traceback):
394
+ self.close()
395
+
396
+ def __del__(self):
397
+ try:
398
+ self.close()
399
+ except Exception:
400
+ pass
401
+
402
+
403
+ def _jsonable(value):
404
+ if hasattr(value, "model_dump"):
405
+ return value.model_dump()
406
+ if hasattr(value, "dict") and _is_pydantic_model(type(value)):
407
+ return value.dict()
408
+ return str(value)
409
+
410
+
411
+ def _schema_parameters(schema):
412
+ raw = pydantic_schema(schema) if _is_pydantic_model(schema) else schema
413
+ return raw.get("parameters", raw) if isinstance(raw, dict) else {}
414
+
415
+
416
+ def _resolve_ref(node, root):
417
+ seen = set()
418
+ while isinstance(node, dict) and "$ref" in node:
419
+ if node["$ref"] in seen:
420
+ break
421
+ seen.add(node["$ref"])
422
+ target = root
423
+ for part in node["$ref"].removeprefix("#/").split("/"):
424
+ target = target.get(part.replace("~1", "/").replace("~0", "~"), {})
425
+ if target is node:
426
+ break
427
+ node = target
428
+ return node
429
+
430
+
431
+ def _source_years(text):
432
+ months = (r"(?:jan(?:uary)?|feb(?:ruary)?|mar(?:ch)?|apr(?:il)?|may|"
433
+ r"jun(?:e)?|jul(?:y)?|aug(?:ust)?|sep(?:t(?:ember)?)?|"
434
+ r"oct(?:ober)?|nov(?:ember)?|dec(?:ember)?)")
435
+ patterns = [
436
+ rf"\b\d{{1,2}}(?:st|nd|rd|th)?\s+{months}[\s,]+(\d{{1,4}})(?![0-9A-Za-z])",
437
+ rf"\b{months}\s+\d{{1,2}}(?:st|nd|rd|th)?\s*,?\s+(\d{{1,4}})(?![0-9A-Za-z])",
438
+ rf"\b{months}[\s,]+(\d{{3,4}})(?![0-9A-Za-z])",
439
+ r"\byear\s+(\d{1,4})(?![0-9A-Za-z])",
440
+ # Year-first numeric dates only (2024-03-15, 2024/03/15). A day- or
441
+ # month-first date such as 5/6/24 must not mint its leading component
442
+ # as a year: no ISO argument can ever match it, so every date would fail.
443
+ r"(?<![0-9])(\d{4})(?=[-/]\d{1,2}[-/]\d{1,2}(?![0-9]))",
444
+ ]
445
+ lowered = text.lower()
446
+ return {int(match.group(1)) for pattern in patterns
447
+ for match in re.finditer(pattern, lowered)}
448
+
449
+
450
+ _ISO_STAMP = re.compile(r"\d{4}-\d{2}-\d{2}")
451
+
452
+
453
+ def _with_date_fact(system: str) -> str:
454
+ """Prefix the local date fact unless the caller already supplied one."""
455
+ if "date:" in system or _ISO_STAMP.search(system):
456
+ return system
457
+ now = datetime.datetime.now()
458
+ fact = now.strftime("date: %Y-%m-%d %a %H:%M")
459
+ if not system.strip():
460
+ return fact
461
+ if system.lstrip().startswith("{"):
462
+ return system
463
+ return fact + "; " + system
464
+
465
+
466
+ _RELATIVE_CUE = re.compile(
467
+ r"\b(today|tonight|tomorrow|yesterday|next|this|coming|now|in \d+ (?:days?|weeks?|months?|years?)|"
468
+ r"monday|tuesday|wednesday|thursday|friday|saturday|sunday)\b", re.IGNORECASE)
469
+
470
+
471
+ def _relative_cue(text) -> bool:
472
+ return bool(text) and _RELATIVE_CUE.search(text) is not None
473
+
474
+
475
+ def _licensed_years(seen_years, system=None, relative=True):
476
+ """Years a date argument may carry: the ones written in the input, plus the
477
+ system date's year when the input reasons relatively (tomorrow, next week)."""
478
+ years = set(seen_years)
479
+ if years and system and relative:
480
+ years |= _source_years(system)
481
+ return years
482
+
483
+
484
+ def _walk_grounding(schema, arguments, years):
485
+ root = _schema_parameters(schema)
486
+ checked, failures = set(), set()
487
+
488
+ def walk(value, node, path):
489
+ node = _resolve_ref(node, root)
490
+ variants = node.get("anyOf") or node.get("oneOf") or []
491
+ concrete = [v for v in variants if _resolve_ref(v, root).get("type") != "null"]
492
+ if len(concrete) == 1:
493
+ node = _resolve_ref(concrete[0], root)
494
+ fmt = node.get("format")
495
+ if fmt in ("date", "date-time") and isinstance(value, str):
496
+ match = re.match(r"^(\d{4})-", value)
497
+ if match and years:
498
+ checked.add(path)
499
+ if int(match.group(1)) not in years:
500
+ failures.add(path)
501
+ return
502
+ if isinstance(value, dict):
503
+ properties = node.get("properties", {})
504
+ for key, item in value.items():
505
+ if key in properties:
506
+ walk(item, properties[key], f"{path}.{key}" if path else key)
507
+ elif isinstance(value, list) and "items" in node:
508
+ for index, item in enumerate(value):
509
+ walk(item, node["items"], f"{path}[{index}]")
510
+
511
+ walk(arguments, root, "")
512
+ return checked, failures
513
+
514
+
515
+ def _temporal_grounding(text, schema, arguments, system=None):
516
+ years = _licensed_years(_source_years(text), system, _relative_cue(text))
517
+ return _walk_grounding(schema, arguments, years)
518
+
519
+
520
+ _NUMBER_TOKEN = re.compile(
521
+ r"(?<![\w.,])(?:[-+]?\d{1,3}(?:,\d{3})+(?:\.\d+)?|[-+]?\d+(?:\.\d+)?)(?!\d)")
522
+
523
+
524
+ def _numeric_paths(arguments, path=""):
525
+ """Yield ``(path, value)`` for every numeric leaf, arguments-relative.
526
+
527
+ Paths match the form both call sites derive after stripping the engine's
528
+ tool prefix, so a flagged ``Tool.path`` resolves to a helper path.
529
+ """
530
+ if isinstance(arguments, bool):
531
+ return
532
+ if isinstance(arguments, (int, float)):
533
+ yield path, arguments
534
+ elif isinstance(arguments, dict):
535
+ for key, item in arguments.items():
536
+ yield from _numeric_paths(item, f"{path}.{key}" if path else str(key))
537
+ elif isinstance(arguments, list):
538
+ for index, item in enumerate(arguments):
539
+ yield from _numeric_paths(item, f"{path}[{index}]")
540
+
541
+
542
+ _DATE_FACT = re.compile(r"date:\s*\d{4}-\d{2}-\d{2}(?:\s+[A-Za-z]{3})?(?:\s+\d{2}:\d{2})?|\d{4}-\d{2}-\d{2}T\d{2}:\d{2}(?::\d{2})?")
543
+
544
+
545
+ def _without_date_facts(text):
546
+ return _DATE_FACT.sub(" ", text or "")
547
+
548
+
549
+ def _source_numbers(*sources):
550
+ sources = tuple(_without_date_facts(source) for source in sources)
551
+ """Numeric values written in the source texts, separators normalized."""
552
+ text = "\n".join(source for source in sources if source)
553
+ return {decimal.Decimal(match.group(0).replace(",", ""))
554
+ for match in _NUMBER_TOKEN.finditer(text)}
555
+
556
+
557
+ def _grounded_number_paths(arguments, *sources):
558
+ numbers = _source_numbers(*sources)
559
+ if not numbers:
560
+ return set()
561
+ grounded = set()
562
+ for path, value in _numeric_paths(arguments):
563
+ try:
564
+ if decimal.Decimal(str(value)) in numbers:
565
+ grounded.add(path)
566
+ except (decimal.InvalidOperation, ValueError):
567
+ continue
568
+ return grounded
569
+
570
+
571
+ def _ungrounded_paths(response):
572
+ validation = response.get("validation") or {}
573
+ grouped = {}
574
+ for name in validation.get("ungrounded") or []:
575
+ tool, _, path = str(name).partition(".")
576
+ grouped.setdefault(tool, set()).add(path or tool)
577
+ return grouped
578
+
579
+
580
+ def _annotate_ungrounded(response, tool_schemas, seen_years, system=None, text=None):
581
+ calls = response.get("function_calls") or []
582
+ years = _licensed_years(seen_years, system, _relative_cue(text) if text is not None else True)
583
+ if not calls or not years:
584
+ return
585
+ schemas = {entry.get("name"): entry for entry in tool_schemas}
586
+ found = []
587
+ for call in calls:
588
+ schema = schemas.get(call.get("name"))
589
+ if schema is None:
590
+ continue
591
+ _, failures = _walk_grounding(schema, call.get("arguments") or {}, years)
592
+ found.extend(f"{call['name']}.{path}" for path in sorted(failures))
593
+ if not found:
594
+ return
595
+ validation = response.setdefault("validation", {}) or {}
596
+ response["validation"] = validation
597
+ existing = list(validation.get("ungrounded") or [])
598
+ validation["ungrounded"] = existing + [
599
+ name for name in found if name not in existing]
600
+
601
+
602
+ def _validate_extraction(text, schema, arguments, response, system=None):
603
+ checked, failures = _temporal_grounding(text, schema, arguments, system)
604
+ validation = response.get("validation") or {}
605
+ flagged = validation.get("ungrounded") or []
606
+ grounded = _grounded_number_paths(arguments, text, system) if flagged else set()
607
+ for name in flagged:
608
+ path = name.split(".", 1)[-1]
609
+ if path in grounded:
610
+ continue
611
+ if path not in checked or path in failures:
612
+ failures.add(path)
613
+ if validation.get("negation"):
614
+ failures.add("negated request")
615
+ if failures:
616
+ detail = ", ".join(sorted(failures))
617
+ raise ExtractionValidationError(
618
+ f"extraction returned values not grounded in the input: {detail}")
619
+
620
+
621
+ def extract(text: str, schema: type | dict, system: str | None = None,
622
+ max_new_tokens: int = 512, weights: str | None = None,
623
+ strict: bool = True, generation: int | None = None) -> object:
624
+ """One-shot structured extraction using the matching native engine.
625
+
626
+ With ``strict=True`` (the default), temporal values that contradict a literal
627
+ year in the input, plus engine-reported fabricated or negated values, raise
628
+ :class:`ExtractionValidationError` instead of being returned silently.
629
+ """
630
+ selected = weights
631
+ generation = _weight_generation(selected) if selected else int(generation or 3)
632
+ _track("extract", {"n_tools": 1, "tuned": bool(selected),
633
+ "generation": generation})
634
+ agent = Needle(tools=[schema], system=system, weights=selected, generation=generation)
635
+ try:
636
+ response = agent._complete(text, max_new_tokens)
637
+ finally:
638
+ agent.close()
639
+ calls = response.get("function_calls") or response.get("suppressed_calls") or []
640
+ if not calls:
641
+ return None
642
+ arguments = calls[0].get("arguments") or {}
643
+ if strict:
644
+ _validate_extraction(text, schema, arguments, response, system)
645
+ return schema(**arguments) if _is_pydantic_model(schema) else arguments