#!/usr/bin/env python3
"""Download a document from a Telegram channel we are a member of (pyrogram).

Resumable: if --out exists, download continues from the last 1MiB boundary
(Telegram GetFile is offset-addressable). Safe to Ctrl-C / crash / rerun —
no progress lost. file_reference expiry mid-download is handled by re-fetching
the message and continuing from the current offset.

Resolves the channel via GetDialogs (populates peer access_hash in storage).

Usage:
  .venv/bin/python scripts/tg_download.py --chat "Wings Daily Updates FREE" \
      --msg-id 2687 --out findings/telegram_logs/wingscloud/wingscloud_ULP_July_13.txt

  .venv/bin/python scripts/tg_download.py --chat "Wings Daily Updates FREE" --list 40
      # list recent messages (id/date/document) without downloading

  --fresh   ignore existing partial file, restart from zero
"""
import argparse
import asyncio
import os
import sys
import time
from pathlib import Path

sys.path.insert(0, str(Path(__file__).parent))
from tg_session_load import build_session_file, load_dotenv, ENV_FILE, SESSION_DIR

CHUNK = 1024 * 1024  # must match pyrogram get_file chunk_size


async def resolve_channel(app, title: str):
    """Find channel by title via GetDialogs and cache it in pyrogram storage.

    Raw invoke() bypasses pyrogram's peer parser, so we must insert the peer
    manually — otherwise resolve_peer() later raises PeerIdInvalid."""
    from pyrogram.raw import functions, types
    r = await app.invoke(functions.messages.GetDialogs(
        offset_date=0, offset_id=0, offset_peer=types.InputPeerEmpty(),
        limit=200, hash=0))
    for c in (getattr(r, "chats", []) or []):
        if isinstance(c, types.Channel) and c.title == title:
            peer_type = "supergroup" if getattr(c, "megagroup", False) else "channel"
            await app.storage.update_peers(
                [(int(f"-100{c.id}"), c.access_hash, peer_type, None, None)])
            return c
    raise RuntimeError(f"channel {title!r} not found in dialogs")


async def fetch_message(app, chat_id: int, msg_id: int):
    msg = await app.get_messages(chat_id=chat_id, message_ids=msg_id)
    if not msg or not msg.document:
        raise RuntimeError(f"message {msg_id} has no document")
    return msg


async def resumable_download(app, chat_id: int, msg_id: int, out: Path,
                             fresh: bool = False, max_refetch: int = 10):
    """Download with resume. Retries on file_reference expiry / transient errors."""
    from pyrogram.file_id import FileId

    total = None
    refetches = 0
    last = [time.time(), 0]

    def prog(cur, size):
        now = time.time()
        if now - last[0] > 15:
            spd = (cur - last[1]) / 2**20 / (now - last[0]) if last[1] else 0
            last[0], last[1] = now, cur
            pct = f" ({spd:.1f} MiB/s)" if spd else ""
            print(f"    {cur/2**20:.0f}/{size/2**20:.0f} MiB{pct}", flush=True)

    while True:
        msg = await fetch_message(app, chat_id, msg_id)
        doc_size = msg.document.file_size
        total = doc_size
        existing = out.stat().st_size if out.exists() and not fresh else 0
        resume_chunks = existing // CHUNK
        have = resume_chunks * CHUNK
        if have >= doc_size:
            print(f"[+] already complete: {out} ({doc_size/2**30:.2f} GiB)")
            return out

        fid = FileId.decode(msg.document.file_id)
        mode = "r+b" if have else "wb"
        with out.open(mode) as f:
            if have:
                f.truncate(have)
                f.seek(have)
                print(f"[*] resuming at {have/2**20:.0f} MiB "
                      f"(of {doc_size/2**20:.0f})")
            try:
                async for chunk in app.get_file(
                        fid, file_size=doc_size, offset=resume_chunks,
                        progress=prog):
                    f.write(chunk)
                print(f"[+] done: {out} ({out.stat().st_size/2**30:.2f} GiB)")
                return out
            except Exception as e:
                refetches += 1
                got = f.tell()
                print(f"    [!] {type(e).__name__}: {e} — at {got/2**20:.0f} MiB, "
                      f"refetch {refetches}/{max_refetch}", flush=True)
                if refetches >= max_refetch:
                    raise
                await asyncio.sleep(5)


async def main_async(args):
    load_dotenv(str(ENV_FILE))
    api_id = int(os.environ["TG_API_ID"])
    api_hash = os.environ["TG_API_HASH"]
    build_session_file(api_id, os.environ["TG_AUTH_KEY"].strip(),
                       int(os.environ["TG_DC_ID"]), int(os.environ["TG_USER_ID"]))
    from pyrogram import Client
    app = Client(name="breach_session", api_id=api_id, api_hash=api_hash,
                 workdir=str(SESSION_DIR))
    await app.start()
    try:
        chat = await resolve_channel(app, args.chat)
        chat_id = int(f"-100{chat.id}")
        print(f"[*] channel: {chat.title!r} id={chat_id}")
        if args.list:
            async for msg in app.get_chat_history(chat_id, limit=args.list):
                doc = ""
                if msg.document:
                    doc = (f" DOC[{msg.document.file_name} "
                           f"{msg.document.file_size/2**20:.0f}MiB]")
                text = (msg.text or msg.caption or "").replace("\n", " ")[:100]
                date = msg.date.strftime("%Y-%m-%d %H:%M") if msg.date else "?"
                print(f"[{msg.id}] {date}{doc} {text}")
            return
        msg = await fetch_message(app, chat_id, args.msg_id)
        print(f"[*] downloading: {msg.document.file_name} "
              f"({msg.document.file_size/2**30:.2f} GiB)")
        t0 = time.time()
        out = Path(args.out)
        out.parent.mkdir(parents=True, exist_ok=True)
        await resumable_download(app, chat_id, args.msg_id, out,
                                 fresh=args.fresh)
        dt = time.time() - t0
        print(f"[+] wall time: {dt:.0f}s")
    finally:
        await app.stop()


def main():
    ap = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    ap.add_argument("--chat", required=True, help="exact channel title (from dialogs)")
    ap.add_argument("--msg-id", type=int, help="message id with document")
    ap.add_argument("--out", help="output file path")
    ap.add_argument("--fresh", action="store_true",
                    help="ignore existing partial file, restart from zero")
    ap.add_argument("--list", type=int, metavar="N", help="list N recent messages, no download")
    args = ap.parse_args()
    if not args.list and not (args.msg_id and args.out):
        ap.error("--msg-id and --out required (or use --list N)")
    asyncio.run(main_async(args))


if __name__ == "__main__":
    main()
