#!/usr/bin/env python3

import argparse as ap
import os
from tkinter import *
from tkinter import filedialog
from tkinter.ttk import *

from unscore import *

class ScrollableFrame(Frame):
    def _on_mousewheel(self, e):
        delta = e.delta
        if e.num == 5 or delta == -120:
            delta = -1
        if e.num == 4 or delta == 120:
            delta = 1
        self.canvas.yview_scroll(-delta, "units")

    def _mouse_enter(self, *args, **kwargs):
        for button in ("<MouseWheel>", "<Button-4>", "<Button-5>",):
            self.bind_all(button, self._on_mousewheel)

    def _mouse_leave(self, *args, **kwargs):
        for button in ("<MouseWheel>", "<Button-4>", "<Button-5>",):
            self.unbind_all(button)

    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self.rowconfigure(0, weight=1)
        self.columnconfigure(0, weight=1)

        self.canvas = Canvas(self)
        self.canvas.grid(row=0, sticky=N+E+S+W)

        self.scrollbar = Scrollbar(self, orient="vertical", command=self.canvas.yview)
        self.scrollbar.grid(row=0, column=1, sticky=N+E+S+W)

        self.inner = Frame(self.canvas)
        self.canvas.create_window((0, 0), window=self.inner, anchor=N+W)
        self.canvas.configure(yscrollcommand=self.scrollbar.set)
        self.bind("<Enter>", self._mouse_enter)
        self.bind("<Leave>", self._mouse_leave)

        self.inner.bind(
            "<Configure>",
            lambda _: self.canvas.configure(scrollregion=self.canvas.bbox("all"))
        )

folder_name = None
folder_name_set = False
def choose_folder(*args, **kwargs):
    global folder_name, folder_name_set

    fname = filedialog.askdirectory(initialdir=folder_name)
    if not fname:
        return

    folder_name_set = True
    folder_name = fname
    folder_name_var.set(fname)

    reload_folder()

video_fname = None
def choose_video(*args, **kwargs):
    global video_fname, folder_name

    fname = filedialog.askopenfilename(initialfile=video_fname)
    if not fname:
        return

    video_fname = fname
    video_fname_var.set(fname)

    if not folder_name_set:
        folder_name = fname + "_scenes"
        folder_name_var.set(folder_name)

    reload_video()

scenes_list = []

## 1. Load scenes using one of the two ways:
def reload_video(*aargs, **kwargs):
    """Find scenes from a video file."""
    global scenes_list

    os.makedirs(folder_name, exist_ok=True)

    scenes_list = find_scenes(
        video_fname, 
        ffmpeg_cutoff_scale.get(),
        line=bool(remove_line.get()),
        scene_cutoff=cutoff_scale.get()
    )
    for n, s in enumerate(scenes_list, 1):
        s.save(os.path.join(folder_name, "out%07d.%s" % (n, "png")))

    update_splits()

def reload_folder(*aargs, **kwargs):
    """Load scenes from a folder."""
    global scenes_list

    print_status("Reading scenes from folder...")
    scenes_list = []
    for f in sorted(os.listdir(folder_name)):
        if not f.endswith('.' + "png"):
            continue
        scenes_list.append(Image.open(os.path.join(folder_name, f)))

    update_splits()

scenes = []
old_split = None
old_number = None

def print_status(s):
    print(s)
    status.set(s)
    root.update()

def update_splits():
    global scenes, old_split, old_number

    print_status("Splitting scenes...")

    padding = int(padding_scale.get())
    split = int(split_scale.get())
    do_number = int(number.get())
    if split == old_split and old_number == do_number:
        return
    old_split = split
    old_number = do_number

    scenes = []

    if scenes_list:
        width = max(s.width for s in scenes_list)
    else:
        width = 0
    newbreaks = []
    for s in scenes_list:
        new = []

        s = crop_scene(s, width)

        todo = [s]
        if split != -1:
            todo = split_scene(s, split)

        for t in todo:
            new.append(pad_scene(t, padding))

        scenes.append(new)

    if do_number:
        number_images([i for j in scenes for i in j], 1, in_place=True)

    update_break_widgets()

    update_button["state"] = "active"
    rescan_button["state"] = "active"

    print_status("Done splitting!")

