"""Load ambient adapter facts once, outside deterministic assignment logic."""
from __future__ import annotations

from collections.abc import Iterable

from .assignment_resolver import (
    AssignmentContext,
    AssignmentEnvironment,
    TerminalBackend,
)
from .domain.provider import ExecutionNormalization, UnknownProviderError
from .domain.role import ROLE_DEFINITIONS
from .domain.worker_exec import WorkerWriteCapability
from .model_pool import ModelPool
from .registry.host_registry import default_host_registry
from .registry.provider_registry import default_provider_registry


def load_assignment_context(
    *,
    host_runtime: str,
    terminal_backend: TerminalBackend,
    execution_provider_ids: Iterable[str] | None = None,
    include_native_provider: bool = True,
) -> AssignmentContext:
    """Snapshot provider, host, terminal, and execution inputs for one caller."""
    provider_registry = default_provider_registry()
    pool = ModelPool.from_registry(provider_registry)
    adapter = default_host_registry(provider_registry).resolve(host_runtime)
    worker_port = adapter.worker_dispatch()
    capability_reader = getattr(worker_port, "worker_write_capability", None)
    native_capability = (
        capability_reader()
        if callable(capability_reader)
        else WorkerWriteCapability("none")
    )
    roles = tuple(row.id for row in ROLE_DEFINITIONS)
    if execution_provider_ids is None:
        provider_ids = provider_registry.ids()
    else:
        requested = [
            provider.strip().lower() for provider in execution_provider_ids
        ]
        if include_native_provider:
            requested.insert(0, adapter.descriptor.native_provider_id)
        requested_ids = dict.fromkeys(requested)
        selected_ids = []
        for provider_id in requested_ids:
            if not provider_id:
                continue
            try:
                provider_registry.resolve(provider_id)
            except UnknownProviderError:
                continue
            selected_ids.append(provider_id)
        provider_ids = tuple(selected_ids)
    execution_resolutions = {}
    for provider_id in provider_ids:
        provider = provider_registry.resolve(provider_id)
        models_by_id = {
            model.model_id: model for model in provider.models.values()
        }
        models = tuple(models_by_id.values())
        normalizer = provider.execution_normalizer
        if normalizer is None:
            rows = {
                (model.model_id, role): ExecutionNormalization(
                    model.execution_value
                )
                for model in models
                for role in roles
            }
        else:
            try:
                rows = normalizer.snapshot(models, roles)
            except Exception as exc:
                rows = {
                    (model.model_id, role): ExecutionNormalization(None, str(exc))
                    for model in models
                    for role in roles
                }
        for model in models:
            for role in roles:
                key = (model.model_id, role)
                outcome = rows.get(key)
                if outcome is None:
                    outcome = ExecutionNormalization(
                        None,
                        f"provider {provider_id!r} did not snapshot model "
                        f"{model.model_id!r} for role {role!r}",
                    )
                execution_resolutions[(f"{provider_id}/{model.model_id}", role)] = (
                    outcome
                )
    environment = AssignmentEnvironment(
        host_descriptor=adapter.descriptor,
        host_model_port=adapter.host_model(),
        native_worker_write_capability=native_capability,
        terminal_backend=terminal_backend,
        execution_resolutions=execution_resolutions,
    )
    return AssignmentContext(pool=pool, environment=environment)
