bl_info = {
    "name": "Sculpt Fill Tool",
    "author": "GG & ChatGPT",
    "version": (1, 1, 0),
    "blender": (4, 5, 0),
    "location": "3D View > Sculpt Mode (Tool Shelf)",
    "description": "Flood-fill sculpt mask or vertex paint confined to area under cursor (cached, optimized, supports CORNER + POINT color)",
    "category": "Sculpt",
    "support": "COMMUNITY",
}

import bpy
from bpy_extras import view3d_utils
from collections import deque

# -----------------------------
# Settings
MASK_STRENGTH = 1.0  # Mask strength
# -----------------------------

# Global caches
ADJACENCY_CACHE = {}
VERTEX_LOOPS_CACHE = {}

# ---------------------------------------------------
# Helper: Active sculpt tool name
# ---------------------------------------------------
def get_active_sculpt_tool_name(context=None):
    context = context or bpy.context
    obj = context.object
    workspace = context.workspace
    if not obj or obj.type != 'MESH' or context.mode != 'SCULPT':
        return None
    try:
        tool = workspace.tools.from_space_view3d_mode(context.mode)
        name = getattr(tool, "idname", str(tool))
        return name.lower().split('.')[-1]
    except:
        return None

# ---------------------------------------------------
# sRGB→Linear
# ---------------------------------------------------
def srgb_to_linear(c):
    if c <= 0.04045:
        return c / 12.92
    else:
        return ((c + 0.055)/1.055) ** 2.4

# ---------------------------------------------------
# Auto cache invalidation
# ---------------------------------------------------
def invalidate_if_topology_changed(obj):
    obj_id = id(obj)
    mesh = obj.data
    if obj_id in ADJACENCY_CACHE:
        if len(ADJACENCY_CACHE[obj_id]) != len(mesh.vertices):
            del ADJACENCY_CACHE[obj_id]
    if obj_id in VERTEX_LOOPS_CACHE:
        if len(VERTEX_LOOPS_CACHE[obj_id]) != len(mesh.vertices):
            del VERTEX_LOOPS_CACHE[obj_id]

# ---------------------------------------------------
# Raycast helper
# ---------------------------------------------------
def raycast_face(context, mouse, obj):
    region = context.region
    rv3d = context.space_data.region_3d
    coord = mouse

    view_vector = view3d_utils.region_2d_to_vector_3d(region, rv3d, coord)
    ray_origin = view3d_utils.region_2d_to_origin_3d(region, rv3d, coord)
    ray_target = ray_origin + view_vector * 10000

    mw = obj.matrix_world
    mwi = mw.inverted()

    ro_local = mwi @ ray_origin
    rt_local = mwi @ ray_target
    direction = (rt_local - ro_local).normalized()

    hit, loc, normal, face_index = obj.ray_cast(ro_local, direction)
    if not hit:
        return None, None

    return face_index, mw @ loc

# ---------------------------------------------------
# Cached adjacency
# ---------------------------------------------------
def get_adjacency(obj):
    obj_id = id(obj)
    if obj_id not in ADJACENCY_CACHE:
        mesh = obj.data
        adjacency = [[] for _ in range(len(mesh.vertices))]
        for poly in mesh.polygons:
            verts = poly.vertices
            n = len(verts)
            for i in range(n):
                a = verts[i]
                b = verts[(i+1)%n]
                adjacency[a].append(b)
                adjacency[b].append(a)
        ADJACENCY_CACHE[obj_id] = adjacency
    return ADJACENCY_CACHE[obj_id]

# ---------------------------------------------------
# Cached vertex → loops
# ---------------------------------------------------
def get_vertex_loops(obj):
    obj_id = id(obj)
    if obj_id not in VERTEX_LOOPS_CACHE:
        mesh = obj.data
        vertex_loops = [[] for _ in range(len(mesh.vertices))]
        for li, loop in enumerate(mesh.loops):
            vertex_loops[loop.vertex_index].append(li)
        VERTEX_LOOPS_CACHE[obj_id] = vertex_loops
    return VERTEX_LOOPS_CACHE[obj_id]

# ---------------------------------------------------
# Mask access
# ---------------------------------------------------
def get_sculpt_mask(obj):
    mesh = obj.data
    attr = mesh.attributes.get(".sculpt_mask")
    if attr is None:
        return None
    return [elem.value for elem in attr.data]

# ---------------------------------------------------
# Color attribute detection/creation
# ---------------------------------------------------
def get_or_create_color_attribute(mesh):
    for attr in mesh.color_attributes:
        if attr.data_type in {'BYTE_COLOR', 'FLOAT_COLOR'}:
            mesh.color_attributes.active_color = attr
            return attr

    attr = mesh.color_attributes.new(
        name="Color",
        type='FLOAT_COLOR',
        domain='POINT'
    )
    for d in attr.data:
        d.color = (1, 1, 1, 1)

    mesh.color_attributes.active_color = attr
    bpy.context.space_data.shading.color_type = 'VERTEX'

    return attr

