Files
6krrt/tests/test_local_decision_parse.py

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,
)