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