583 lines
17 KiB
Python
583 lines
17 KiB
Python
"""Tests for local_decision.parse_logprobs and local_decision.classify_choice.
|
|
|
|
Covers (parse_logprobs):
|
|
- Basic option-letter extraction from Ollama logprobs.
|
|
- Token stripping: trailing dots, leading whitespace.
|
|
- Non-option tokens are silently ignored.
|
|
- Confidence = winner_mass / total_option_mass.
|
|
- Coverage = total_option_mass (sum of all option letter masses).
|
|
- RuntimeError on coverage below minimum.
|
|
- RuntimeError on zero total mass.
|
|
- RuntimeError on missing logprobs.
|
|
- Multiple positions and overlapping top_logprobs entries.
|
|
|
|
Covers (classify_choice):
|
|
- POSTs to the Ollama native /api/chat endpoint with the right body.
|
|
- timeout_s maps to the requests timeout.
|
|
- Returns (label, confidence, coverage) from parse_logprobs.
|
|
- Propagates requests.exceptions.RequestException.
|
|
- Forwards coverage_min to parse_logprobs.
|
|
"""
|
|
|
|
import json
|
|
import math
|
|
from typing import Any
|
|
|
|
import pytest
|
|
|
|
from local_decision import parse_logprobs
|
|
|
|
# --- Fixtures ----------------------------------------------------------------
|
|
|
|
|
|
def _make_logprob(
|
|
token: str,
|
|
logprob: float,
|
|
top_logprobs: list[dict[str, Any]],
|
|
) -> dict[str, Any]:
|
|
"""Convenience builder for a logprob position entry."""
|
|
return {
|
|
"token": token,
|
|
"logprob": logprob,
|
|
"top_logprobs": top_logprobs,
|
|
}
|
|
|
|
|
|
def _make_response(
|
|
positions: list[dict[str, Any]],
|
|
) -> dict[str, Any]:
|
|
"""Build a minimal Ollama-style response dict from position entries."""
|
|
return {"logprobs": positions}
|
|
|
|
|
|
# --- Basic parsing -----------------------------------------------------------
|
|
|
|
|
|
def test_simple_two_option_a_wins():
|
|
"""A has much higher mass than B -> (A, ~1.0, total)."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A",
|
|
-0.01,
|
|
[
|
|
{"token": "A", "logprob": -0.01},
|
|
{"token": "B", "logprob": -3.0},
|
|
],
|
|
),
|
|
_make_logprob(
|
|
".",
|
|
-0.5,
|
|
[{"token": ".", "logprob": -0.5}],
|
|
),
|
|
])
|
|
label, confidence, coverage = parse_logprobs(response, ["A", "B"])
|
|
assert label == "A"
|
|
assert confidence > 0.95
|
|
assert coverage > 0
|
|
|
|
|
|
def test_simple_two_option_b_wins():
|
|
"""B has much higher mass than A -> (B, ~1.0, total)."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"B",
|
|
-0.02,
|
|
[
|
|
{"token": "A", "logprob": -4.0},
|
|
{"token": "B", "logprob": -0.02},
|
|
],
|
|
),
|
|
])
|
|
label, confidence, _ = parse_logprobs(response, ["A", "B"])
|
|
assert label == "B"
|
|
assert confidence > 0.98
|
|
|
|
|
|
def test_four_options_a_wins():
|
|
"""A wins among A, B, C, D."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A",
|
|
-0.05,
|
|
[
|
|
{"token": "A", "logprob": -0.05},
|
|
{"token": "B", "logprob": -2.0},
|
|
{"token": "C", "logprob": -3.0},
|
|
{"token": "D", "logprob": -4.0},
|
|
],
|
|
),
|
|
])
|
|
label, confidence, _ = parse_logprobs(response, ["A", "B", "C", "D"])
|
|
assert label == "A"
|
|
assert confidence > 0.82
|
|
|
|
|
|
# --- Token stripping ---------------------------------------------------------
|
|
|
|
|
|
def test_trailing_dot_stripped():
|
|
"""Ollama sends 'A.' as the option separator token — must match 'A'."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A.",
|
|
-0.01,
|
|
[
|
|
{"token": "A.", "logprob": -0.01},
|
|
],
|
|
),
|
|
])
|
|
label, _, _ = parse_logprobs(response, ["A"])
|
|
assert label == "A"
|
|
|
|
|
|
def test_leading_whitespace_token_ignored():
|
|
"""Token ' thinking' — after stripping whitespace it is not an option."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
" thinking",
|
|
-0.1,
|
|
[
|
|
{"token": " thinking", "logprob": -0.1},
|
|
{"token": " response", "logprob": -0.2},
|
|
],
|
|
),
|
|
])
|
|
# Should not raise; no option letters found so coverage < 0.3.
|
|
with pytest.raises(RuntimeError, match="coverage"):
|
|
parse_logprobs(response, ["A", "B"])
|
|
|
|
|
|
def test_thinking_token_does_not_pollute_mass():
|
|
"""'thinking' and 'response' tokens are ignored, not counted."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A",
|
|
-0.01,
|
|
[
|
|
{"token": "A", "logprob": -0.01},
|
|
{"token": "thinking", "logprob": -0.5},
|
|
{"token": "response", "logprob": -0.6},
|
|
],
|
|
),
|
|
])
|
|
label, confidence, _ = parse_logprobs(response, ["A", "B"])
|
|
assert label == "A"
|
|
# Only A's mass contributes — confidence should be 1.0
|
|
assert confidence == pytest.approx(1.0)
|
|
|
|
|
|
# --- Confidence and coverage -------------------------------------------------
|
|
|
|
|
|
def test_confidence_is_normalized_mass_ratio():
|
|
"""confidence = winner_mass / total_option_mass."""
|
|
# Use -ln(2) ≈ -0.693 so exp ≈ 0.5 for A, and -ln(1.5) ≈ -0.405 for B
|
|
logprob_a = math.log(2.0) # ≈ 0.693, exp = 0.5
|
|
logprob_b = math.log(1.5) # ≈ 0.405, exp = 0.667
|
|
# A's mass = 0.5, B's mass = 0.667 → total = 1.167, confidence(A) = 0.5/1.167
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"B",
|
|
-logprob_b,
|
|
[
|
|
{"token": "A", "logprob": -logprob_a},
|
|
{"token": "B", "logprob": -logprob_b},
|
|
],
|
|
),
|
|
])
|
|
label, confidence, coverage = parse_logprobs(response, ["A", "B"])
|
|
assert label == "B"
|
|
expected_confidence = math.exp(-logprob_b) / (math.exp(-logprob_a) + math.exp(-logprob_b))
|
|
assert confidence == pytest.approx(expected_confidence, rel=1e-9)
|
|
assert coverage == pytest.approx(math.exp(-logprob_a) + math.exp(-logprob_b), rel=1e-9)
|
|
|
|
|
|
def test_coverage_is_total_mass():
|
|
"""Coverage = sum of exp(logprob) across all option letters."""
|
|
logprob_val = -0.1
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A",
|
|
logprob_val,
|
|
[
|
|
{"token": "A", "logprob": logprob_val},
|
|
{"token": "B", "logprob": logprob_val},
|
|
],
|
|
),
|
|
])
|
|
_, _, coverage = parse_logprobs(response, ["A", "B"])
|
|
expected = math.exp(logprob_val) + math.exp(logprob_val)
|
|
assert coverage == pytest.approx(expected, rel=1e-9)
|
|
|
|
|
|
# --- Multiple positions ------------------------------------------------------
|
|
|
|
|
|
def test_multiple_positions_accumulate_mass():
|
|
"""Logprobs from multiple positions add up per option."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A",
|
|
-0.1,
|
|
[{"token": "A", "logprob": -0.1}, {"token": "B", "logprob": -0.2}],
|
|
),
|
|
_make_logprob(
|
|
"B",
|
|
-0.05,
|
|
[{"token": "A", "logprob": -0.3}, {"token": "B", "logprob": -0.05}],
|
|
),
|
|
])
|
|
label, _, coverage = parse_logprobs(response, ["A", "B"])
|
|
assert label == "B"
|
|
# A mass = exp(-0.1) + exp(-0.3), B mass = exp(-0.2) + exp(-0.05)
|
|
expected_a = math.exp(-0.1) + math.exp(-0.3)
|
|
expected_b = math.exp(-0.2) + math.exp(-0.05)
|
|
assert coverage == pytest.approx(expected_a + expected_b, rel=1e-9)
|
|
|
|
|
|
def test_mixed_positions_with_non_option_tokens():
|
|
"""Mix of option letters and noise tokens across multiple positions."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A",
|
|
-0.01,
|
|
[
|
|
{"token": "A", "logprob": -0.01},
|
|
{"token": "thinking", "logprob": -1.0},
|
|
],
|
|
),
|
|
_make_logprob(
|
|
" response",
|
|
-0.3,
|
|
[
|
|
{"token": "B", "logprob": -0.3},
|
|
{"token": "assistant", "logprob": -2.0},
|
|
],
|
|
),
|
|
])
|
|
label, _, _ = parse_logprobs(response, ["A", "B"])
|
|
assert label == "A" # A: exp(-0.01) ≈ 0.99 > B: exp(-0.3) ≈ 0.74
|
|
|
|
|
|
# --- Coverage threshold errors -----------------------------------------------
|
|
|
|
|
|
def test_runtime_error_on_low_coverage():
|
|
"""Coverage below minimum raises RuntimeError."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A",
|
|
-10.0, # very low probability -> very small mass
|
|
[{"token": "A", "logprob": -10.0}],
|
|
),
|
|
])
|
|
with pytest.raises(RuntimeError, match="coverage"):
|
|
parse_logprobs(response, ["A", "B"], coverage_min=0.3)
|
|
|
|
|
|
def test_runtime_error_on_zero_coverage():
|
|
"""No option letters in logprobs → zero coverage → raise."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"thinking",
|
|
-0.1,
|
|
[
|
|
{"token": "thinking", "logprob": -0.1},
|
|
{"token": "response", "logprob": -0.2},
|
|
],
|
|
),
|
|
])
|
|
with pytest.raises(RuntimeError, match="coverage"):
|
|
parse_logprobs(response, ["A", "B"])
|
|
|
|
|
|
def test_runtime_error_on_empty_logprobs():
|
|
"""Empty logprobs list raises RuntimeError about missing logprobs."""
|
|
response: dict[str, Any] = {"logprobs": []}
|
|
with pytest.raises(RuntimeError, match="no logprobs"):
|
|
parse_logprobs(response, ["A", "B"])
|
|
|
|
|
|
def test_runtime_error_on_missing_logprobs_key():
|
|
"""Missing logprobs key → default to [] → raise."""
|
|
response: dict[str, Any] = {}
|
|
with pytest.raises(RuntimeError, match="no logprobs"):
|
|
parse_logprobs(response, ["A", "B"])
|
|
|
|
|
|
def test_custom_coverage_min():
|
|
"""Higher coverage_min requires stronger evidence."""
|
|
logprob_val = -0.1
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A",
|
|
logprob_val,
|
|
[{"token": "A", "logprob": logprob_val}],
|
|
),
|
|
])
|
|
# Default min (0.3) should pass since exp(-0.1) ≈ 0.905
|
|
_, _, _ = parse_logprobs(response, ["A"])
|
|
# Very high min should fail
|
|
with pytest.raises(RuntimeError, match="coverage"):
|
|
parse_logprobs(response, ["A"], coverage_min=10.0)
|
|
|
|
|
|
# --- Edge cases --------------------------------------------------------------
|
|
|
|
|
|
def test_single_option_letter():
|
|
"""Single option letter in the list — must still work."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A",
|
|
-0.01,
|
|
[{"token": "A", "logprob": -0.01}],
|
|
),
|
|
])
|
|
label, confidence, _ = parse_logprobs(response, ["A"])
|
|
assert label == "A"
|
|
assert confidence == pytest.approx(1.0)
|
|
|
|
|
|
def test_option_letters_not_present_in_response():
|
|
"""When none of the option letters appear in top_logprobs, coverage is 0."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"X",
|
|
-0.5,
|
|
[{"token": "X", "logprob": -0.5}],
|
|
),
|
|
])
|
|
with pytest.raises(RuntimeError, match="coverage"):
|
|
parse_logprobs(response, ["A", "B"])
|
|
|
|
|
|
def test_dot_token_only_is_ignored():
|
|
"""A bare '.' token (not a stripped option) should be ignored."""
|
|
response = _make_response([
|
|
_make_logprob(
|
|
".",
|
|
-0.5,
|
|
[{"token": ".", "logprob": -0.5}],
|
|
),
|
|
])
|
|
with pytest.raises(RuntimeError, match="coverage"):
|
|
parse_logprobs(response, ["A", "B"])
|
|
|
|
|
|
def test_empty_top_logprobs_list():
|
|
"""Position with an empty top_logprobs list is skipped gracefully."""
|
|
response = _make_response([
|
|
_make_logprob("A", -0.1, []),
|
|
_make_logprob(
|
|
"A",
|
|
-0.01,
|
|
[{"token": "A", "logprob": -0.01}],
|
|
),
|
|
])
|
|
label, _, _ = parse_logprobs(response, ["A"])
|
|
assert label == "A"
|
|
|
|
|
|
def test_missing_top_logprobs_key():
|
|
"""Position entry missing top_logprobs key does not crash."""
|
|
response = _make_response([
|
|
{"token": "A", "logprob": -0.1}, # no top_logprobs key
|
|
_make_logprob(
|
|
"B",
|
|
-0.01,
|
|
[{"token": "B", "logprob": -0.01}],
|
|
),
|
|
])
|
|
label, _, _ = parse_logprobs(response, ["A", "B"])
|
|
assert label == "B"
|
|
|
|
|
|
def test_missing_logprob_entry_is_skipped():
|
|
response = _make_response([
|
|
_make_logprob(
|
|
"A",
|
|
-0.01,
|
|
[
|
|
{"token": "A", "logprob": -0.01},
|
|
{"token": "B"}, # no logprob key
|
|
],
|
|
),
|
|
])
|
|
label, confidence, _ = parse_logprobs(response, ["A", "B"])
|
|
assert label == "A"
|
|
# Only A's mass contributes — B's missing logprob is skipped
|
|
assert confidence == pytest.approx(1.0)
|
|
|
|
|
|
def test_fixture_real_response_parses():
|
|
import pathlib
|
|
|
|
fixture_path = pathlib.Path(__file__).parent / "fixtures" / "real_ollama_response.json"
|
|
with open(fixture_path) as f:
|
|
response = json.load(f)
|
|
|
|
label, confidence, coverage = parse_logprobs(response, ["A", "B", "C", "D"])
|
|
assert label == "A"
|
|
assert confidence > 0.9
|
|
assert coverage > 0
|
|
|
|
|
|
# --- classify_choice --------------------------------------------------------
|
|
|
|
|
|
class _FakeResponse:
|
|
def __init__(self, payload: dict[str, Any]) -> None:
|
|
self._payload = payload
|
|
|
|
def json(self) -> dict[str, Any]:
|
|
return self._payload
|
|
|
|
def raise_for_status(self) -> None:
|
|
return None
|
|
|
|
|
|
def test_classify_choice_posts_native_api_and_parses(monkeypatch):
|
|
"""classify_choice POSTs to /api/chat and returns parse_logprobs result."""
|
|
import local_decision
|
|
|
|
captured: dict[str, Any] = {}
|
|
|
|
def fake_post(url, *, json, timeout):
|
|
captured["url"] = url
|
|
captured["json"] = json
|
|
captured["timeout"] = timeout
|
|
return _FakeResponse(_make_response([
|
|
_make_logprob(
|
|
"A",
|
|
-0.01,
|
|
[
|
|
{"token": "A", "logprob": -0.01},
|
|
{"token": "B", "logprob": -3.0},
|
|
],
|
|
),
|
|
]))
|
|
|
|
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
|
|
|
label, confidence, coverage = local_decision.classify_choice(
|
|
"Write a function",
|
|
{"A": "writing new code", "B": "refactoring"},
|
|
base_url="http://ollama:11434",
|
|
model="qwen3.5:4b",
|
|
num_ctx=8192,
|
|
timeout_s=120,
|
|
)
|
|
|
|
assert captured["url"] == "http://ollama:11434/api/chat"
|
|
assert captured["timeout"] == 120
|
|
|
|
payload = captured["json"]
|
|
assert payload["model"] == "qwen3.5:4b"
|
|
assert payload["stream"] is False
|
|
assert payload["think"] is False
|
|
assert payload["logprobs"] is True
|
|
assert payload["top_logprobs"] == 20
|
|
assert payload["options"] == {
|
|
"num_predict": 1,
|
|
"temperature": 0,
|
|
"num_ctx": 8192,
|
|
}
|
|
assert payload["messages"][0] == {"role": "system", "content": "answer with the letter only"}
|
|
user_content = payload["messages"][1]["content"]
|
|
assert "Write a function" in user_content
|
|
assert "A. writing new code" in user_content
|
|
assert "B. refactoring" in user_content
|
|
assert user_content.endswith("Answer with the letter only.")
|
|
|
|
assert label == "A"
|
|
assert confidence > 0.95
|
|
assert coverage > 0
|
|
|
|
|
|
def test_classify_choice_sorted_option_letters(monkeypatch):
|
|
"""option_letters passed to parse_logprobs are sorted."""
|
|
import local_decision
|
|
|
|
captured: dict[str, Any] = {}
|
|
|
|
def fake_post(url, *, json, timeout):
|
|
captured["url"] = url
|
|
captured["json"] = json
|
|
captured["timeout"] = timeout
|
|
return _FakeResponse(_make_response([
|
|
_make_logprob(
|
|
"B",
|
|
-0.02,
|
|
[
|
|
{"token": "A", "logprob": -4.0},
|
|
{"token": "B", "logprob": -0.02},
|
|
],
|
|
),
|
|
]))
|
|
|
|
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
|
|
|
# Unsorted option insertion order; parse must sort to A, B.
|
|
label, _, _ = local_decision.classify_choice(
|
|
"Refactor",
|
|
{"B": "second", "A": "first"},
|
|
base_url="http://ollama:11434",
|
|
model="m",
|
|
num_ctx=4096,
|
|
timeout_s=30,
|
|
)
|
|
assert label == "B"
|
|
|
|
|
|
def test_classify_choice_propagates_request_exception(monkeypatch):
|
|
"""requests.exceptions.RequestException propagates out of classify_choice."""
|
|
import requests as real_requests
|
|
|
|
import local_decision
|
|
|
|
def fake_post(url, *, json, timeout):
|
|
raise real_requests.exceptions.ConnectionError("boom")
|
|
|
|
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
|
|
|
with pytest.raises(real_requests.exceptions.RequestException):
|
|
local_decision.classify_choice(
|
|
"x",
|
|
{"A": "a"},
|
|
base_url="http://ollama:11434",
|
|
model="m",
|
|
num_ctx=4096,
|
|
timeout_s=30,
|
|
)
|
|
|
|
|
|
def test_classify_choice_coverage_min_passthrough(monkeypatch):
|
|
"""coverage_min is forwarded to parse_logprobs."""
|
|
import local_decision
|
|
|
|
captured: dict[str, Any] = {}
|
|
|
|
def fake_post(url, *, json, timeout):
|
|
captured["timeout"] = timeout
|
|
return _FakeResponse(_make_response([
|
|
_make_logprob(
|
|
"A",
|
|
-10.0, # very low mass -> fails high coverage_min
|
|
[{"token": "A", "logprob": -10.0}],
|
|
),
|
|
]))
|
|
|
|
monkeypatch.setattr(local_decision.requests, "post", fake_post)
|
|
|
|
with pytest.raises(RuntimeError, match="coverage"):
|
|
local_decision.classify_choice(
|
|
"x",
|
|
{"A": "a"},
|
|
base_url="http://ollama:11434",
|
|
model="m",
|
|
num_ctx=4096,
|
|
timeout_s=30,
|
|
coverage_min=10.0,
|
|
)
|