72 lines
3.3 KiB
Python
72 lines
3.3 KiB
Python
|
|
"""
|
||
|
|
Tests for DAAT — Direction-Anchored Ambiguity Triage.
|
||
|
|
"""
|
||
|
|
import pytest
|
||
|
|
from malkhut.daat.core import DaatQuery, DaatVerdict, daat_classify, _cosine_similarity
|
||
|
|
|
||
|
|
|
||
|
|
class TestCosineSimilarity:
|
||
|
|
def test_identical_vectors(self):
|
||
|
|
assert abs(_cosine_similarity([1, 0, 0], [1, 0, 0]) - 1.0) < 1e-9
|
||
|
|
|
||
|
|
def test_orthogonal_vectors(self):
|
||
|
|
assert abs(_cosine_similarity([1, 0], [0, 1])) < 1e-9
|
||
|
|
|
||
|
|
def test_opposite_vectors(self):
|
||
|
|
assert abs(_cosine_similarity([1, 0], [-1, 0]) - (-1.0)) < 1e-9
|
||
|
|
|
||
|
|
def test_zero_vector(self):
|
||
|
|
assert _cosine_similarity([0, 0], [1, 1]) == 0.0
|
||
|
|
|
||
|
|
|
||
|
|
class TestDaatClassify:
|
||
|
|
def test_known_state(self):
|
||
|
|
"""Live state matches explored state closely."""
|
||
|
|
explored = [DaatQuery(1.0, 1000000, 0.1, 5.0, 0.5, 0.3, 50.0, 0.5)]
|
||
|
|
magnitudes = [sum(abs(x) for x in [1.0, 1000000, 0.1, 5.0, 0.5, 0.3, 50.0, 0.5])]
|
||
|
|
query = DaatQuery(1.1, 900000, 0.12, 5.2, 0.48, 0.31, 52.0, 0.48)
|
||
|
|
result = daat_classify(query, explored, magnitudes)
|
||
|
|
assert result.verdict == DaatVerdict.KNOWN
|
||
|
|
assert result.cosine_sim > 0.99
|
||
|
|
|
||
|
|
def test_out_of_distribution(self):
|
||
|
|
"""Opposite-direction vector → low cosine → OOD."""
|
||
|
|
explored = [DaatQuery(1.0, 1000000, 0.1, 5.0, 0.5, 0.3, 50.0, 0.5)]
|
||
|
|
magnitudes = [sum(abs(x) for x in [1.0, 1000000, 0.1, 5.0, 0.5, 0.3, 50.0, 0.5])]
|
||
|
|
# Very different features → low cosine → OOD
|
||
|
|
query = DaatQuery(100.0, 1.0, 0.9, 0.1, 10.0, 0.1, 1.0, 0.1)
|
||
|
|
result = daat_classify(query, explored, magnitudes)
|
||
|
|
# Direction is very different (depth is tiny, spread is huge)
|
||
|
|
assert result.verdict == DaatVerdict.OUT_OF_DISTRIBUTION
|
||
|
|
|
||
|
|
def test_empty_explored(self):
|
||
|
|
query = DaatQuery(1.0, 1000000, 0.1, 5.0, 0.5, 0.3, 50.0, 0.5)
|
||
|
|
result = daat_classify(query, [], [])
|
||
|
|
assert result.verdict == DaatVerdict.OUT_OF_DISTRIBUTION
|
||
|
|
|
||
|
|
def test_marginal_state(self):
|
||
|
|
"""High cosine but extreme magnitude → magnitude gate catches it."""
|
||
|
|
explored = [DaatQuery(1.0, 1000000, 0.1, 5.0, 0.5, 0.3, 50.0, 0.5)]
|
||
|
|
magnitudes = [sum(abs(x) for x in [1.0, 1000000, 0.1, 5.0, 0.5, 0.3, 50.0, 0.5])]
|
||
|
|
# Same direction (all positive) but very different spread/depth ratio
|
||
|
|
query = DaatQuery(50.0, 10000.0, 0.5, 0.5, 5.0, 0.5, 5.0, 0.5)
|
||
|
|
result = daat_classify(query, explored, magnitudes)
|
||
|
|
# Magnitude gate should catch this — query magnitude is very different
|
||
|
|
# The cosine is high (same direction), but magnitude_ratio should be != 1.0
|
||
|
|
print(f' cosine={result.cosine_sim:.6f} mag_ratio={result.magnitude_ratio:.4f}')
|
||
|
|
print(f' verdict={result.verdict}')
|
||
|
|
# At minimum: magnitude_ratio should NOT be 1.0
|
||
|
|
assert result.magnitude_ratio != 1.0
|
||
|
|
|
||
|
|
def test_result_fields(self):
|
||
|
|
"""Result has all required fields."""
|
||
|
|
explored = [DaatQuery(1.0, 1000000, 0.1, 5.0, 0.5, 0.3, 50.0, 0.5)]
|
||
|
|
magnitudes = [1.0]
|
||
|
|
query = DaatQuery(1.0, 1000000, 0.1, 5.0, 0.5, 0.3, 50.0, 0.5)
|
||
|
|
result = daat_classify(query, explored, magnitudes)
|
||
|
|
assert isinstance(result.verdict, DaatVerdict)
|
||
|
|
assert isinstance(result.cosine_sim, float)
|
||
|
|
assert isinstance(result.magnitude_ratio, float)
|
||
|
|
assert isinstance(result.confidence, float)
|
||
|
|
assert 0.0 <= result.confidence <= 1.0
|