514 lines
20 KiB
Python
514 lines
20 KiB
Python
import os
|
|
import sys
|
|
import json
|
|
import time
|
|
import math
|
|
import asyncio
|
|
import inspect
|
|
import logging
|
|
from telethon import TelegramClient, errors, utils
|
|
from telethon.tl.functions.auth import ExportAuthorizationRequest, ImportAuthorizationRequest
|
|
from telethon.tl.functions import InvokeWithLayerRequest
|
|
from telethon.tl.types import MessageMediaPhoto, MessageMediaDocument
|
|
from telethon.network import MTProtoSender
|
|
from telethon.tl.alltlobjects import LAYER
|
|
from FastTelethonhelper.FastTelethon import DownloadSender
|
|
|
|
# Suppress Telethon's internal network warnings (e.g., transient server-closed socket resets that auto-reconnect)
|
|
logging.basicConfig(level=logging.ERROR)
|
|
logging.getLogger('telethon').setLevel(logging.ERROR)
|
|
logging.getLogger('asyncio').setLevel(logging.ERROR)
|
|
|
|
# Global cache for persistent downloaders per DC ID to prevent connection churn
|
|
downloaders = {}
|
|
downloader_lock = asyncio.Lock()
|
|
|
|
# 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 -1
|
|
|
|
async def update_slot(self, slot_idx, filename, received, total, start_time):
|
|
if not total:
|
|
total = 1
|
|
percent = (received / total) * 100
|
|
bar_len = 15
|
|
filled_len = int(bar_len * received // total)
|
|
bar = '█' * filled_len + '░' * (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:]
|
|
|
|
async with self.lock:
|
|
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:
|
|
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")
|
|
sys.stdout.write("\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 = '█' * filled_len + ' ' * (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):
|
|
if sys.platform == 'win32':
|
|
import ctypes
|
|
kernel32 = ctypes.windll.kernel32
|
|
kernel32.SetConsoleMode(kernel32.GetStdHandle(-11), 7)
|
|
|
|
while True:
|
|
await renderer.render()
|
|
await asyncio.sleep(0.2)
|
|
|
|
# Persistent parallel connection downloader to avoid connection churn
|
|
class PersistentParallelDownloader:
|
|
def __init__(self, client, dc_id, connection_count):
|
|
self.client = client
|
|
self.dc_id = dc_id
|
|
self.connection_count = connection_count
|
|
self.senders = []
|
|
self.auth_key = None
|
|
|
|
async def initialize(self):
|
|
self.auth_key = (
|
|
None
|
|
if self.dc_id and self.client.session.dc_id != self.dc_id
|
|
else self.client.session.auth_key
|
|
)
|
|
|
|
# Connect persistent MTProtoSenders sequentially to prevent session ID conflicts
|
|
for i in range(self.connection_count):
|
|
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 i > 0 and not sender.auth_key:
|
|
sender.auth_key = self.auth_key
|
|
|
|
self.senders.append(sender)
|
|
|
|
async def download_file(self, input_file_location, size, out, progress_callback=None):
|
|
part_size_kb = utils.get_appropriated_part_size(size)
|
|
part_size = part_size_kb * 1024
|
|
part_count = math.ceil(size / part_size)
|
|
|
|
connections = self.connection_count
|
|
minimum, remainder = divmod(part_count, connections)
|
|
|
|
def get_part_count():
|
|
nonlocal remainder
|
|
if remainder > 0:
|
|
remainder -= 1
|
|
return minimum + 1
|
|
return minimum
|
|
|
|
download_senders = []
|
|
for i in range(connections):
|
|
ds = DownloadSender(
|
|
self.client,
|
|
self.senders[i],
|
|
input_file_location, # Pass the correct InputFileLocation subclass, not the Document TLObject
|
|
offset=i * part_size,
|
|
limit=part_size,
|
|
stride=connections * part_size,
|
|
count=get_part_count()
|
|
)
|
|
download_senders.append(ds)
|
|
|
|
part = 0
|
|
while part < part_count:
|
|
tasks = []
|
|
for ds in download_senders:
|
|
tasks.append(self.client.loop.create_task(ds.next()))
|
|
try:
|
|
for task in tasks:
|
|
data = await task
|
|
if not data:
|
|
break
|
|
out.write(data)
|
|
part += 1
|
|
if progress_callback:
|
|
r = progress_callback(out.tell(), size)
|
|
if inspect.isawaitable(r):
|
|
await r
|
|
except Exception:
|
|
for task in tasks:
|
|
if not task.done():
|
|
task.cancel()
|
|
await asyncio.gather(*tasks, return_exceptions=True)
|
|
raise
|
|
|
|
async def disconnect_all(self):
|
|
await asyncio.gather(*[sender.disconnect() for sender in self.senders if sender], return_exceptions=True)
|
|
self.senders = []
|
|
|
|
async def get_downloader(client, dc_id, connection_count):
|
|
async with downloader_lock:
|
|
if dc_id in downloaders:
|
|
dl = downloaders[dc_id]
|
|
# Ensure existing downloaders have all active senders matching connection count
|
|
all_alive = len(dl.senders) == connection_count and all(s.is_connected() for s in dl.senders)
|
|
if not all_alive:
|
|
await dl.disconnect_all()
|
|
del downloaders[dc_id]
|
|
if dc_id not in downloaders:
|
|
dl = PersistentParallelDownloader(client, dc_id, connection_count)
|
|
await dl.initialize()
|
|
downloaders[dc_id] = dl
|
|
return downloaders[dc_id]
|
|
|
|
async def download_media_fast(client, chat, message, filepath, callback, connection_count=8):
|
|
if message.document:
|
|
dc_id, input_file_location = utils.get_input_location(message.document)
|
|
dl = await get_downloader(client, dc_id, connection_count)
|
|
with open(filepath, "wb") as f:
|
|
await dl.download_file(input_file_location, message.document.size, f, progress_callback=callback)
|
|
else:
|
|
# Photos are small enough that native download works perfectly
|
|
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=8):
|
|
msg_id_str = str(message.id)
|
|
slot_idx = await renderer.get_free_slot()
|
|
|
|
async with semaphore:
|
|
start_time = time.time()
|
|
|
|
def callback(received, total):
|
|
asyncio.create_task(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
|
|
try:
|
|
if message.document:
|
|
dc_id, _ = utils.get_input_location(message.document)
|
|
async with downloader_lock:
|
|
dl = downloaders.pop(dc_id, None)
|
|
if dl:
|
|
await dl.disconnect_all()
|
|
except Exception:
|
|
pass
|
|
if retries >= max_retries:
|
|
break
|
|
try:
|
|
if os.path.exists(filepath):
|
|
os.truncate(filepath, 0)
|
|
except Exception:
|
|
pass
|
|
await asyncio.sleep(retries * 3)
|
|
|
|
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 1
|
|
connections_per_file = int(sys.argv[7]) if len(sys.argv) > 7 else 8
|
|
|
|
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}")
|
|
|
|
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 (common for deleted accounts or raw user IDs), iterate dialogs to find entity and access hash
|
|
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)
|
|
|
|
if os.path.exists(filepath):
|
|
local_size = os.path.getsize(filepath)
|
|
|
|
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("\nStarting downloads...")
|
|
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 downloaders
|
|
for dl in downloaders.values():
|
|
await dl.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())
|