603 lines
23 KiB
Python
603 lines
23 KiB
Python
import os
|
|
import sys
|
|
import json
|
|
import time
|
|
import math
|
|
import asyncio
|
|
import inspect
|
|
import logging
|
|
|
|
# Ensure UTF-8 output and enable Windows VT100 / ANSI escape sequences
|
|
if sys.platform == 'win32':
|
|
os.system('')
|
|
try:
|
|
sys.stdout.reconfigure(encoding='utf-8', errors='replace')
|
|
sys.stderr.reconfigure(encoding='utf-8', errors='replace')
|
|
except Exception:
|
|
pass
|
|
|
|
from telethon import TelegramClient, errors, utils
|
|
from telethon.tl.functions.auth import ExportAuthorizationRequest, ImportAuthorizationRequest
|
|
from telethon.tl.functions import InvokeWithLayerRequest
|
|
from telethon.tl.functions.upload import GetFileRequest
|
|
from telethon.tl.types import MessageMediaPhoto, MessageMediaDocument
|
|
from telethon.network import MTProtoSender
|
|
from telethon.tl.alltlobjects import LAYER
|
|
|
|
# Suppress Telethon's internal network warnings
|
|
logging.basicConfig(level=logging.ERROR)
|
|
logging.getLogger('telethon').setLevel(logging.ERROR)
|
|
logging.getLogger('asyncio').setLevel(logging.ERROR)
|
|
|
|
# Maximum chunk size supported by Telegram MTProto upload.getFile is 512 KB
|
|
CHUNK_SIZE = 512 * 1024
|
|
|
|
# Helper function to format sizes
|
|
def format_size(bytes_count):
|
|
if bytes_count < 1024:
|
|
return f"{bytes_count} B"
|
|
elif bytes_count < 1024 * 1024:
|
|
return f"{bytes_count / 1024:.2f} KB"
|
|
elif bytes_count < 1024 * 1024 * 1024:
|
|
return f"{bytes_count / (1024 * 1024):.2f} MB"
|
|
else:
|
|
return f"{bytes_count / (1024 * 1024 * 1024):.2f} GB"
|
|
|
|
# Shared statistics and UI rendering
|
|
class ConcurrentProgressRenderer:
|
|
def __init__(self, total_files, concurrency):
|
|
self.total_files = total_files
|
|
self.concurrency = concurrency
|
|
self.completed_files = 0
|
|
self.downloaded_files = 0
|
|
self.skipped_files = 0
|
|
self.skipped_bytes = 0
|
|
self.total_downloaded_bytes = 0
|
|
self.start_time = time.time()
|
|
self.slots = [""] * concurrency
|
|
self.lock = asyncio.Lock()
|
|
self.initialized = False
|
|
|
|
async def get_free_slot(self):
|
|
async with self.lock:
|
|
for i in range(self.concurrency):
|
|
if self.slots[i] == "":
|
|
self.slots[i] = "Initializing..."
|
|
return i
|
|
return 0
|
|
|
|
def update_slot(self, slot_idx, filename, received, total, start_time):
|
|
if not (0 <= slot_idx < len(self.slots)):
|
|
return
|
|
if not total:
|
|
total = 1
|
|
percent = (received / total) * 100
|
|
bar_len = 15
|
|
filled_len = int(bar_len * received // total) if total > 0 else 0
|
|
bar = '█' * min(filled_len, bar_len) + '░' * max(0, bar_len - filled_len)
|
|
|
|
now = time.time()
|
|
elapsed = now - start_time
|
|
speed = received / elapsed if elapsed > 0 else 0
|
|
speed_str = f"{format_size(speed)}/s"
|
|
size_str = f"{format_size(received)}/{format_size(total)}"
|
|
|
|
display_name = filename
|
|
if len(display_name) > 20:
|
|
display_name = display_name[:9] + "..." + display_name[-8:]
|
|
|
|
self.slots[slot_idx] = f"Slot {slot_idx+1}: {display_name} [{bar}] {percent:.1f}% ({size_str}) @ {speed_str}"
|
|
|
|
async def release_slot(self, slot_idx):
|
|
async with self.lock:
|
|
if 0 <= slot_idx < len(self.slots):
|
|
self.slots[slot_idx] = ""
|
|
|
|
async def add_bytes(self, size):
|
|
async with self.lock:
|
|
self.total_downloaded_bytes += size
|
|
self.downloaded_files += 1
|
|
self.completed_files += 1
|
|
|
|
async def increment_skipped(self, size=0):
|
|
async with self.lock:
|
|
self.skipped_files += 1
|
|
self.skipped_bytes += size
|
|
self.completed_files += 1
|
|
|
|
async def log(self, message):
|
|
async with self.lock:
|
|
if self.initialized:
|
|
num_lines = 2 + len(self.slots)
|
|
sys.stdout.write(f"\033[{num_lines}A\033[J")
|
|
sys.stdout.write(message + "\n")
|
|
sys.stdout.flush()
|
|
self.initialized = False
|
|
|
|
async def render(self):
|
|
async with self.lock:
|
|
elapsed = time.time() - self.start_time
|
|
overall_speed = self.total_downloaded_bytes / elapsed if elapsed > 0 else 0
|
|
|
|
percent = (self.completed_files / self.total_files * 100) if self.total_files > 0 else 0.0
|
|
bar_len = 50
|
|
filled_len = int(bar_len * self.completed_files // self.total_files) if self.total_files > 0 else 0
|
|
bar = '█' * min(filled_len, bar_len) + ' ' * max(0, bar_len - filled_len)
|
|
|
|
lines = []
|
|
lines.append(
|
|
f"Overall Progress: {self.completed_files}/{self.total_files} |{bar}| {percent:.2f}%"
|
|
)
|
|
lines.append(
|
|
f"Skipped: {self.skipped_files} files equaling {format_size(self.skipped_bytes)} | "
|
|
f"Downloaded: {self.downloaded_files} files equaling {format_size(self.total_downloaded_bytes)} (Avg: {format_size(overall_speed)}/s)"
|
|
)
|
|
|
|
for slot_str in self.slots:
|
|
lines.append(slot_str.ljust(100))
|
|
|
|
if self.initialized:
|
|
sys.stdout.write(f"\033[{len(lines)}A")
|
|
else:
|
|
self.initialized = True
|
|
|
|
sys.stdout.write("\n".join(lines) + "\n")
|
|
sys.stdout.flush()
|
|
|
|
async def ui_loop(renderer):
|
|
while True:
|
|
try:
|
|
await renderer.render()
|
|
except asyncio.CancelledError:
|
|
break
|
|
except Exception:
|
|
pass
|
|
await asyncio.sleep(0.25)
|
|
|
|
# Thread-safe persistent MTProto sender pool per Telegram Data Center (DC)
|
|
class SenderPool:
|
|
def __init__(self, client, dc_id, max_senders=20):
|
|
self.client = client
|
|
self.dc_id = dc_id
|
|
self.max_senders = max_senders
|
|
self.auth_key = (
|
|
None
|
|
if dc_id and client.session.dc_id != dc_id
|
|
else client.session.auth_key
|
|
)
|
|
self.available = []
|
|
self.total_created = 0
|
|
self.condition = asyncio.Condition()
|
|
|
|
async def _create_sender(self):
|
|
dc = await self.client._get_dc(self.dc_id)
|
|
sender = MTProtoSender(self.auth_key, loggers=self.client._log, retries=10, delay=1, auto_reconnect=True)
|
|
await sender.connect(
|
|
self.client._connection(
|
|
dc.ip_address,
|
|
dc.port,
|
|
dc.id,
|
|
loggers=self.client._log,
|
|
proxy=self.client._proxy,
|
|
)
|
|
)
|
|
if not self.auth_key:
|
|
auth = await self.client(ExportAuthorizationRequest(self.dc_id))
|
|
self.client._init_request.query = ImportAuthorizationRequest(
|
|
id=auth.id, bytes=auth.bytes
|
|
)
|
|
req = InvokeWithLayerRequest(LAYER, self.client._init_request)
|
|
await sender.send(req)
|
|
self.auth_key = sender.auth_key
|
|
elif not sender.auth_key:
|
|
sender.auth_key = self.auth_key
|
|
return sender
|
|
|
|
async def acquire_senders(self, count):
|
|
async with self.condition:
|
|
senders = []
|
|
while len(senders) < count:
|
|
# 1. Drain alive senders from available pool
|
|
while self.available and len(senders) < count:
|
|
s = self.available.pop()
|
|
if s.is_connected():
|
|
senders.append(s)
|
|
else:
|
|
try:
|
|
await s.disconnect()
|
|
except Exception:
|
|
pass
|
|
self.total_created = max(0, self.total_created - 1)
|
|
|
|
if len(senders) == count:
|
|
break
|
|
|
|
# 2. Create new connections up to max_senders
|
|
if self.total_created < self.max_senders:
|
|
needed = min(count - len(senders), self.max_senders - self.total_created)
|
|
for _ in range(needed):
|
|
try:
|
|
s = await self._create_sender()
|
|
self.total_created += 1
|
|
senders.append(s)
|
|
except Exception:
|
|
break
|
|
|
|
# If we have at least 1 sender, proceed without blocking
|
|
if senders:
|
|
break
|
|
|
|
# If 0 senders available, wait for another task to release
|
|
await self.condition.wait()
|
|
|
|
return senders
|
|
|
|
async def release_senders(self, senders):
|
|
async with self.condition:
|
|
for s in senders:
|
|
if s and s.is_connected():
|
|
self.available.append(s)
|
|
else:
|
|
if s:
|
|
try:
|
|
await s.disconnect()
|
|
except Exception:
|
|
pass
|
|
self.total_created = max(0, self.total_created - 1)
|
|
self.condition.notify_all()
|
|
|
|
async def disconnect_all(self):
|
|
async with self.condition:
|
|
all_senders = list(self.available)
|
|
self.available.clear()
|
|
for s in all_senders:
|
|
try:
|
|
await s.disconnect()
|
|
except Exception:
|
|
pass
|
|
self.total_created = 0
|
|
|
|
dc_pools = {}
|
|
pool_lock = asyncio.Lock()
|
|
|
|
async def get_dc_pool(client, dc_id):
|
|
async with pool_lock:
|
|
if dc_id not in dc_pools:
|
|
dc_pools[dc_id] = SenderPool(client, dc_id)
|
|
return dc_pools[dc_id]
|
|
|
|
# High-throughput asynchronous pipelined downloader with sequential streaming disk writes
|
|
async def download_file_parallel(pool, client, input_file_location, size, filepath, callback, desired_connections=4):
|
|
part_count = math.ceil(size / CHUNK_SIZE) if size > 0 else 0
|
|
if part_count == 0:
|
|
with open(filepath, "wb") as f:
|
|
pass
|
|
return
|
|
|
|
num_senders = min(desired_connections, part_count)
|
|
senders = await pool.acquire_senders(num_senders)
|
|
|
|
try:
|
|
queue = asyncio.Queue()
|
|
for i in range(part_count):
|
|
queue.put_nowait((i, i * CHUNK_SIZE))
|
|
|
|
current_part = 0
|
|
buffer = {}
|
|
bytes_written = 0
|
|
write_lock = asyncio.Lock()
|
|
stop_event = asyncio.Event()
|
|
|
|
# Open file in write-binary mode directly (no slow pre-allocation or seek stalls)
|
|
with open(filepath, "wb") as f:
|
|
async def handle_chunk(part_idx, data):
|
|
nonlocal current_part, bytes_written
|
|
async with write_lock:
|
|
buffer[part_idx] = data
|
|
while current_part in buffer:
|
|
chunk = buffer.pop(current_part)
|
|
f.write(chunk)
|
|
bytes_written += len(chunk)
|
|
current_part += 1
|
|
if callback:
|
|
callback(bytes_written, size)
|
|
|
|
async def worker(sender):
|
|
while not queue.empty() and not stop_event.is_set():
|
|
try:
|
|
part_idx, offset = queue.get_nowait()
|
|
except asyncio.QueueEmpty:
|
|
break
|
|
|
|
data = None
|
|
last_err = None
|
|
for attempt in range(4):
|
|
if stop_event.is_set():
|
|
break
|
|
try:
|
|
req = GetFileRequest(input_file_location, offset=offset, limit=CHUNK_SIZE)
|
|
res = await client._call(sender, req)
|
|
data = res.bytes
|
|
break
|
|
except errors.FloodWaitError as e:
|
|
await asyncio.sleep(e.seconds)
|
|
except errors.FileReferenceExpiredError:
|
|
stop_event.set()
|
|
raise
|
|
except Exception as e:
|
|
last_err = e
|
|
await asyncio.sleep(0.5 * (attempt + 1))
|
|
|
|
if data is None and not stop_event.is_set():
|
|
stop_event.set()
|
|
if last_err:
|
|
raise last_err
|
|
raise IOError(f"Failed to fetch chunk at offset {offset}")
|
|
|
|
if data:
|
|
await handle_chunk(part_idx, data)
|
|
|
|
queue.task_done()
|
|
|
|
await asyncio.gather(*[worker(s) for s in senders])
|
|
f.flush()
|
|
finally:
|
|
await pool.release_senders(senders)
|
|
|
|
async def download_media_fast(client, chat, message, filepath, callback, connection_count=4):
|
|
if message.document:
|
|
dc_id, input_file_location = utils.get_input_location(message.document)
|
|
pool = await get_dc_pool(client, dc_id)
|
|
await download_file_parallel(pool, client, input_file_location, message.document.size, filepath, callback, desired_connections=connection_count)
|
|
else:
|
|
# Photos are small enough that native download works smoothly
|
|
await client.download_media(message, file=filepath, progress_callback=callback)
|
|
|
|
async def download_file_task(client, chat, message, filepath, filename, file_size, index,
|
|
renderer, semaphore, history_file, downloaded_history, history_lock, connection_count=4):
|
|
msg_id_str = str(message.id)
|
|
|
|
async with semaphore:
|
|
slot_idx = await renderer.get_free_slot()
|
|
start_time = time.time()
|
|
# Immediately display slot info and starting 0% progress bar
|
|
renderer.update_slot(slot_idx, filename, 0, file_size, start_time)
|
|
|
|
def callback(received, total):
|
|
renderer.update_slot(slot_idx, filename, received, total or file_size, start_time)
|
|
|
|
success = False
|
|
retries = 0
|
|
max_retries = 5
|
|
|
|
while retries < max_retries:
|
|
try:
|
|
await download_media_fast(client, chat, message, filepath, callback, connection_count=connection_count)
|
|
success = True
|
|
break
|
|
except errors.FileReferenceExpiredError:
|
|
retries += 1
|
|
try:
|
|
refreshed = await client.get_messages(chat, ids=message.id)
|
|
if refreshed and refreshed.media:
|
|
message = refreshed
|
|
except Exception:
|
|
pass
|
|
if retries >= max_retries:
|
|
break
|
|
try:
|
|
if os.path.exists(filepath):
|
|
os.truncate(filepath, 0)
|
|
except Exception:
|
|
pass
|
|
await asyncio.sleep(1)
|
|
except errors.FloodWaitError as e:
|
|
await asyncio.sleep(e.seconds)
|
|
except (errors.RPCError, asyncio.TimeoutError, ConnectionError, Exception):
|
|
retries += 1
|
|
if retries >= max_retries:
|
|
break
|
|
try:
|
|
if os.path.exists(filepath):
|
|
os.truncate(filepath, 0)
|
|
except Exception:
|
|
pass
|
|
await asyncio.sleep(retries * 2)
|
|
|
|
actual_size = os.path.getsize(filepath) if success and os.path.exists(filepath) else 0
|
|
await renderer.release_slot(slot_idx)
|
|
|
|
if success:
|
|
await renderer.add_bytes(actual_size)
|
|
await renderer.log(f"[FINISHED] '{filename}' downloaded successfully ({format_size(actual_size)}).")
|
|
async with history_lock:
|
|
downloaded_history[msg_id_str] = {
|
|
"filename": filename,
|
|
"size": actual_size,
|
|
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S")
|
|
}
|
|
try:
|
|
with open(history_file, 'w', encoding='utf-8') as f:
|
|
json.dump(downloaded_history, f, indent=4, ensure_ascii=False)
|
|
except Exception:
|
|
pass
|
|
else:
|
|
await renderer.add_bytes(0)
|
|
|
|
async def main():
|
|
if len(sys.argv) < 5:
|
|
print("Usage: python download_telegram.py <api_id> <api_hash> <phone> <chat_username> [output_dir] [concurrency] [connections_per_file]")
|
|
sys.exit(1)
|
|
|
|
api_id = int(sys.argv[1])
|
|
api_hash = sys.argv[2]
|
|
phone = sys.argv[3]
|
|
try:
|
|
chat_username = int(sys.argv[4])
|
|
except ValueError:
|
|
chat_username = sys.argv[4]
|
|
output_dir = sys.argv[5] if len(sys.argv) > 5 else "telegram_downloads"
|
|
concurrency = int(sys.argv[6]) if len(sys.argv) > 6 else 3
|
|
connections_per_file = int(sys.argv[7]) if len(sys.argv) > 7 else 4
|
|
|
|
os.makedirs(output_dir, exist_ok=True)
|
|
history_file = os.path.join(output_dir, "download_history.json")
|
|
|
|
downloaded_history = {}
|
|
if os.path.exists(history_file):
|
|
try:
|
|
with open(history_file, 'r', encoding='utf-8') as f:
|
|
downloaded_history = json.load(f)
|
|
except Exception as e:
|
|
print(f"Warning: Failed to load download history: {e}")
|
|
|
|
# Fast directory scan to cache local files in memory
|
|
existing_files = {}
|
|
try:
|
|
with os.scandir(output_dir) as it:
|
|
for entry in it:
|
|
if entry.is_file():
|
|
existing_files[entry.name] = entry.stat().st_size
|
|
except Exception as e:
|
|
print(f"Warning scanning output directory: {e}")
|
|
|
|
client = TelegramClient('session_dumdum', api_id, api_hash)
|
|
|
|
await client.start(phone=phone)
|
|
print("LOGGED_IN")
|
|
|
|
chat = None
|
|
try:
|
|
chat = await client.get_entity(chat_username)
|
|
except Exception as e:
|
|
if isinstance(chat_username, int) and chat_username < 0 and not str(chat_username).startswith("-100"):
|
|
try:
|
|
alternative_id = int(f"-100{abs(chat_username)}")
|
|
print(f"Failed to resolve {chat_username}. Retrying with channel ID format {alternative_id}...")
|
|
chat = await client.get_entity(alternative_id)
|
|
except Exception:
|
|
pass
|
|
|
|
# If entity not found in session cache, search dialogs
|
|
if not chat:
|
|
print(f"Direct get_entity failed ({e}). Searching dialogs for ID {chat_username}...")
|
|
target_id = chat_username if isinstance(chat_username, int) else None
|
|
if target_id is None:
|
|
try:
|
|
target_id = int(str(chat_username).strip())
|
|
except ValueError:
|
|
target_id = None
|
|
|
|
async for dialog in client.iter_dialogs():
|
|
if target_id is not None and dialog.id == target_id:
|
|
chat = dialog.input_entity
|
|
break
|
|
elif getattr(dialog.entity, 'username', None) and dialog.entity.username.lower() == str(chat_username).lower().lstrip('@'):
|
|
chat = dialog.input_entity
|
|
break
|
|
elif target_id is not None and getattr(dialog.entity, 'id', None) == target_id:
|
|
chat = dialog.input_entity
|
|
break
|
|
|
|
if not chat:
|
|
print(f"Error getting chat: {e}")
|
|
await client.disconnect()
|
|
sys.exit(1)
|
|
|
|
print(f"Connected to chat: {chat_username}")
|
|
print("Scanning chat history to count files. Please wait...")
|
|
|
|
media_messages = []
|
|
async for message in client.iter_messages(chat):
|
|
if message.media:
|
|
media_messages.append(message)
|
|
|
|
media_messages.reverse()
|
|
total_files = len(media_messages)
|
|
print(f"Found {total_files} media files in total.")
|
|
|
|
renderer = ConcurrentProgressRenderer(total_files, concurrency)
|
|
semaphore = asyncio.Semaphore(concurrency)
|
|
history_lock = asyncio.Lock()
|
|
|
|
tasks_to_run = []
|
|
assigned_filenames = {}
|
|
|
|
for message in media_messages:
|
|
msg_id_str = str(message.id)
|
|
filename = None
|
|
file_size = 0
|
|
|
|
if isinstance(message.media, MessageMediaPhoto):
|
|
filename = f"photo_{message.id}.jpg"
|
|
if hasattr(message.media, 'photo') and message.media.photo:
|
|
file_size = getattr(message.media.photo, 'sizes', [None])[-1]
|
|
file_size = getattr(file_size, 'size', 0) if file_size else 0
|
|
else:
|
|
if message.file:
|
|
filename = message.file.name
|
|
file_size = message.file.size
|
|
if not filename:
|
|
ext = message.file.ext if message.file and message.file.ext else '.bin'
|
|
filename = f"file_{message.id}{ext}"
|
|
|
|
if filename in assigned_filenames and assigned_filenames[filename] != msg_id_str:
|
|
base, extension = os.path.splitext(filename)
|
|
filename = f"{base}_{message.id}{extension}"
|
|
|
|
assigned_filenames[filename] = msg_id_str
|
|
filepath = os.path.join(output_dir, filename)
|
|
|
|
local_size = existing_files.get(filename)
|
|
if local_size is not None:
|
|
if file_size > 0 and local_size == file_size:
|
|
print(f"[SKIP] '{filename}' already exists and is complete ({format_size(local_size)}).")
|
|
if msg_id_str not in downloaded_history:
|
|
downloaded_history[msg_id_str] = {
|
|
"filename": filename,
|
|
"size": local_size,
|
|
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S")
|
|
}
|
|
await renderer.increment_skipped(local_size)
|
|
continue
|
|
elif file_size > 0 and local_size != file_size:
|
|
print(f"[OVERWRITE] '{filename}' is incomplete (local: {format_size(local_size)}, expected: {format_size(file_size)}). Overwriting...")
|
|
else:
|
|
print(f"[OVERWRITE] '{filename}' size unknown or conflict. Overwriting...")
|
|
|
|
task = download_file_task(
|
|
client, chat, message, filepath, filename, file_size, len(tasks_to_run) + 1,
|
|
renderer, semaphore, history_file, downloaded_history, history_lock,
|
|
connection_count=connections_per_file
|
|
)
|
|
tasks_to_run.append(task)
|
|
|
|
print(f"\nStarting downloads with {concurrency} concurrent files and {connections_per_file} streams per file...")
|
|
ui_task = asyncio.create_task(ui_loop(renderer))
|
|
|
|
if tasks_to_run:
|
|
await asyncio.gather(*tasks_to_run)
|
|
|
|
ui_task.cancel()
|
|
await renderer.render()
|
|
|
|
# Close persistent DC pools
|
|
for pool in dc_pools.values():
|
|
await pool.disconnect_all()
|
|
|
|
try:
|
|
with open(history_file, 'w', encoding='utf-8') as f:
|
|
json.dump(downloaded_history, f, indent=4, ensure_ascii=False)
|
|
except Exception:
|
|
pass
|
|
|
|
print(f"\n\nFINISHED.")
|
|
print(f"Total media files: {total_files}")
|
|
print(f"Newly downloaded/retried: {len(tasks_to_run)}")
|
|
print(f"Skipped: {renderer.skipped_files}")
|
|
print(f"Output folder: {os.path.abspath(output_dir)}")
|
|
await client.disconnect()
|
|
|
|
if __name__ == '__main__':
|
|
asyncio.run(main())
|