"""
Trinary Encoding for EMF Breath
Schema: INHALE=0, EXHALE=1, HOLD=2
Applied across Layer 1 (Schumann) + Layer 2 (Individual) for each breath cycle.
"""
import math
from typing import Literal

Trit = Literal[0, 1, 2]
Trits = list[Trit]

TRINARY_MAP = {"INHALE": 0, "EXHALE": 1, "HOLD": 2}
REVERSE_MAP = {0: "INHALE", 1: "EXHALE", 2: "HOLD"}


def encode_trinary(phases: list[str]) -> list[Trit]:
    """Convert phase names to trinary digits."""
    result = []
    for p in phases:
        if p not in TRINARY_MAP:
            raise ValueError(f"Unknown phase: {p}. Must be INHALE, EXHALE, or HOLD.")
        result.append(TRINARY_MAP[p])
    return result


def decode_trinary(trits: list[int]) -> list[str]:
    """Convert trinary digits back to phase names."""
    return [REVERSE_MAP.get(t, "HOLD") for t in trits]


def trits_to_bits(trits: list[Trit]) -> list[int]:
    """
    Convert trits to bits using balanced trinary expansion.
    Each trit encodes 1.585 bits of information.
    Returns a list of 0/1 bits.
    """
    bits = []
    for trit in trits:
        # Write trit as 2 bits with one redundant code for trinary
        if trit == 0:
            bits.extend([0, 0])
        elif trit == 1:
            bits.extend([0, 1])
        else:  # 2
            bits.extend([1, 0])
    return bits


def bits_to_trits(bits: list[int]) -> list[Trit]:
    """Convert bits back to trits."""
    trits = []
    for i in range(0, len(bits) - 1, 2):
        a, b = bits[i], bits[i + 1]
        if a == 0 and b == 0:
            trits.append(0)
        elif a == 0 and b == 1:
            trits.append(1)
        elif a == 1 and b == 0:
            trits.append(2)
        else:
            trits.append(2)  # redundant code
    return trits


def encode_message(trits: list[Trit]) -> dict:
    """
    Encode a message as trinary breath sequence.
    Each breath cycle produces 3 trits (one per beat).
    trits_per_breath_cycle = 3.
    """
    bits = trits_to_bits(trits)
    return {
        "trits": trits,
        "bits": bits,
        "trit_count": len(trits),
        "bit_count": len(bits),
        "bits_per_trit": len(bits) / len(trits) if trits else 0,
    }


def decode_message(bits: list[int]) -> dict:
    """Decode bits back to trits then to phase names."""
    trits = bits_to_trits(bits)
    phases = decode_trinary(trits)
    return {
        "trits": trits,
        "phases": phases,
        "trit_count": len(trits),
    }


def example_cycle() -> dict:
    """Demonstrate a full breath cycle encoding."""
    # 3 trits per breath cycle (inhale, exhale, hold)
    phases = ["INHALE", "EXHALE", "HOLD"]
    trits = encode_trinary(phases)
    encoded = encode_message(trits)
    decoded = decode_message(encoded["bits"])
    return {
        "input_phases": phases,
        "trits": trits,
        "bits": encoded["bits"],
        "decoded_phases": decoded["phases"],
        "bits_per_trit": f"{encoded['bits_per_trit']:.3f}",
    }
