botp
/

File size: 3,393 Bytes
1d2de8a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
import math
import threading

import pytest

from solomon_mlx._vendor.contract import parse_questions
from solomon_mlx._vendor.semantics import p_yes
from solomon_mlx.api import Solomon, branches, distributions, ordering_score


def test_collapse_before_temperature():
    p = p_yes([0, 0, 0, 0], 2)
    assert p == pytest.approx(1 / (1 + math.sqrt(3)))
    assert p != pytest.approx(0.25)


def test_routing_and_candidate_order():
    specs = parse_questions(
        {
            "single": {"type": "choice", "instructions": "Who?", "options": ["A", "B"]},
            "ordered": {"type": "score", "instructions": "Level?", "levels": ["low", "high"]},
            "entity": {"instructions": "Is {candidate} certified?", "candidates": ["Z", "X"]},
        }
    )
    assert branches(specs[0])[0][1:] == (4, "single/choiceR")
    assert branches(specs[1])[0][1:] == (2, "ordered/choiceS")
    assert "Is Z certified?" in branches(specs[2])[0][0]
    assert distributions(specs[0], [{"letter_logits": [0, 0, 100, 100]}], 1) == [[0.5, 0.5]]


def test_ordering_score_product():
    assert ordering_score([[0.8, 0.2], [0.1, 0.9]]) == pytest.approx(0.72)
    with pytest.raises(ValueError):
        ordering_score([[0.8, 0.3]])


class FakeEngine:
    def __init__(self):
        self.identity = {"fingerprint": "test-only"}
        self.lock = threading.RLock()
        self.prefills = []

    def prefill(self, parts):
        self.prefills.append(parts)
        return {"parts": parts, "prefix_ids": [1, 2]}

    def ask(self, state, block, width, head, **kwargs):
        return {
            "letter_logits": [2.0] + [0.0] * (width - 1),
            "branch_tokens": 3,
            "prompt_tokens": 5,
            "head_key": head,
        }


def test_state_ownership_close_and_replay(tmp_path):
    model = Solomon(FakeEngine())
    other = Solomon(FakeEngine())
    state = model.prefill("A fact.")
    recipe = tmp_path / "state.json"
    state.save(recipe)
    with pytest.raises(ValueError):
        other.decide(state=state, questions={"a": "Fact?"})
    state.close()
    with pytest.raises(ValueError):
        model.decide(state=state, questions={"a": "Fact?"})
    with model.replay(recipe) as restored:
        assert restored.prefix_tokens == 2
    recipe.write_text(recipe.read_text().replace("A fact.", "Bad fact."))
    with pytest.raises(ValueError):
        model.replay(recipe)


def test_evidence_budget_stops_fresh_calls():
    engine = FakeEngine()
    model = Solomon(engine)
    with model.prefill("Alice is certified.\nBob is not certified.") as state:
        result = model.decide(
            state=state, questions={"a": "Is Alice certified?"}, evidence="removal", evidence_max_calls=0
        )
    assert result["answers"]["a"]["evidence_status"] == "budget_exhausted"
    assert len(engine.prefills) == 1


def test_evidence_spans_and_fresh_verification():
    engine = FakeEngine()
    model = Solomon(engine)
    text = "Alice is certified.\nBob is not certified."
    with model.prefill(text) as state:
        out = model.decide(state=state, questions={"a": "Is Alice certified?"}, evidence="removal")[
            "answers"
        ]["a"]
    assert len(engine.prefills) == 3
    for span in out["evidence"]:
        assert text[span["start"] : span["end"]] == span["text"]
    assert out["evidence_detail"]["verification"] == "fresh_source_reencoding"