Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions needle/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -205,8 +205,9 @@ def run(self, query: str, max_steps: int = 8, max_new_tokens: int = 256) -> dict

def extract(self, text: str, schema: type | dict, max_new_tokens: int = 256,
strict: bool = True) -> object:
return extract(text, schema, max_new_tokens=max_new_tokens,
weights=self._weights, strict=strict)
return extract(text, schema, system=self._system.decode("utf-8") or None,
max_new_tokens=max_new_tokens, weights=self._weights,
strict=strict)

def reset(self):
self._bind()
Expand Down
46 changes: 46 additions & 0 deletions tests/test_weights.py
Original file line number Diff line number Diff line change
Expand Up @@ -202,3 +202,49 @@ class Invoice(pydantic.BaseModel):
assert needle._source_years("due September 5") == set()
assert needle._source_years("due 5th September 42") == {42}
assert needle._source_years("Invoice 42 is due tomorrow at 5") == set()


def test_agent_extract_carries_its_own_system_facts(engine, monkeypatch):
import needle

seen = {}

def spy(text, schema, system=None, max_new_tokens=256, weights=None, strict=True):
seen.update(text=text, system=system, weights=weights, strict=strict)
return None

monkeypatch.setattr(needle, "extract", spy)
facts = "date: 2026-07-21 Tue 14:30; locale: en-US"
agent = needle.Needle(tools="[]", system=facts)

assert agent.extract("dinner tomorrow at 7", {"type": "object"}) is None
assert seen["system"] == facts
assert seen["text"] == "dinner tomorrow at 7"


def test_agent_extract_without_system_facts_sends_none(engine, monkeypatch):
import needle

seen = {}
monkeypatch.setattr(
needle, "extract",
lambda text, schema, system=None, **kwargs: seen.update(system=system))
needle.Needle(tools="[]").extract("anything", {"type": "object"})

assert seen["system"] is None


def test_agent_extract_still_carries_weights_and_strict(engine, tuned, monkeypatch):
import needle

seen = {}
monkeypatch.setattr(
needle, "extract",
lambda text, schema, system=None, max_new_tokens=256, weights=None, strict=True:
seen.update(system=system, weights=weights, strict=strict))
with warnings.catch_warnings():
warnings.simplefilter("ignore")
agent = needle.Needle(tools="[]", weights=tuned, system="device: phone")
agent.extract("anything", {"type": "object"}, strict=False)

assert seen == {"system": "device: phone", "weights": tuned, "strict": False}