Skip to content

Commit 3a60e44

Browse files
fix: address adaptive context review findings
1 parent 6f43e8e commit 3a60e44

4 files changed

Lines changed: 96 additions & 11 deletions

File tree

‎engraphis/service.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2035,6 +2035,7 @@ def adaptive_context(
20352035
"title": chunk.get("title"),
20362036
"scope": chunk.get("scope"),
20372037
"mtype": chunk.get("mtype"),
2038+
"provenance": _compact_provenance(chunk.get("provenance")),
20382039
}
20392040
for packed in result.recall.packed_chunks
20402041
if (chunk := chunks_by_id.get(packed.id)) is not None

‎eval/productivity.py‎

Lines changed: 48 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,6 @@
3333
from engraphis.core.interfaces import MemoryType, Scope
3434
from engraphis.core.store import Store
3535
from engraphis.core.textutil import tokenize
36-
from eval import metrics
3736
from eval.harness import _seed_case_graph, load_dataset
3837

3938

@@ -169,10 +168,35 @@ def _percentile(values: list[float], percentile: float) -> float:
169168
return ordered[index]
170169

171170

172-
def _completed(response: str, expected: str) -> bool:
173-
if not str(expected or "").strip():
174-
return bool(str(response or "").strip())
175-
return metrics.answer_token_recall([str(response or "")], str(expected)) >= 1.0
171+
AnswerEvaluator = Callable[[str, dict, tuple[str, ...]], bool]
172+
173+
174+
def _normalized_answer(value: object) -> str:
175+
"""Return a punctuation-insensitive canonical answer for fixture comparison."""
176+
return " ".join(re.findall(r"[\w-]+", str(value or "").casefold()))
177+
178+
179+
def _completed(response: str, question: dict, supporting_evidence: tuple[str, ...]) -> bool:
180+
"""Evaluate task success against a case's explicit answer and source evidence.
181+
182+
Productivity completion is a correctness metric, not a retrieval metric: token
183+
containment lets statements such as ``the release manager does not approve``
184+
count as a successful answer to ``release manager``. The offline oracle accepts
185+
only a case's canonical answer, an explicitly listed acceptable answer, or an
186+
exact supporting evidence sentence. Hosted or paraphrasing benchmarks can
187+
inject an ``answer_evaluator`` into :func:`run` with richer semantics.
188+
"""
189+
normalized_response = _normalized_answer(response)
190+
expected = str(question.get("answer", question.get("evidence", "")))
191+
if not _normalized_answer(expected):
192+
return bool(normalized_response)
193+
acceptable = [expected, *supporting_evidence]
194+
configured = question.get("acceptable_answers", ())
195+
if isinstance(configured, (list, tuple)):
196+
acceptable.extend(str(value) for value in configured)
197+
return normalized_response in {
198+
candidate for value in acceptable if (candidate := _normalized_answer(value))
199+
}
176200

177201

