import re
import uuid
from enum import StrEnum
from typing import Any

from pydantic import BaseModel, Field, model_validator


class RiskLevel(StrEnum):
    low = "low"
    medium = "medium"
    high = "high"


class AccessType(StrEnum):
    read = "read"
    write = "write"


class PlanStep(BaseModel):
    step_id: str = Field(min_length=1, max_length=80, pattern=r"^[A-Za-z0-9._-]+$")
    tool_name: str = Field(min_length=3, max_length=120, pattern=r"^[a-z][a-z0-9_.-]+$")
    arguments: dict[str, Any] = Field(default_factory=dict)
    concise_rationale: str = Field(min_length=1, max_length=500)
    depends_on: list[str] = Field(default_factory=list, max_length=20)


class SupervisorPlan(BaseModel):
    summary: str = Field(min_length=1, max_length=1000)
    steps: list[PlanStep] = Field(default_factory=list, max_length=25)
    response_outline: str = Field(min_length=1, max_length=1000)

    @model_validator(mode="after")
    def validate_graph(self):
        step_ids = [step.step_id for step in self.steps]
        if len(step_ids) != len(set(step_ids)):
            raise ValueError("Plan step identifiers must be unique")
        known = set(step_ids)
        dependencies = {step.step_id: set(step.depends_on) for step in self.steps}
        for step_id, required in dependencies.items():
            if step_id in required or not required <= known:
                raise ValueError("Plan dependencies must reference other plan steps")

        visiting: set[str] = set()
        visited: set[str] = set()

        def visit(step_id: str) -> None:
            if step_id in visiting:
                raise ValueError("Plan dependencies must be acyclic")
            if step_id in visited:
                return
            visiting.add(step_id)
            for dependency in dependencies[step_id]:
                visit(dependency)
            visiting.remove(step_id)
            visited.add(step_id)

        for step_id in step_ids:
            visit(step_id)
        return self


class PlanRequest(BaseModel):
    execution_id: uuid.UUID
    workspace_id: uuid.UUID
    actor_id: uuid.UUID
    instruction: str = Field(min_length=1, max_length=20_000)
    model: str | None = Field(default=None, min_length=1, max_length=160)
    allowed_tools: list[str] = Field(default_factory=list, max_length=100)
    safety_identifier: str = Field(min_length=16, max_length=64, pattern=r"^[a-f0-9]+$")

    @model_validator(mode="after")
    def normalize_tools(self):
        if len(self.allowed_tools) != len(set(self.allowed_tools)):
            raise ValueError("Allowed tools must be unique")
        for name in self.allowed_tools:
            if not re.fullmatch(r"[a-z][a-z0-9_.-]{2,119}", name):
                raise ValueError("Allowed tool name is invalid")
        return self


class Usage(BaseModel):
    input_tokens: int = Field(default=0, ge=0)
    output_tokens: int = Field(default=0, ge=0)
    total_tokens: int = Field(default=0, ge=0)


class PlanResult(BaseModel):
    provider: str
    model: str
    provider_response_id: str | None = None
    plan: SupervisorPlan
    usage: Usage = Field(default_factory=Usage)
