import random
from dataclasses import dataclass
from js import document, window

@dataclass(frozen=True)
class Literal:
    var: str
    negated: bool = False

@dataclass(frozen=True)
class Clause:
    literals: frozenset
    def is_empty(self):
        return len(self.literals) == 0


class KnowledgeBase:
    def __init__(self):
        self.clauses = set()
        self.inference_steps = 0

    def tell(self, clause):
        if clause.literals:
            self.clauses.add(clause)

    def ask_safe(self, r, c):
        self.inference_steps = 0
        return self._prove_not(f"P_{r}_{c}") and self._prove_not(f"W_{r}_{c}")

    def ask_true(self, var):
        self.inference_steps = 0
        return self._prove(var)

    def _prove(self, var):
        clauses = set(self.clauses)
        clauses.add(Clause(frozenset({Literal(var, True)})))
        return self._resolve_all(clauses)

    def _prove_not(self, var):
        clauses = set(self.clauses)
        clauses.add(Clause(frozenset({Literal(var, False)})))
        return self._resolve_all(clauses)

    def _resolve_all(self, clauses):
        new = set(clauses)
        while True:
            added = False
            lst = list(new)

            for i in range(len(lst)):
                for j in range(i + 1, len(lst)):
                    self.inference_steps += 1
                    res = self._resolve(lst[i], lst[j])
                    if res is None:
                        continue
                    if res.is_empty():
                        return True
                    if res not in new:
                        new.add(res)
                        added = True

            if not added:
                return False

    def _resolve(self, c1, c2):
        for l1 in c1.literals:
            for l2 in c2.literals:
                if l1.var == l2.var and l1.negated != l2.negated:
                    return Clause(frozenset((c1.literals | c2.literals) - {l1, l2}))
        return None


class WumpusEnv:
    def __init__(self, r, c):
        self.r = r
        self.c = c
        self.grid = [[{"pit": False, "wumpus": False, "visited": False} for _ in range(c)] for _ in range(r)]
        self.agent = (0, 0)
        self.kb = KnowledgeBase()
        self.moves = 0
        self.current_percepts = {"breeze": False, "stench": False}
        self.confirmed_pits = set()
        self.confirmed_wumpus = set()

    def generate(self):
        wr, wc = random.randint(0, self.r-1), random.randint(0, self.c-1)
        self.grid[wr][wc]["wumpus"] = True

        for _ in range((self.r * self.c) // 6):
            pr, pc = random.randint(0, self.r-1), random.randint(0, self.c-1)
            if (pr, pc) != (0, 0):
                self.grid[pr][pc]["pit"] = True

        self.kb.tell(Clause(frozenset({Literal("P_0_0", True)})))
        self.kb.tell(Clause(frozenset({Literal("W_0_0", True)})))
        self.update_percepts()

    def percepts(self, r, c):
        breeze = stench = False
        for dr, dc in [(-1,0),(1,0),(0,-1),(0,1)]:
            nr, nc = r+dr, c+dc
            if 0 <= nr < self.r and 0 <= nc < self.c:
                breeze |= self.grid[nr][nc]["pit"]
                stench |= self.grid[nr][nc]["wumpus"]
        return breeze, stench

    def update_percepts(self):
        r, c = self.agent
        breeze, stench = self.percepts(r, c)

        self.current_percepts = {"breeze": breeze, "stench": stench}
        neighbors = self.get_neighbors(r, c)

        if breeze:
            self.kb.tell(Clause(frozenset({Literal(f"P_{nr}_{nc}", False) for nr, nc in neighbors})))
        else:
            for nr, nc in neighbors:
                self.kb.tell(Clause(frozenset({Literal(f"P_{nr}_{nc}", True)})))

        if stench:
            self.kb.tell(Clause(frozenset({Literal(f"W_{nr}_{nc}", False) for nr, nc in neighbors})))
        else:
            for nr, nc in neighbors:
                self.kb.tell(Clause(frozenset({Literal(f"W_{nr}_{nc}", True)})))

        for nr, nc in neighbors:
            if self.kb.ask_true(f"P_{nr}_{nc}"):
                self.confirmed_pits.add((nr, nc))
            if self.kb.ask_true(f"W_{nr}_{nc}"):
                self.confirmed_wumpus.add((nr, nc))

    def get_neighbors(self, r, c):
        return [(r+dr, c+dc) for dr, dc in [(-1,0),(1,0),(0,-1),(0,1)]
                if 0 <= r+dr < self.r and 0 <= c+dc < self.c]

    def move(self, r, c):
        self.agent = (r, c)
        self.grid[r][c]["visited"] = True
        self.moves += 1
        self.update_percepts()


class Agent:
    def __init__(self, env):
        self.env = env
        self.visited = {(0, 0)}
        self.stack = [(0, 0)]

    def safe_moves(self, r, c):
        res = []
        for nr, nc in self.env.get_neighbors(r, c):
            if (nr, nc) not in self.visited and self.env.kb.ask_safe(nr, nc):
                res.append((nr, nc))
        return res

    def step(self):
        r, c = self.env.agent
        options = self.safe_moves(r, c)

        if options:
            nr, nc = options[0]
            self.stack.append((r, c))
        elif self.stack:
            nr, nc = self.stack.pop()
        else:
            return "❌ Stopped"

        self.env.move(nr, nc)
        self.visited.add((nr, nc))
        return f"{nr},{nc}"


env = None
agent = None


def on_start(*args):
    global env, agent
    r = int(document.getElementById("rows").value)
    c = int(document.getElementById("cols").value)
    env = WumpusEnv(r, c)
    env.generate()
    agent = Agent(env)
    render()
    document.getElementById("status-message").innerText = "Started"


def on_step(*args):
    document.getElementById("status-message").innerText = agent.step()
    render()


def render():
    if not env:
        return

    canvas = document.getElementById("canvas")
    ctx = canvas.getContext("2d")
    cw = canvas.width / env.c
    ch = canvas.height / env.r

    ctx.clearRect(0, 0, canvas.width, canvas.height)

    for r in range(env.r):
        for c in range(env.c):
            x, y = c * cw, r * ch
            cell = env.grid[r][c]

            if (r, c) == env.agent:
                ctx.fillStyle = "blue"
            elif cell["visited"]:
                ctx.fillStyle = "green"
            elif (r, c) in env.confirmed_pits or (r, c) in env.confirmed_wumpus:
                ctx.fillStyle = "red"
            else:
                ctx.fillStyle = "gray"

            ctx.fillRect(x, y, cw, ch)
            ctx.strokeRect(x, y, cw, ch)

    document.getElementById("inference-count").innerText = str(env.kb.inference_steps)
    document.getElementById("moves-count").innerText = str(env.moves)

    p = env.current_percepts
    document.getElementById("percepts-list").innerHTML = "<br>".join(
        [x for x in ["Breeze" if p["breeze"] else "", "Stench" if p["stench"] else ""] if x]
    ) or "None"


window.on_start = on_start
window.on_step = on_step