"""Replaceable local repository for immutable private activity records."""

from __future__ import annotations

from typing import Iterable, Protocol

from .contracts import (
    Activity,
    ActivityDuplicateRelationship,
    ActivityTelemetry,
    ExternalProviderIdentity,
    RawProviderRecord,
    User,
    UserProfile,
)


class ActivityRepository(Protocol):
    def get_user(self, user_id: str) -> User: ...
    def get_profile(self, user_id: str) -> UserProfile | None: ...
    def list_activities(self, user_id: str) -> tuple[Activity, ...]: ...
    def get_activity(self, user_id: str, activity_id: str) -> Activity: ...
    def list_activity_telemetry(self, user_id: str, activity_id: str) -> tuple[ActivityTelemetry, ...]: ...
    def list_activity_sources(self, user_id: str) -> tuple[ExternalProviderIdentity, ...]: ...


class LocalActivityRepository:
    """Append-only local records; raw records and canonical activities are never rewritten."""

    def __init__(
        self,
        *,
        users: Iterable[User] = (),
        profiles: Iterable[UserProfile] = (),
        provider_identities: Iterable[ExternalProviderIdentity] = (),
        raw_records: Iterable[RawProviderRecord] = (),
        activities: Iterable[Activity] = (),
        telemetry: Iterable[ActivityTelemetry] = (),
        duplicate_relationships: Iterable[ActivityDuplicateRelationship] = (),
    ):
        users, profiles, provider_identities = tuple(users), tuple(profiles), tuple(provider_identities)
        raw_records, activities, telemetry, duplicate_relationships = tuple(raw_records), tuple(activities), tuple(telemetry), tuple(duplicate_relationships)
        self._users = self._index_unique(users, lambda item: item.user_id, "user IDs")
        self._profiles = self._index_unique(profiles, lambda item: item.user_id, "profile user IDs")
        self._provider_identities = self._index_unique(
            provider_identities, lambda item: item.provider_identity_id, "provider identity IDs"
        )
        self._raw_records = self._index_unique(raw_records, lambda item: item.raw_record_id, "raw record IDs")
        self._activities = self._index_unique(activities, lambda item: item.activity_id, "activity IDs")
        self._telemetry = self._index_unique(telemetry, lambda item: item.telemetry_id, "telemetry IDs")
        self._duplicates = self._index_unique(
            duplicate_relationships, lambda item: item.relationship_id, "duplicate relationship IDs"
        )
        if any(profile.user_id not in self._users for profile in profiles):
            raise ValueError("profile belongs to unknown user")
        if any(identity.user_id not in self._users for identity in provider_identities):
            raise ValueError("provider identity belongs to unknown user")
        for activity in self._activities.values():
            self._validate_activity(activity)
        for record in self._raw_records.values():
            if record.user_id not in self._users:
                raise ValueError("raw record belongs to unknown user")
        for series in self._telemetry.values():
            self._validate_telemetry(series)
        for relationship in self._duplicates.values():
            self._validate_duplicate_relationship(relationship)

    def get_user(self, user_id: str) -> User:
        return self._users[user_id]

    def get_profile(self, user_id: str) -> UserProfile | None:
        self.get_user(user_id)
        return self._profiles.get(user_id)

    def list_activities(self, user_id: str) -> tuple[Activity, ...]:
        self.get_user(user_id)
        return tuple(item for item in self._activities.values() if item.user_id == user_id)

    def get_activity(self, user_id: str, activity_id: str) -> Activity:
        activity = self._activities[activity_id]
        if activity.user_id != user_id:
            raise KeyError(activity_id)
        return activity

    def list_activity_telemetry(self, user_id: str, activity_id: str) -> tuple[ActivityTelemetry, ...]:
        self.get_activity(user_id, activity_id)
        return tuple(item for item in self._telemetry.values() if item.activity_id == activity_id)

    def list_activity_sources(self, user_id: str) -> tuple[ExternalProviderIdentity, ...]:
        self.get_user(user_id)
        return tuple(item for item in self._provider_identities.values() if item.user_id == user_id)

    def add_raw_record(self, record: RawProviderRecord) -> RawProviderRecord:
        if record.raw_record_id in self._raw_records:
            raise ValueError("raw provider record is immutable and already exists")
        if record.user_id not in self._users:
            raise ValueError("raw record belongs to unknown user")
        self._raw_records[record.raw_record_id] = record
        return record

    def add_activity(self, activity: Activity) -> Activity:
        if activity.activity_id in self._activities:
            raise ValueError("canonical activity is immutable and already exists")
        self._validate_activity(activity)
        self._activities[activity.activity_id] = activity
        return activity

    def add_telemetry(self, telemetry: ActivityTelemetry) -> ActivityTelemetry:
        if telemetry.telemetry_id in self._telemetry:
            raise ValueError("telemetry dataset is immutable and already exists")
        self._validate_telemetry(telemetry)
        self._telemetry[telemetry.telemetry_id] = telemetry
        return telemetry

    def add_duplicate_relationship(self, relationship: ActivityDuplicateRelationship) -> ActivityDuplicateRelationship:
        if relationship.relationship_id in self._duplicates:
            raise ValueError("duplicate relationship is immutable and already exists")
        self._validate_duplicate_relationship(relationship)
        self._duplicates[relationship.relationship_id] = relationship
        return relationship

    def _validate_activity(self, activity: Activity) -> None:
        if activity.user_id not in self._users:
            raise ValueError("activity belongs to unknown user")
        for source_reference in activity.provenance.source_references:
            raw_record = self._raw_records.get(source_reference.raw_record_id)
            if raw_record is None:
                raise ValueError("activity provenance references unknown raw record")
            self._validate_source_reference(activity.user_id, source_reference, raw_record)

    def _validate_telemetry(self, telemetry: ActivityTelemetry) -> None:
        activity = self._activities.get(telemetry.activity_id)
        provenance = telemetry.provenance
        if activity is None:
            raise ValueError("telemetry refers to an unknown canonical activity")
        raw_record = self._raw_records.get(provenance.raw_record_id)
        if raw_record is None:
            raise ValueError("telemetry provenance references unknown raw record")
        self._validate_source_reference(activity.user_id, provenance.source_reference, raw_record)
        activity_sources = activity.provenance.source_references
        if not any(
            source.source_kind is provenance.source_reference.source_kind
            and source.provider == provenance.source_reference.provider
            and source.source_activity_id == provenance.source_reference.source_activity_id
            for source in activity_sources
        ):
            raise ValueError("telemetry source identity differs from canonical activity provenance")

    def _validate_duplicate_relationship(self, relationship: ActivityDuplicateRelationship) -> None:
        activity = self._activities.get(relationship.activity_id)
        related_activity = self._activities.get(relationship.related_activity_id)
        if activity is None or related_activity is None:
            raise ValueError("duplicate relationship refers to unknown activity")
        if activity.user_id != related_activity.user_id:
            raise ValueError("duplicate relationship cannot cross user boundaries")

    @staticmethod
    def _index_unique(items: tuple, identity, label: str) -> dict:
        indexed = {identity(item): item for item in items}
        if len(indexed) != len(items):
            raise ValueError(f"{label} must be unique")
        return indexed

    @staticmethod
    def _validate_source_reference(
        user_id: str, source_reference, raw_record: RawProviderRecord
    ) -> None:
        if raw_record.user_id != user_id:
            raise ValueError("source record belongs to a different user than canonical activity")
        if source_reference.source_kind is not raw_record.source_kind:
            raise ValueError("source reference kind differs from raw record")
        if source_reference.provider != raw_record.provider:
            raise ValueError("source reference provider differs from raw record")
        if source_reference.source_activity_id != raw_record.source_activity_id:
            raise ValueError("source reference activity identity differs from raw record")