178202
def _seed_case(
@@ -326,6 +350,7 @@ def run(
326350
dim: int = 256,
327351
embedder: Optional[object] = None,
328352
agent: Optional[Callable[[str, str], Union[str, AgentTurn]]] = None,
353+
answer_evaluator: Optional[AnswerEvaluator] = None,
329354
clock: Callable[[], float] = time.perf_counter,
330355
strategy_order: tuple[str, ...] = STRATEGIES,
331356
) -> dict:
@@ -370,10 +395,15 @@ def run(
370395
counter = RegexTokenCounter()
371396
selected_embedder = embedder or DeterministicEmbedder(dim=dim)
372397
selected_agent = agent or DeterministicTaskAgent()
398+
selected_answer_evaluator = answer_evaluator or _completed
373399
rows = {name: [] for name in STRATEGIES}
374400
task_offset = 0
375401

376402
for case in dataset:
403+
evidence_by_tag = {
404+
str(memory.get("tag")): str(memory.get("text", ""))
405+
for memory in case.get("memories", [])
406+
}
377407
# Each strategy gets an independently seeded engine. This keeps recall
378408
# caches, reinforcement bugs, or future mutable read state from making a
379409
# later strategy look artificially faster or more accurate.
@@ -386,8 +416,10 @@ def run(
386416
try:
387417
for number, question_row in enumerate(case.get("questions", [])):
388418
question = str(question_row.get("q", ""))
389-
expected = str(
390-
question_row.get("answer", question_row.get("evidence", ""))
419+
supporting_evidence = tuple(
420+
evidence_by_tag[str(tag)]
421+
for tag in question_row.get("supporting", [])
422+
if str(tag) in evidence_by_tag
391423
)
392424
task_id = str(
393425
question_row.get("id") or f"{case.get('id', 'case')}:{number}"
@@ -415,7 +447,9 @@ def run(
415447
)
416448
first_turn = _turn(selected_agent(question, context))
417449
first_response = first_turn.answer
418-
first_completed = _completed(first_response, expected)
450+
first_completed = selected_answer_evaluator(
451+
first_response, question_row, supporting_evidence
452+
)
419453
first_abstained = not first_response.strip()
420454
agent_turns = 1
421455
input_tokens = counter(question) + counter(context)
@@ -440,12 +474,16 @@ def run(
440474
selected_agent(corrected_question, correction_history)
441475
)
442476
corrected_response = corrected_turn.answer
443-
successful_correction = _completed(corrected_response, expected)
477+
successful_correction = selected_answer_evaluator(
478+
corrected_response, question_row, supporting_evidence
479+
)
444480
final_response = corrected_response
445481
agent_turns += 1
446482
input_tokens += counter(corrected_question) + counter(correction_history)
447483
output_tokens += counter(corrected_response)
448-
completed = _completed(final_response, expected)
484+
completed = selected_answer_evaluator(
485+
final_response, question_row, supporting_evidence
486+
)
449487
elapsed_ms = max(0.0, (clock() - started) * 1000.0)
450488
provider_turns = [first_turn]
451489
if correction_attempted:

‎tests/test_adaptive_context.py‎

Lines changed: 16 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -252,7 +252,17 @@ def test_service_adaptive_context_keeps_sources_in_packed_citation_order(monkeyp
252252
recall = RecallResult(
253253
chunks=[
254254
{"id": "mem_first", "title": "First", "scope": "repo", "mtype": "episodic"},
255-
{"id": "mem_second", "title": "Second", "scope": "repo", "mtype": "semantic"},
255+
{
256+
"id": "mem_second",
257+
"title": "Second",
258+
"scope": "repo",
259+
"mtype": "semantic",
260+
"provenance": {
261+
"source": "agent:review",
262+
"trusted": True,
263+
"secret": "must not be forwarded",
264+
},
265+
},
256266
],
257267
packed_chunks=[
258268
PackedChunk("mem_second", "second evidence", 2),
@@ -286,6 +296,11 @@ def test_service_adaptive_context_keeps_sources_in_packed_citation_order(monkeyp
286296
assert [source["id"] for source in result["sources"]] == [
287297
"mem_second", "mem_first",
288298
]
299+
assert result["sources"][0]["provenance"] == {
300+
"source": "agent:review",
301+
"trusted": True,
302+
}
303+
assert result["sources"][1]["provenance"] == {}
289304

290305

291306
def test_service_bounds_adaptive_prompt_budgets() -> None:

‎tests/test_productivity_eval.py‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,37 @@ def test_productivity_report_measures_outcomes_corrections_turns_and_all_tokens(
7373
assert adaptive["context_modes"] == {"history_bypass": 1}
7474

7575

76+
def test_productivity_completion_oracle_rejects_a_negated_answer() -> None:
77+
class NegatingAgent:
78+
def __call__(self, question, context):
79+
del question, context
80+
return "The release manager does not own deployment approval."
81+
82+
report = run(_small_dataset(), agent=NegatingAgent())
83+
84+
for method in report["methods"].values():
85+
assert method["completion_rate"] == 0.0
86+
assert method["wrong_answers"] == 1
87+
assert method["corrections"] == 1
88+
assert method["successful_corrections"] == 0
89+
90+
91+
def test_productivity_accepts_an_injected_case_aware_answer_evaluator() -> None:
92+
def evaluator(response, question, supporting_evidence):
93+
return response == question["answer"].upper() and not supporting_evidence
94+
95+
data = _small_dataset()
96+
data[0]["questions"][0].pop("supporting")
97+
98+
report = run(
99+
data,
100+
agent=lambda _question, _context: "RELEASE MANAGER",
101+
answer_evaluator=evaluator,
102+
)
103+
104+
assert all(method["completion_rate"] == 1.0 for method in report["methods"].values())
105+
106+
76107
def test_large_history_routes_between_strong_retrieval_and_weak_widening() -> None:
77108
memories = [
78109
{

0 commit comments

Comments
 (0)