## We are assembling two types of things: scenes and lines, their subdivisions. We
## run into two situations when updating the available line breaks: 
##
## 1. We reloaded the scenes
## 2. We re-split the scenes into lines
##
## In case 2, in all likelihood most of the scenes will remain the same; thus, we
## should preserve as many of the existing forced breaks as possible. In situation
## 1, we just clear them all.
breaks = []
break_widgets = []

def update_break_widgets(clear=False):
    global breaks, break_widgets

    ## Remove old breaks
    break_values = []
    for lines, lines_v in zip(break_widgets, breaks):
        break_values.append([i.get() == 1 for i in lines_v])
        for w in lines:
            w.grid_forget()

    if clear:
        break_values = [[-1]*len(i) for i in scenes]

    if len(break_values) < len(scenes):
        break_values = [[]] * len(scenes)

    breaks = []
    break_widgets = []
    i = 0
    for lines, old_values in zip(scenes, break_values):
        breaks.append([])
        break_widgets.append([])
        old_values = old_values + [-1]*max(0, len(lines) - len(old_values))
        for l, ov in zip(lines, old_values):
            v = IntVar()
            v.set(ov)
            a = Checkbutton(frame.inner, text="Line " + str(i+1), onvalue=1, offvalue=-1, variable=v)
            a.grid(row=i+1, sticky=N+E+S+W)
            frame.inner.rowconfigure(i+1, weight=1)
            break_widgets[-1].append(a)
            breaks[-1].append(v)
            i += 1

old_split = -1

## 2. Update output on user change
def update(*aargs, **kwargs):
    global output_fname

    aspect = aspect_scale.get()
    savgol = savgol_scale.get()
    do_otsu = finalize.get()

    if output_fname is None:
        output_fname = filedialog.asksaveasfilename(initialdir=folder_name)

    ## Handles numbering and line split distance
    update_splits()

    ## We need to handle merging and any postprocessing
    flat_scenes = [i for j in scenes for i in j]
    forced_breaks = [n for n, v in enumerate((i for j in breaks for i in j)) if v.get() > 0]

    print_status("Writing temporary scenes...")

    with tempfile.TemporaryDirectory() as outdir:
        outfiles = []
        fmt = "out%07d.png"

        for n, done in enumerate(merge_lines(flat_scenes, aspect=aspect, center=False, center_breaks=False, breaks=forced_breaks), 1):
            name = os.path.join(outdir, fmt % n)
            outfiles.append(name)
            if do_otsu:
                done = otsu(done.resize((done.width*2, done.height*2), resample=Image.Resampling.LANCZOS))
            if savgol > 0:
                done = do_savgol(done, savgol)
            done.save(name)

        print_status("Merging output...")

        subp.check_call(["magick", "-density", "300", *outfiles, os.path.join(folder_name, output_fname)])
        
        print_status("Done!")

parser = ap.ArgumentParser()
parser.add_argument("-i", "--input", help="Input video to process")
parser.add_argument("-o", "--output", default=None, help="Output filename (defaults to <outdir>/out.pdf)")
parser.add_argument("-f", "--format", help="Output format", default="png")
parser.add_argument("-t", "--dir", help="Intermediate directory")

args = parser.parse_args()

input_file = args.input
mid_dir = args
output_file = args.output

root = Tk()
root.columnconfigure(0, weight=1)

video_frame = Frame(root)
video_frame.grid(row=0, sticky=N+E+S+W)
video_frame.columnconfigure(0, weight=1)

Label(video_frame, text="FFmpeg cutoff:").grid(row=0, sticky=N+E+S+W)
ffmpeg_cutoff_scale = Scale(video_frame, from_=0.01, to=0.1)
ffmpeg_cutoff_scale.grid(row=1, sticky=N+E+S+W)
ffmpeg_cutoff_scale.set(0.03)