# ---------------------------------------------------
# Unified Cached Flood-Fill Engine
# ---------------------------------------------------
def flood_fill_cached(start_vertex, adjacency, accept_fn, rim_accept_fn=None):
    visited = set()
    result = set()

    stack = deque([start_vertex])
    rim_stack = set()

    while stack or rim_stack:
        if stack:
            v = stack.pop()
            if v in visited:
                continue
            visited.add(v)
            result.add(v)

            if accept_fn(v):
                for n in adjacency[v]:
                    if n not in visited:
                        stack.append(n)
            else:
                if rim_accept_fn:
                    for n in adjacency[v]:
                        if n not in visited:
                            rim_stack.add(n)

        else:
            v = rim_stack.pop()
            if v in visited:
                continue
            visited.add(v)
            result.add(v)

    return result

# ---------------------------------------------------
# Mask Flood Adapter
# ---------------------------------------------------
def flood_fill_mask(obj, start_vertices):
    mesh = obj.data
    mask_attr = mesh.attributes.get(".sculpt_mask")
    if mask_attr is None:
        return []

    mask = [elem.value for elem in mask_attr.data]
    adjacency = get_adjacency(obj)

    result = set()
    for sv in start_vertices:
        result |= flood_fill_cached(
            sv,
            adjacency,
            accept_fn=lambda v: mask[v] < 0.85,
            rim_accept_fn=lambda v: True,
        )
    return list(result)

# ---------------------------------------------------
# Paint Flood Adapter
# ---------------------------------------------------
def flood_fill_color(obj, start_vertex, base_color, COLOR_THRESHOLD=0.02):
    mesh = obj.data
    color_attr = get_or_create_color_attribute(mesh)
    adjacency = get_adjacency(obj)
    vertex_loops = get_vertex_loops(obj)

    vcolors_cache = {}

    def get_vcolor(v):
        if v in vcolors_cache:
            return vcolors_cache[v]

        if color_attr.domain == 'CORNER':
            loops = vertex_loops[v]
            r = g = b = 0.0
            for li in loops:
                c = color_attr.data[li].color
                r += c[0]; g += c[1]; b += c[2]
            n = len(loops)
            col = (r/n, g/n, b/n)
        else:
            c = color_attr.data[v].color
            col = (c[0], c[1], c[2])

        vcolors_cache[v] = col
        return col

    def color_diff(c1, c2):
        return abs(c1[0]-c2[0]) + abs(c1[1]-c2[1]) + abs(c1[2]-c2[2])

    verts = flood_fill_cached(
        start_vertex,
        adjacency,
        accept_fn=lambda v: color_diff(get_vcolor(v), base_color) <= COLOR_THRESHOLD,
        rim_accept_fn=None
    )

    if color_attr.domain == 'CORNER':
        loops = []
        for v in verts:
            loops.extend(vertex_loops[v])
        return loops

    return list(verts)

