diff --git a/needle/__init__.py b/needle/__init__.py index 67c1f9b7..a8f2e5ff 100644 --- a/needle/__init__.py +++ b/needle/__init__.py @@ -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() diff --git a/tests/test_weights.py b/tests/test_weights.py index 437a8f12..68d86a37 100644 --- a/tests/test_weights.py +++ b/tests/test_weights.py @@ -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}