import Combine
import ConcurrencyExtras
import FirebaseClient
import Foundation

// MARK: - AppStartUpOrchestratorProtocol

protocol AppStartUpOrchestratorProtocol {
    func start(
        eventHandler: @escaping (AppStartUpOrchestrationEvent) async -> Void
    ) async
}

// MARK: - AppStartUpOrchestrationEvent

enum AppStartUpOrchestrationEvent: Equatable {
    case phaseStarted(AppStartUpPhase)
    case sdkInitialized(SDKIdentifier)
    case criticalFailure([SDKIdentifier])
    case canShowUI
    case completed(AppStartUpMetrics)
}

// MARK: - AppStartUpOrchestrationError

enum AppStartUpOrchestrationError: LocalizedError {
    case timeout(duration: TimeInterval)
    case sdkInitializationFailed(sdk: String, underlying: Error)

    public var errorDescription: String? {
        switch self {
        case .timeout(let duration):
            return "Operation timed out after \(duration) seconds"
        case .sdkInitializationFailed(let sdk, let underlying):
            return "SDK \(sdk) initialization failed: \(underlying.localizedDescription)"
        }
    }
}

// MARK: - AppStartUpOrchestrator

final class AppStartUpOrchestrator: AppStartUpOrchestratorProtocol {
    private let sdkStates: LockIsolated<[SDKIdentifier: SDKState]>
    private let activeTasks: LockIsolated<[SDKIdentifier: Task<Void, Error>]>
    private let tracesInitialized: LockIsolated<TracesState>

    private let sdkManager: SDKManagerProtocol
    private let metricsCollector: MetricsCollectorProtocol
    private let dependencies: LockIsolated<[SDKIdentifier: SDKDependency]>
    private let firebaseClient: FirebaseClient
    private let now: () -> Date

    // MARK: - TracesState

    private struct TracesState {
        // Start startupOrchestration span for entire process
        var startupTrace: AnyCancellable?

        // Start timeToInteractive span until UI can be shown
        var timeToInteractiveTrace: AnyCancellable?
    }

    init(
        sdkManager: SDKManagerProtocol,
        metricsCollector: MetricsCollectorProtocol,
        dependencies: [SDKIdentifier: SDKDependency],
        sdkStates: [SDKIdentifier: SDKState],
        firebaseClient: FirebaseClient,
        now: @escaping () -> Date = { .now }
    ) {
        self.sdkManager = sdkManager
        self.metricsCollector = metricsCollector
        self.dependencies = .init(dependencies)
        self.sdkStates = .init(sdkStates)
        self.firebaseClient = firebaseClient
        self.activeTasks = .init([:])
        self.tracesInitialized = .init(TracesState())
        self.now = now
    }

    deinit {
        let tasks = activeTasks.value
        for (_, task) in tasks {
            task.cancel()
        }
    }

    func start(
        eventHandler: @escaping (AppStartUpOrchestrationEvent) async -> Void
    ) async {
        let startTime = now()
        metricsCollector.startupBegan(at: startTime)

        let enhancedEventHandler = createTraceAwareEventHandler(eventHandler: eventHandler)

        await executePhase(.essential, eventHandler: enhancedEventHandler)
        await executePhase(.critical, eventHandler: enhancedEventHandler)

        await eventHandler(.canShowUI)

        // End timeToInteractive span when UI can be shown
        tracesInitialized.withValue { state in
            state.timeToInteractiveTrace?.cancel()
            state.timeToInteractiveTrace = nil
        }

        await executePhase(.important, eventHandler: enhancedEventHandler)
        await executePhase(.optional, eventHandler: enhancedEventHandler)

        let totalTime = now().timeIntervalSince(startTime)
        let metrics = metricsCollector.getMetrics(totalDuration: totalTime)
        await eventHandler(.completed(metrics))

        // End startup trace
        tracesInitialized.withValue { state in
            state.startupTrace?.cancel()
            state.startupTrace = nil
        }
    }

    private func createTraceAwareEventHandler(
        eventHandler: @escaping (AppStartUpOrchestrationEvent) async -> Void
    ) -> (AppStartUpOrchestrationEvent) async -> Void {
        return { [weak self] event in
            guard let self else { return }

            // Check if we need to initialize traces when Firebase is ready
            if case .sdkInitialized(let sdkId) = event, sdkId == .firebase {
                self.initializeTracesIfNeeded()
            }

            // Always forward the event to the original handler
            await eventHandler(event)
        }
    }

    private func initializeTracesIfNeeded() {
        // Only initialize traces if Firebase performance is ready
        guard firebaseClient.isPerformanceReady() else { return }

        tracesInitialized.withValue { state in
            // Initialize startup trace if not already created
            if state.startupTrace == nil {
                state.startupTrace = firebaseClient.trace(.startupOrchestration)
            }

            // Initialize time to interactive trace if not already created
            if state.timeToInteractiveTrace == nil {
                state.timeToInteractiveTrace = firebaseClient.trace(.timeToInteractive)
            }
        }
    }