Label(video_frame, text="Scene cutoff:").grid(row=2, sticky=N+E+S+W)
cutoff_scale = Scale(video_frame, from_=0.1, to=3)
cutoff_scale.grid(row=3, sticky=N+E+S+W)
cutoff_scale.set(2)

remove_line = IntVar()
remove_line.set(0)
Checkbutton(root, text="Remove line", variable=remove_line, onvalue=1, offvalue=0).grid(row=4, sticky=N+E+S+W)

video_fname_var = StringVar(video_frame)
video_fname_var.set("Select a video to load...")
if args.input:
    video_fname_var.set(args.input)
    video_fname = args.input
    root.after(0, reload_video)
Label(video_frame, textvariable=video_fname_var).grid(row=5, sticky=N+E+S+W)
Button(video_frame, text="Choose Video", command=choose_video).grid(row=6, sticky=N+E+S+W)

folder_frame = Frame(root)
folder_frame.grid(row=1, sticky=N+E+S+W)
folder_frame.columnconfigure(0, weight=1)

folder_name_var = StringVar(video_frame)
folder_name_var.set("Select a folder to load...")
if args.dir:
    folder_name_var.set(args.dir)
    folder_name = args.dir
    folder_name_set = True
    if not args.input:
        root.after(0, reload_folder)
Label(folder_frame, textvariable=folder_name_var).grid(row=0, sticky=N+E+S+W)
Button(folder_frame, text="Choose Folder", command=choose_folder).grid(row=1, sticky=N+E+S+W)
rescan_button = Button(folder_frame, text="Rescan", command=reload_folder)
rescan_button.grid(row=2, sticky=N+E+S+W)
rescan_button["state"] = "disabled"

output_fname = None
if args.output:
    output_fname = args.output

rest_frame = Frame(root)
rest_frame.columnconfigure(0, weight=1)
rest_frame.grid(row=2, sticky=N+E+S+W)

Label(rest_frame, text="Padding:").grid(row=0, sticky=N+E+S+W)
padding_scale = Scale(rest_frame, from_=5, to=100)
padding_scale.grid(row=1, sticky=N+E+S+W)
padding_scale.set(10)

Label(rest_frame, text="Aspect:").grid(row=2, sticky=N+E+S+W)
aspect_scale = Scale(rest_frame, from_=0.5, to=0.9)
aspect_scale.grid(row=3, sticky=N+E+S+W)
aspect_scale.set(8.5/11)

Label(rest_frame, text="Split:").grid(row=4, sticky=N+E+S+W)
split_scale = Scale(rest_frame, from_=1, to=100)
split_scale.grid(row=5, sticky=N+E+S+W)
split_scale.set(50)

Label(rest_frame, text="Savgol:").grid(row=6, sticky=N+E+S+W)
savgol_scale = Scale(rest_frame, from_=0, to=1)
savgol_scale.grid(row=7, sticky=N+E+S+W)
savgol_scale.set(0.7)
try:
    import scipy
except ModuleNotFoundError:
    savgol_scale.set(0)
    savgol_scale["state"] = "disabled"

number = IntVar()
number.set(0)
Checkbutton(rest_frame, text="Number lines", variable=number, onvalue=1, offvalue=0).grid(row=8, sticky=N+E+S+W)

finalize = IntVar()
finalize.set(0)
Checkbutton(rest_frame, text="Finalize", variable=finalize, onvalue=1, offvalue=0).grid(row=9, sticky=N+E+S+W)

frame = ScrollableFrame(root)
frame.grid(row=3, sticky=N+E+S+W)
frame.inner.columnconfigure(0, weight=1)
root.rowconfigure(3, weight=1)

Label(frame.inner, text="Force page break after:").grid(row=0, sticky=N+E+S+W)

update_button = Button(root, text="Update", command=update, state="disabled")
update_button.grid(row=4, sticky=N+E+S+W)

status = StringVar(root)
status.set("Ready!")

Label(root, textvariable=status).grid(row=5, sticky=N+E+S+W)

root.wm_title("Unscore")
root.mainloop()