# ---------------------------------------------------
# Operator
# ---------------------------------------------------
class SCULPT_OT_confined_fill(bpy.types.Operator):
    bl_idname = "sculpt.confined_fill"
    bl_label = "Confined Mask / Paint Fill"
    bl_options = {'REGISTER','UNDO'}

    mouse = None

    def invoke(self, context, event):
        if context.mode != 'SCULPT':
            self.report({'ERROR'}, "Must be in Sculpt Mode")
            return {'CANCELLED'}
        brush = context.tool_settings.sculpt.brush
        if not brush:
            self.report({'ERROR'}, "No active brush")
            return {'CANCELLED'}
        self.mouse = (event.mouse_region_x, event.mouse_region_y)
        return self.execute(context)

    def execute(self, context):
        obj = context.active_object
        
        # -----------------------------
        # Multi-Res Check
        # -----------------------------
        if any(m.type == 'MULTIRES' for m in obj.modifiers):
            self.report({'WARNING'}, "Fill tool is not supported on a Multi-Res mesh")
            return {'CANCELLED'}

        invalidate_if_topology_changed(obj)
        brush = context.tool_settings.sculpt.brush
        s = brush.strength

        # ---------------------------------------------------
        # MASK
        # ---------------------------------------------------
        if brush.sculpt_tool == 'MASK':
            mesh = obj.data
            mask_attr = mesh.attributes.get(".sculpt_mask")

            if mask_attr is None:
                mask_attr = mesh.attributes.new(
                    name=".sculpt_mask",
                    type='FLOAT',
                    domain='POINT'
                )
                for elem in mask_attr.data:
                    elem.value = 0.0

            bpy.ops.ed.undo_push(message="Confined Mask Fill")

            face_index, hit_loc = raycast_face(context, self.mouse, obj)
            if face_index is None:
                self.report({'WARNING'}, "Click on the mesh surface.")
                return {'CANCELLED'}

            poly = mesh.polygons[face_index]
            mask_values = get_sculpt_mask(obj)

            start_vertices = [v for v in poly.vertices if mask_values[v] < 0.0001]
            if not start_vertices:
                start_vertices = list(poly.vertices)

            verts_to_mask = flood_fill_mask(obj, start_vertices)

            for v in verts_to_mask:
                mask_attr.data[v].value = MASK_STRENGTH

            mesh.attributes.update()
            obj.update_from_editmode()
            context.view_layer.update()
            bpy.ops.sculpt.mask_filter(filter_type='CONTRAST_INCREASE',
                                       iterations=3,
                                       auto_iteration_count=False)

            self.report({'INFO'}, f"Masked {len(verts_to_mask)} vertices (Strength={MASK_STRENGTH})")
            return {'FINISHED'}

        # ---------------------------------------------------
        # PAINT
        # ---------------------------------------------------
        elif brush.sculpt_tool == 'PAINT':
            mesh = obj.data
            color_attr = get_or_create_color_attribute(mesh)
            bpy.ops.ed.undo_push(message="Confined Paint Fill")

            face_index, hit_loc = raycast_face(context, self.mouse, obj)
            if face_index is None:
                self.report({'WARNING'}, "Click on mesh surface")
                return {'CANCELLED'}

            poly = mesh.polygons[face_index]
            start_vertex = min(poly.vertices, key=lambda v: (mesh.vertices[v].co - hit_loc).length)
            vertex_loops = get_vertex_loops(obj)

            if color_attr.domain == 'CORNER':
                loops = vertex_loops[start_vertex]
                r = g = b = 0.0
                for li in loops:
                    c = color_attr.data[li].color
                    r+=c[0]; g+=c[1]; b+=c[2]
                n = len(loops)
                base_color = (r/n, g/n, b/n)
            else:
                c = color_attr.data[start_vertex].color
                base_color = (c[0], c[1], c[2])

            # Step 1: Flood-fill
            targets = set(flood_fill_color(obj, start_vertex, base_color))

            # Step 2: One extra outer ring
            adjacency = get_adjacency(obj)
            extra_layer = set()
            for v in targets:
                for n in adjacency[v]:
                    if n not in targets:
                        extra_layer.add(n)
            targets.update(extra_layer)
            targets = list(targets)

            brush_color = brush.color
            r_lin = srgb_to_linear(brush_color[0])
            g_lin = srgb_to_linear(brush_color[1])
            b_lin = srgb_to_linear(brush_color[2])

            if color_attr.domain == 'CORNER':
                for li in targets:
                    old = color_attr.data[li].color
                    color_attr.data[li].color = (
                        old[0]*(1-s) + r_lin*s,
                        old[1]*(1-s) + g_lin*s,
                        old[2]*(1-s) + b_lin*s,
                        1.0
                    )
            else:
                for v in targets:
                    old = color_attr.data[v].color
                    color_attr.data[v].color = (
                        old[0]*(1-s) + r_lin*s,
                        old[1]*(1-s) + g_lin*s,
                        old[2]*(1-s) + b_lin*s,
                        1.0
                    )

            mesh.color_attributes.update()
            obj.update_from_editmode()
            context.view_layer.update()

            self.report({'INFO'},
                f"Painted {len(targets)} {'loops' if color_attr.domain=='CORNER' else 'vertices'}")
            return {'FINISHED'}

        else:
            self.report({'ERROR'}, "Active brush must be Mask or Paint type")
            return {'CANCELLED'}

# ---------------------------------------------------
# Workspace Tool
# ---------------------------------------------------
class SCULPT_TOOL_confined_fill(bpy.types.WorkSpaceTool):
    bl_space_type = 'VIEW_3D'
    bl_context_mode = 'SCULPT'
    bl_idname = "sculpt.confined_fill_tool"
    bl_label = "Fill"
    bl_description = "Flood-fill mask or vertex paint confined to area under cursor"
    bl_icon = "brush.paint_texture.fill"
    bl_keymap = (
        ("sculpt.confined_fill", {"type":'LEFTMOUSE','value':'PRESS'}, None),
    )

# ---------------------------------------------------
# Register
# ---------------------------------------------------
def register():
    bpy.utils.register_class(SCULPT_OT_confined_fill)
    bpy.utils.register_tool(SCULPT_TOOL_confined_fill,
                            after={"builtin.select_box"},
                            separator=True)

def unregister():
    bpy.utils.unregister_tool(SCULPT_TOOL_confined_fill)
    bpy.utils.unregister_class(SCULPT_OT_confined_fill)