    private func executePhase(
        _ phase: AppStartUpPhase,
        eventHandler: @escaping (AppStartUpOrchestrationEvent) async -> Void
    ) async {
        await eventHandler(.phaseStarted(phase))

        let sdksForPhase = getSDKsForPhase(phase)
        guard !sdksForPhase.isEmpty else { return }

        // Start phase-specific span (only if Firebase is ready)
        let phaseTrace = if let span = phase.span, firebaseClient.isPerformanceReady() {
            firebaseClient.trace(span)
        } else {
            AnyCancellable {}
        }
        defer { phaseTrace.cancel() }

        // Initialize SDKs concurrently with dependency management
        await withTaskGroup(of: Void.self) { group in
            for sdk in sdksForPhase {
                group.addTask {
                    await self.initializeSDKWithDependencies(
                        sdk: sdk,
                        eventHandler: eventHandler
                    )
                    await eventHandler(.sdkInitialized(sdk))
                }
            }
        }
    }

    private func initializeSDKWithDependencies(
        sdk: SDKIdentifier,
        eventHandler: @escaping (AppStartUpOrchestrationEvent) async -> Void
    ) async {
        guard let dependency = dependencies[sdk] else { return }

        // Wait for dependencies with span tracking (only if Firebase is ready)
        for requiredSDK in dependency.requirements {
            let dependencyWaitTrace = if firebaseClient.isPerformanceReady() {
                firebaseClient.trace(.sdk_dependency_wait_time)
            } else {
                AnyCancellable {}
            }
            await waitForSDKCompletion(requiredSDK)
            dependencyWaitTrace.cancel()
        }

        // Initialize SDK with span tracking
        let startTime = now()
        metricsCollector.sdkStarted(sdk, at: startTime)
        sdkStates.withValue { $0[sdk] = .initializing(startTime: startTime) }

        // Start SDK-specific span (only if Firebase is ready)
        let sdkTrace = if firebaseClient.isPerformanceReady() {
            firebaseClient.trace(sdk.span)
        } else {
            AnyCancellable {}
        }
        defer { sdkTrace.cancel() }

        let task = Task<Void, Error> {
            try await withTimeout(dependency.timeout) { [weak self] in
                guard let self else { return }
                try await sdkManager.initialize(
                    sdk,
                    onMainThread: dependency.mustBeOnMainThread
                )
            }
        }

        activeTasks.withValue { $0[sdk] = task }

        do {
            try await task.value
            let endTime = now()
            let duration = endTime.timeIntervalSince(startTime)

            sdkStates.withValue { $0[sdk] = .completed(duration: duration) }
            metricsCollector.sdkCompleted(sdk, duration: duration, success: true)
        } catch {
            let endTime = now()
            let duration = endTime.timeIntervalSince(startTime)

            if dependency.isOptional {
                sdkStates.withValue { $0[sdk] = .skipped(reason: error.localizedDescription) }
                metricsCollector.sdkSkipped(sdk, reason: error.localizedDescription)
            } else {
                sdkStates.withValue { $0[sdk] = .failed(error: error.localizedDescription, retryCount: 0) }
                metricsCollector.sdkFailed(sdk, error: error, duration: duration)
                await eventHandler(.criticalFailure([sdk]))
            }
        }

        _ = activeTasks.withValue { $0.removeValue(forKey: sdk) }
    }

    private func waitForSDKCompletion(_ sdk: SDKIdentifier) async {
        while let sdkState = sdkStates[sdk], !sdkState.isComplete {
            try? await Task.sleep(for: .milliseconds(10))
        }
    }

    private func getSDKsForPhase(_ phase: AppStartUpPhase) -> [SDKIdentifier] {
        SDKIdentifier.allCases.filter { sdk in
            dependencies[sdk]?.phase == phase
        }
    }
}

private extension AppStartUpPhase {
    var span: TracePerformanceSpan? {
        switch self {
        case .essential:
            return .essential_phase_completion
        case .critical:
            return .critical_phase_completion
        case .important:
            return .important_phase_completion
        case .optional:
            return .optional_phase_completion
        case .notStarted, .complete:
            return nil
        }
    }
}

private extension SDKIdentifier {
    var span: TracePerformanceSpan {
        switch self {
        case .firebase:
            return .sdk_Firebase_init
        case .braze:
            return .sdk_Braze_Init
        case .clerk:
            return .sdk_Clerk_init
        case .revenueCat:
            return .sdk_RevenueCat_Init
        case .statsig:
            return .sdk_Statsig_Init
        case .adamantium:
            return .sdk_Adamantium_init
        case .shareAsset:
            return .sdk_ShareAsset_init
        case .attribution:
            return .sdk_Attribution_init
        case .rageshake:
            return .sdk_Rageshake_init
        }
    }
}

// MARK: - Timeout Utility

private func withTimeout<T>(
    _ duration: TimeInterval,
    operation: @escaping () async throws -> T
) async throws -> T {
    try await withThrowingTaskGroup(of: T.self) { group in
        group.addTask {
            try await operation()
        }

        group.addTask {
            try await Task.sleep(for: .seconds(duration))
            throw AppStartUpOrchestrationError.timeout(duration: duration)
        }

        guard let result = try await group.next() else {
            throw AppStartUpOrchestrationError.timeout(duration: duration)
        }

        group.cancelAll()
        return result
    }
}

#if DEBUG

    // MARK: - Mock

    extension AppStartUpOrchestrator {
        final class Mock: AppStartUpOrchestratorProtocol {
            private let events: [AppStartUpOrchestrationEvent]

            init(events: [AppStartUpOrchestrationEvent]) {
                self.events = events
            }

            func start(eventHandler: @escaping (AppStartUpOrchestrationEvent) async -> Void) async {
                for event in events {
                    await eventHandler(event)
                }
            }
        }
    }
#endif
