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 +645 -0
- needle/_telemetry.py +86 -0
- needle/_worker.py +285 -0
- needle/agent/__init__.py +0 -0
- needle/agent/fetch.py +225 -0
- needle/agent/tools.py +180 -0
- needle/cli.py +352 -0
- needle/environments/__init__.py +45 -0
- needle/environments/_harness.py +60 -0
- needle/environments/data_capture.py +132 -0
- needle/environments/kitchen_appliance.py +121 -0
- needle/environments/media_player.py +104 -0
- needle/environments/productivity.py +133 -0
- needle/environments/smart_home.py +139 -0
- needle/environments/wearable.py +111 -0
- needle/model/__init__.py +0 -0
- needle/model/architecture.py +1136 -0
- needle/model/checkpoints.py +121 -0
- needle/model/export.py +599 -0
- needle/model/finetune.py +577 -0
- needle/model/quantize.py +235 -0
- needle/model/run.py +125 -0
- needle/model/tokenizer.py +130 -0
- needle/platform.py +429 -0
- needle/playground/__init__.py +0 -0
- needle/playground/app.js +331 -0
- needle/playground/index.html +111 -0
- needle/playground/server.py +185 -0
- needle/playground/style.css +513 -0
- neuralos-3.0.2.dist-info/METADATA +168 -0
- neuralos-3.0.2.dist-info/RECORD +35 -0
- neuralos-3.0.2.dist-info/WHEEL +5 -0
- neuralos-3.0.2.dist-info/entry_points.txt +3 -0
- neuralos-3.0.2.dist-info/licenses/LICENSE +202 -0
- neuralos-3.0.2.dist-info/top_level.txt +1 -0
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
|