import os
import cv2
import json

# -------------------- ПАРАМЕТРЫ --------------------
INPUT_DIR = "input"
OUTPUT_DIR = "output"
LOGO_PATH = "logo/logo.png"

FPS_ANALYZE = 1       # кадр/сек для редкого сканирования
BOOST_FPS = 5         # кадр/сек для локального ускоренного сканирования
BOOST_WINDOW = 2      # секунды до и после найденного кадра
THRESHOLD = 0.8       # confidence для template matching
MIN_DURATION = 1.0    # минимальная длительность появления логотипа (сек)
# ---------------------------------------------------

os.makedirs(OUTPUT_DIR, exist_ok=True)

# Загружаем эталон логотипа
logo = cv2.imread(LOGO_PATH, cv2.IMREAD_GRAYSCALE)

# -------------------- ФУНКЦИИ ----------------------
def analyze_video(video_path):
    cap = cv2.VideoCapture(video_path)
    fps = cap.get(cv2.CAP_PROP_FPS)
    frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
    frame_id = 0
    hits = []

    print(f"\nАнализ: {os.path.basename(video_path)}")

    while True:
        ret, frame = cap.read()
        if not ret:
            break

        frame_id += 1
        if frame_id % int(fps / FPS_ANALYZE) != 0:
            continue

        gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
        res = cv2.matchTemplate(gray, logo, cv2.TM_CCOEFF_NORMED)
        _, max_val, _, _ = cv2.minMaxLoc(res)

        if max_val >= THRESHOLD:
            # локальное ускоренное сканирование
            local_hits = boost_scan(cap, fps, frame_id, BOOST_WINDOW, BOOST_FPS)
            hits.extend(local_hits)

    cap.release()
    hits = sorted(list(set(hits)))  # убираем дубликаты
    intervals = merge_intervals(hits)
    # фильтруем короткие интервалы
    intervals = [i for i in intervals if i[1] - i[0] >= MIN_DURATION]
    return intervals

def boost_scan(cap, fps, frame_id, window_sec, boost_fps):
    """Ускоренная проверка вокруг найденного кадра"""
    hits = []
    start_frame = max(frame_id - int(window_sec * fps), 0)
    end_frame = min(frame_id + int(window_sec * fps), int(cap.get(cv2.CAP_PROP_FRAME_COUNT))-1)

    current_pos = cap.get(cv2.CAP_PROP_POS_FRAMES)

    step = max(int(fps / boost_fps), 1)

    for fid in range(start_frame, end_frame + 1, step):
        cap.set(cv2.CAP_PROP_POS_FRAMES, fid)
        ret, frame = cap.read()
        if not ret:
            continue
        gray = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
        res = cv2.matchTemplate(gray, logo, cv2.TM_CCOEFF_NORMED)
        _, max_val, _, _ = cv2.minMaxLoc(res)
        if max_val >= THRESHOLD:
            hits.append(round(fid / fps, 2))

    cap.set(cv2.CAP_PROP_POS_FRAMES, current_pos)
    return hits

def merge_intervals(timestamps, max_gap=1.0):
    if not timestamps:
        return []

    timestamps.sort()
    intervals = []
    start = timestamps[0]
    end = timestamps[0]

    for t in timestamps[1:]:
        if t - end <= max_gap:
            end = t
        else:
            intervals.append([round(start,2), round(end,2)])
            start = end = t
    intervals.append([round(start,2), round(end,2)])
    return intervals

def print_intervals(video_name, intervals):
    if not intervals:
        print("  Логотип не найден.")
    else:
        print("  Интервалы появления логотипа:")
        for start, end in intervals:
            print(f"    [{start} – {end}] сек")

# -------------------- MAIN ------------------------
def main():
    videos = [f for f in os.listdir(INPUT_DIR) if f.endswith(".mp4")]
    all_results = {"videos": []}

    if not videos:
        print("В папке input нет видео.")
        return

    for v in videos:
        path = os.path.join(INPUT_DIR, v)
        intervals = analyze_video(path)
        print_intervals(v, intervals)
        all_results["videos"].append({
            "video": v,
            "logo_intervals_sec": intervals
        })

    # сохраняем JSON для всех видео
    out_path = os.path.join(OUTPUT_DIR, "results.json")
    with open(out_path, "w", encoding="utf-8") as f:
        json.dump(all_results, f, indent=2)

    print(f"\nГотово! JSON сохранен: {out_path}")

if __name__ == "__main__":
    main()
