# SPDX-License-Identifier: MIT
"""Contract tests for the Langfuse v4 fixture adapter."""

from dataclasses import dataclass, field
import json
from pathlib import Path
import sys
from typing import Dict, List, Mapping, Optional
import unittest

sys.path.insert(0, str(Path(__file__).parent))

from langfuse_adapter import LangfuseLevel, LangfuseTraceSink  # noqa: E402
from workflow import (  # noqa: E402
    FixtureTools,
    ObservationKind,
    SupportRequest,
    run_support_workflow,
)


def fixture_request() -> SupportRequest:
    return SupportRequest(
        question="Can roman.test@example.com refund acct_DEMO123?",
        email="roman.test@example.com",
        account_id="acct_DEMO123",
    )


@dataclass
class FakeObservation:
    trace_id: str
    id: str
    updates: List[Dict[str, object]] = field(default_factory=list)
    ended: bool = False

    def update(
        self,
        *,
        metadata: Mapping[str, object],
        level: LangfuseLevel,
        status_message: Optional[str] = None,
    ) -> "FakeObservation":
        self.updates.append(
            {
                "metadata": dict(metadata),
                "level": level,
                "status_message": status_message,
            }
        )
        return self

    def end(self) -> "FakeObservation":
        self.ended = True
        return self


@dataclass
class FakeClient:
    starts: List[Dict[str, object]] = field(default_factory=list)
    observations: List[FakeObservation] = field(default_factory=list)
    flush_count: int = 0

    def start_observation(
        self,
        *,
        name: str,
        as_type: ObservationKind,
        trace_context: Optional[Mapping[str, str]],
        metadata: Mapping[str, object],
        version: str,
        model: Optional[str] = None,
    ) -> FakeObservation:
        index = len(self.observations) + 1
        trace_id = trace_context["trace_id"] if trace_context else "a" * 32
        observation = FakeObservation(trace_id, f"{index:016x}")
        self.starts.append(
            {
                "name": name,
                "as_type": as_type,
                "trace_context": dict(trace_context) if trace_context else None,
                "metadata": dict(metadata),
                "version": version,
                "model": model,
            }
        )
        self.observations.append(observation)
        return observation

    def flush(self) -> None:
        self.flush_count += 1


class FailingClient(FakeClient):
    def start_observation(
        self,
        *,
        name: str,
        as_type: ObservationKind,
        trace_context: Optional[Mapping[str, str]],
        metadata: Mapping[str, object],
        version: str,
        model: Optional[str] = None,
    ) -> FakeObservation:
        del name, as_type, trace_context, metadata, version, model
        raise ConnectionError("fixture Langfuse unavailable")


class LangfuseAdapterTest(unittest.TestCase):
    def test_children_use_explicit_root_parentage(self) -> None:
        client = FakeClient()
        sink = LangfuseTraceSink(client)
        run_support_workflow(fixture_request(), FixtureTools(), sink)

        root = client.observations[0]
        self.assertIsNone(client.starts[0]["trace_context"])
        self.assertTrue(
            all(
                start["trace_context"]
                == {
                    "trace_id": root.trace_id,
                    "parent_span_id": root.id,
                }
                for start in client.starts[1:]
            )
        )

    def test_retry_creates_separate_tool_observations(self) -> None:
        client = FakeClient()
        sink = LangfuseTraceSink(client)
        run_support_workflow(
            fixture_request(),
            FixtureTools(subscription_failures_remaining=1),
            sink,
        )

        attempts = [
            (start, observation)
            for start, observation in zip(client.starts, client.observations)
            if start["name"] == "tool:get-subscription"
        ]
        self.assertEqual(len(attempts), 2)
        self.assertEqual(attempts[0][1].updates[0]["level"], "ERROR")
        self.assertEqual(attempts[1][1].updates[0]["level"], "DEFAULT")
        self.assertTrue(all(observation.ended for _, observation in attempts))

    def test_adapter_payload_is_redacted_before_export(self) -> None:
        client = FakeClient()
        run_support_workflow(
            fixture_request(),
            FixtureTools(),
            LangfuseTraceSink(client),
        )

        serialized = json.dumps(
            {
                "starts": client.starts,
                "updates": [item.updates for item in client.observations],
            },
            sort_keys=True,
        )
        self.assertNotIn("roman.test@example.com", serialized)
        self.assertNotIn("acct_DEMO123", serialized)
        self.assertIn("[EMAIL_REDACTED]", serialized)
        self.assertIn("[ACCOUNT_REDACTED]", serialized)

    def test_terminal_error_does_not_fail_siblings(self) -> None:
        client = FakeClient()
        sink = LangfuseTraceSink(client)
        run_support_workflow(
            fixture_request(),
            FixtureTools(refund_terminal_error=True),
            sink,
        )

        levels = {
            start["metadata"]["logical_observation_id"]: observation.updates[0][
                "level"
            ]
            for start, observation in zip(client.starts, client.observations)
        }
        self.assertEqual(levels["refund-1"], "ERROR")
        self.assertEqual(levels["retrieval-1"], "DEFAULT")
        self.assertEqual(levels["plan-1"], "DEFAULT")
        self.assertEqual(levels["answer-1"], "DEFAULT")
        self.assertEqual(levels["task-1"], "WARNING")

    def test_client_failure_cannot_change_application_result(self) -> None:
        expected = run_support_workflow(
            fixture_request(),
            FixtureTools(),
            LangfuseTraceSink(FakeClient()),
        )
        actual = run_support_workflow(
            fixture_request(),
            FixtureTools(),
            LangfuseTraceSink(FailingClient()),
        )
        self.assertEqual(actual, expected)

    def test_flush_is_explicit(self) -> None:
        client = FakeClient()
        sink = LangfuseTraceSink(client)
        sink.flush()
        self.assertEqual(client.flush_count, 1)


if __name__ == "__main__":
    unittest.main()
