Ztorrent/internal/torrent/engine.go
2026-03-06 14:09:24 +03:00

866 lines
19 KiB
Go

package torrent
import (
"context"
"crypto/rand"
"crypto/sha1"
"encoding/hex"
"errors"
"fmt"
"io"
"net"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
"sync"
"time"
"unicode"
"github.com/veggiedefender/torrent-client/internal/torrentfile"
"github.com/veggiedefender/torrent-client/internal/tracker"
)
type Engine struct {
mu sync.RWMutex
torrent *torrentfile.TorrentFile
peers []PeerStatus
trackers []TrackerStatus
peerID [20]byte
cancel context.CancelFunc
lastErr error
phase string
downloadedBytes int64
completedPieces int
totalPieces int
outputPath string
}
type Status struct {
Loaded bool
Name string
Length int
PieceLength int
PieceCount int
CompletedPieces int
DownloadedBytes int64
TotalBytes int64
OutputPath string
Files []torrentfile.File
Announce string
PeerCount int
Peers []PeerStatus
Trackers []TrackerStatus
PeerID string
Progress float64
Phase string
LastError string
}
type PeerStatus struct {
Address string
Port uint16
Source string
State string
Error string
DownloadedPieces int
}
type TrackerStatus struct {
URL string
State string
PeerCount int
Error string
}
const (
trackerPort = 6881
trackerNumWant = 200
trackerAnnounceTimeout = 6 * time.Second
downloadStallTimeout = 90 * time.Second
reannounceInterval = 25 * time.Second
idleRetryDelay = 2 * time.Second
)
func NewEngine() *Engine {
return &Engine{
peerID: generatePeerID(),
phase: "idle",
}
}
func (e *Engine) LoadTorrent(path string) error {
e.Stop()
e.resetStateForNewLoad()
e.setPhase("loading_metadata")
tf, err := torrentfile.Open(path)
if err != nil {
e.setError(err)
e.setPhase("failed")
return err
}
downloadCtx, cancel := context.WithCancel(context.Background())
e.swapCancel(cancel)
queryCtx, queryCancel := context.WithTimeout(downloadCtx, 30*time.Second)
defer queryCancel()
e.setPhase("querying_trackers")
peers, trackerStatuses := e.queryTrackers(queryCtx, tf, tracker.AnnounceOptions{
PeerID: e.peerID,
Port: trackerPort,
NumWant: trackerNumWant,
Timeout: trackerAnnounceTimeout,
})
e.mu.Lock()
e.torrent = tf
e.trackers = trackerStatuses
e.peers = peers
e.totalPieces = len(tf.PieceHashes)
e.downloadedBytes = 0
e.completedPieces = 0
e.outputPath = buildOutputPath(tf, "downloads")
if len(peers) == 0 {
warnErr := fmt.Errorf("no peers yet via %d trackers", len(trackerStatuses))
if trErr := firstTrackerError(trackerStatuses); trErr != nil {
warnErr = fmt.Errorf("no peers yet: %w", trErr)
}
e.lastErr = warnErr
} else {
e.lastErr = nil
}
e.phase = "ready"
e.mu.Unlock()
go e.runDownload(downloadCtx, tf)
return nil
}
func (e *Engine) Stop() {
e.mu.Lock()
cancel := e.cancel
e.cancel = nil
if cancel != nil {
e.phase = "stopped"
}
e.mu.Unlock()
if cancel != nil {
cancel()
}
}
func (e *Engine) Progress() float64 {
e.mu.RLock()
defer e.mu.RUnlock()
return e.progressUnsafe()
}
func (e *Engine) Status() Status {
e.mu.RLock()
defer e.mu.RUnlock()
status := Status{
Loaded: e.torrent != nil,
PeerCount: len(e.peers),
Peers: append([]PeerStatus(nil), e.peers...),
Trackers: append([]TrackerStatus(nil), e.trackers...),
PeerID: string(e.peerID[:]),
Progress: e.progressUnsafe(),
Phase: e.phase,
CompletedPieces: e.completedPieces,
DownloadedBytes: e.downloadedBytes,
OutputPath: e.outputPath,
}
if e.torrent != nil {
status.Name = e.torrent.Name
status.Length = e.torrent.Length
status.TotalBytes = int64(e.torrent.Length)
status.PieceLength = e.torrent.PieceLength
status.PieceCount = len(e.torrent.PieceHashes)
status.Files = append(status.Files, e.torrent.Files...)
status.Announce = e.torrent.Announce
}
if e.lastErr != nil {
status.LastError = e.lastErr.Error()
}
return status
}
func (e *Engine) queryTrackers(ctx context.Context, tf *torrentfile.TorrentFile, opts tracker.AnnounceOptions) ([]PeerStatus, []TrackerStatus) {
trackersToTry := collectTrackersToTry(tf)
peerMap := make(map[string]PeerStatus)
statuses := make([]TrackerStatus, 0, len(trackersToTry))
for _, announce := range trackersToTry {
perTrackerCtx, cancel := context.WithTimeout(ctx, opts.Timeout)
peers, err := tracker.GetPeersFromURL(perTrackerCtx, announce, tf, opts)
cancel()
st := TrackerStatus{URL: announce}
if err != nil {
st.State = "error"
st.Error = err.Error()
statuses = append(statuses, st)
continue
}
if len(peers) == 0 {
st.State = "empty"
} else {
st.State = "ok"
st.PeerCount = len(peers)
}
statuses = append(statuses, st)
for _, peer := range peers {
addPeer(peerMap, peer, announce)
}
}
peerList := make([]PeerStatus, 0, len(peerMap))
for _, peer := range peerMap {
peerList = append(peerList, peer)
}
sort.Slice(peerList, func(i, j int) bool {
if peerList[i].Address == peerList[j].Address {
return peerList[i].Port < peerList[j].Port
}
return peerList[i].Address < peerList[j].Address
})
return peerList, statuses
}
func (e *Engine) runDownload(ctx context.Context, tf *torrentfile.TorrentFile) {
defer e.clearCancel()
e.setPhase("preparing_download")
partPath, partFile, err := createPartFile(tf)
if err != nil {
e.setTerminalError("failed", fmt.Errorf("create temp file: %w", err))
return
}
defer partFile.Close()
scheduler := newPieceScheduler(len(tf.PieceHashes))
workerCtx, workerCancel := context.WithCancel(ctx)
var workersWG sync.WaitGroup
startedWorkers := make(map[string]struct{})
workerDoneCh := make(chan string, 512)
startWorkers := func(peers []PeerStatus) {
for _, peer := range peers {
key := peerStatusKey(peer.Address, peer.Port)
if _, exists := startedWorkers[key]; exists {
continue
}
startedWorkers[key] = struct{}{}
workersWG.Add(1)
go func(workerKey string, p PeerStatus) {
defer workersWG.Done()
e.runPeerWorker(workerCtx, tf, p, partFile, scheduler)
select {
case workerDoneCh <- workerKey:
default:
}
}(key, peer)
}
}
var cleanupOnce sync.Once
cleanupWorkers := func() {
cleanupOnce.Do(func() {
workerCancel()
scheduler.Stop()
workersWG.Wait()
})
}
defer cleanupWorkers()
e.setPhase("downloading")
lastProgressAt := time.Now()
lastReannounceAt := time.Now()
startWorkers(e.snapshotPeers())
heartbeatTicker := time.NewTicker(idleRetryDelay)
defer heartbeatTicker.Stop()
downloadCompleted := false
for !downloadCompleted {
select {
case <-ctx.Done():
e.setPhase("stopped")
e.setError(nil)
return
case _, ok := <-scheduler.Progress():
if !ok {
downloadCompleted = true
continue
}
lastProgressAt = time.Now()
case <-scheduler.Done():
downloadCompleted = true
case workerKey := <-workerDoneCh:
delete(startedWorkers, workerKey)
case <-heartbeatTicker.C:
if time.Since(lastReannounceAt) >= reannounceInterval {
e.setPhase("querying_trackers")
refreshCtx, cancel := context.WithTimeout(ctx, 20*time.Second)
_ = e.refreshPeersFromTrackers(refreshCtx, tf)
cancel()
lastReannounceAt = time.Now()
e.setPhase("downloading")
}
startWorkers(e.snapshotPeers())
if time.Since(lastProgressAt) >= downloadStallTimeout {
e.setTerminalError("stalled", fmt.Errorf("download stalled: no piece progress for %s", downloadStallTimeout.Round(time.Second)))
return
}
}
}
cleanupWorkers()
e.setPhase("writing_files")
if err := materializeDownloadedFiles(tf, partPath, "downloads"); err != nil {
e.setTerminalError("failed", fmt.Errorf("write output files: %w", err))
return
}
_ = os.Remove(partPath)
e.mu.Lock()
e.downloadedBytes = int64(tf.Length)
e.completedPieces = len(tf.PieceHashes)
e.phase = "completed"
e.lastErr = nil
e.mu.Unlock()
}
func (e *Engine) runPeerWorker(
ctx context.Context,
tf *torrentfile.TorrentFile,
peer PeerStatus,
partFile *os.File,
scheduler *pieceScheduler,
) {
e.setPeerState(peer.Address, peer.Port, "connecting", "")
peerAddr := net.JoinHostPort(peer.Address, strconv.Itoa(int(peer.Port)))
client, err := newPeerClient(ctx, peerAddr, tf.InfoHash, e.peerID, len(tf.PieceHashes))
if err != nil {
if ctx.Err() != nil {
e.setPeerState(peer.Address, peer.Port, "stopped", "")
} else {
e.setPeerState(peer.Address, peer.Port, "error", err.Error())
}
return
}
defer client.Close()
e.setPeerState(peer.Address, peer.Port, "ready", "")
consecutiveFailures := 0
for {
if ctx.Err() != nil {
e.setPeerState(peer.Address, peer.Port, "stopped", "")
return
}
have, hasInfo := client.PieceAvailability()
task, ok, err := scheduler.Acquire(ctx, have, hasInfo)
if err != nil {
if ctx.Err() != nil {
e.setPeerState(peer.Address, peer.Port, "stopped", "")
} else {
e.setPeerState(peer.Address, peer.Port, "error", err.Error())
}
return
}
if !ok {
select {
case <-ctx.Done():
e.setPeerState(peer.Address, peer.Port, "stopped", "")
return
case <-scheduler.Done():
e.setPeerState(peer.Address, peer.Port, "done", "")
return
case <-time.After(500 * time.Millisecond):
}
continue
}
e.setPeerState(peer.Address, peer.Port, "requesting", "")
pieceSize := pieceSizeForIndex(tf, task.Index)
pieceData, err := client.DownloadPiece(ctx, task.Index, pieceSize)
if err != nil {
_, _ = scheduler.Report(ctx, task.Index, false)
if ctx.Err() != nil {
e.setPeerState(peer.Address, peer.Port, "stopped", "")
return
}
consecutiveFailures++
e.setPeerState(peer.Address, peer.Port, "error", err.Error())
if shouldDropPeer(err) || consecutiveFailures >= 3 {
return
}
continue
}
hash := sha1.Sum(pieceData)
if hash != tf.PieceHashes[task.Index] {
_, _ = scheduler.Report(ctx, task.Index, false)
consecutiveFailures++
e.setPeerState(peer.Address, peer.Port, "error", fmt.Sprintf("piece %d hash mismatch", task.Index))
if consecutiveFailures >= 3 {
return
}
continue
}
offset := int64(task.Index * tf.PieceLength)
if _, err := partFile.WriteAt(pieceData, offset); err != nil {
_, _ = scheduler.Report(ctx, task.Index, false)
e.setPeerState(peer.Address, peer.Port, "error", fmt.Sprintf("write piece %d: %v", task.Index, err))
return
}
accepted, reportErr := scheduler.Report(ctx, task.Index, true)
if reportErr != nil {
if ctx.Err() != nil {
e.setPeerState(peer.Address, peer.Port, "stopped", "")
} else {
e.setPeerState(peer.Address, peer.Port, "error", reportErr.Error())
}
return
}
if accepted {
e.recordPieceComplete(peer.Address, peer.Port, len(pieceData))
}
consecutiveFailures = 0
e.setPeerState(peer.Address, peer.Port, "active", "")
}
}
func shouldDropPeer(err error) bool {
if err == nil {
return false
}
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
return true
}
if isTimeout(err) {
return false
}
msg := strings.ToLower(err.Error())
switch {
case strings.Contains(msg, "peer choked"):
return false
case strings.Contains(msg, "peer did not unchoke"):
return false
case strings.Contains(msg, "hash mismatch"):
return false
default:
return true
}
}
func (e *Engine) snapshotPeers() []PeerStatus {
e.mu.RLock()
defer e.mu.RUnlock()
return append([]PeerStatus(nil), e.peers...)
}
func (e *Engine) refreshPeersFromTrackers(ctx context.Context, tf *torrentfile.TorrentFile) int {
peers, trackerStatuses := e.queryTrackers(ctx, tf, tracker.AnnounceOptions{
PeerID: e.peerID,
Port: trackerPort,
Downloaded: e.currentDownloadedBytes(),
NumWant: trackerNumWant,
Timeout: trackerAnnounceTimeout,
})
e.mu.Lock()
defer e.mu.Unlock()
if len(trackerStatuses) > 0 {
e.trackers = trackerStatuses
}
if len(peers) == 0 {
waitErr := fmt.Errorf("waiting for peers via %d trackers", len(trackerStatuses))
if trErr := firstTrackerError(trackerStatuses); trErr != nil {
waitErr = fmt.Errorf("waiting for peers: %w", trErr)
}
e.lastErr = waitErr
return 0
}
e.lastErr = nil
seen := make(map[string]struct{}, len(e.peers))
for _, existingPeer := range e.peers {
seen[peerStatusKey(existingPeer.Address, existingPeer.Port)] = struct{}{}
}
added := 0
for _, peer := range peers {
key := peerStatusKey(peer.Address, peer.Port)
if _, ok := seen[key]; ok {
continue
}
e.peers = append(e.peers, peer)
seen[key] = struct{}{}
added++
}
if added > 0 {
sort.Slice(e.peers, func(i, j int) bool {
if e.peers[i].Address == e.peers[j].Address {
return e.peers[i].Port < e.peers[j].Port
}
return e.peers[i].Address < e.peers[j].Address
})
}
return added
}
func (e *Engine) currentDownloadedBytes() int64 {
e.mu.RLock()
defer e.mu.RUnlock()
return e.downloadedBytes
}
func (e *Engine) recordPieceComplete(address string, port uint16, pieceSize int) {
e.mu.Lock()
defer e.mu.Unlock()
e.downloadedBytes += int64(pieceSize)
e.completedPieces++
for i := range e.peers {
if e.peers[i].Address == address && e.peers[i].Port == port {
e.peers[i].DownloadedPieces++
break
}
}
}
func (e *Engine) setPeerState(address string, port uint16, state, errText string) {
e.mu.Lock()
defer e.mu.Unlock()
for i := range e.peers {
if e.peers[i].Address == address && e.peers[i].Port == port {
e.peers[i].State = state
e.peers[i].Error = errText
break
}
}
}
func (e *Engine) setTerminalError(phase string, err error) {
e.mu.Lock()
defer e.mu.Unlock()
e.phase = phase
e.lastErr = err
}
func (e *Engine) progressUnsafe() float64 {
if e.torrent != nil && e.torrent.Length > 0 {
if e.phase == "completed" {
return 1
}
if e.downloadedBytes > 0 {
progress := float64(e.downloadedBytes) / float64(e.torrent.Length)
if progress > 1 {
return 1
}
return progress
}
}
switch e.phase {
case "idle", "failed", "stopped", "stalled", "tracker_errors":
return 0
case "loading_metadata":
return 0.05
case "querying_trackers":
return 0.12
case "ready":
return 0.2
case "preparing_download":
return 0.25
case "downloading":
return 0.3
case "writing_files":
return 0.95
default:
return 0
}
}
func (e *Engine) resetStateForNewLoad() {
e.mu.Lock()
defer e.mu.Unlock()
e.torrent = nil
e.peers = nil
e.trackers = nil
e.lastErr = nil
e.phase = "idle"
e.downloadedBytes = 0
e.completedPieces = 0
e.totalPieces = 0
e.outputPath = ""
}
func (e *Engine) setError(err error) {
e.mu.Lock()
defer e.mu.Unlock()
e.lastErr = err
}
func (e *Engine) setPhase(phase string) {
e.mu.Lock()
defer e.mu.Unlock()
e.phase = phase
}
func (e *Engine) swapCancel(cancel context.CancelFunc) {
e.mu.Lock()
defer e.mu.Unlock()
if e.cancel != nil {
e.cancel()
}
e.cancel = cancel
}
func (e *Engine) clearCancel() {
e.mu.Lock()
defer e.mu.Unlock()
if e.cancel != nil {
e.cancel()
e.cancel = nil
}
}
func generatePeerID() [20]byte {
const prefix = "-ZT0001-"
const alphabet = "0123456789abcdefghijklmnopqrstuvwxyz"
var id [20]byte
copy(id[:], prefix)
buf := make([]byte, len(id)-len(prefix))
if _, err := rand.Read(buf); err != nil {
now := time.Now().UnixNano()
for i := range buf {
buf[i] = byte(now >> (i * 8))
}
}
for i, b := range buf {
id[len(prefix)+i] = alphabet[int(b)%len(alphabet)]
}
return id
}
func addPeer(peerMap map[string]PeerStatus, peer tracker.Peer, source string) {
key := peerStatusKey(peer.IP.String(), peer.Port)
if _, exists := peerMap[key]; exists {
return
}
peerMap[key] = PeerStatus{
Address: peer.IP.String(),
Port: peer.Port,
Source: source,
State: "discovered",
}
}
func peerStatusKey(address string, port uint16) string {
return net.JoinHostPort(address, strconv.Itoa(int(port)))
}
func collectTrackersToTry(tf *torrentfile.TorrentFile) []string {
const fallbackTrackers = `
udp://open.stealth.si:80/announce
udp://tracker.opentrackr.org:1337/announce
udp://tracker.openbittorrent.com:6969/announce
udp://tracker.torrent.eu.org:451/announce
https://tracker.opentrackr.org:443/announce
http://tracker.opentrackr.org:1337/announce
`
seen := make(map[string]struct{})
list := make([]string, 0, len(tf.Trackers)+8)
add := func(tr string) {
if tr == "" {
return
}
if _, ok := seen[tr]; ok {
return
}
seen[tr] = struct{}{}
list = append(list, tr)
}
for _, tr := range tf.Trackers {
add(tr)
}
for _, tr := range splitLines(fallbackTrackers) {
add(tr)
}
return list
}
func splitLines(s string) []string {
parts := strings.Split(s, "\n")
out := make([]string, 0, len(parts))
for _, part := range parts {
trimmed := strings.TrimSpace(part)
if trimmed == "" {
continue
}
out = append(out, trimmed)
}
return out
}
func firstTrackerError(statuses []TrackerStatus) error {
for _, st := range statuses {
if st.State == "error" && st.Error != "" {
return errors.New(st.Error)
}
}
return nil
}
func createPartFile(tf *torrentfile.TorrentFile) (string, *os.File, error) {
partsDir := filepath.Join("downloads", ".parts")
if err := os.MkdirAll(partsDir, 0o755); err != nil {
return "", nil, err
}
hashPrefix := hex.EncodeToString(tf.InfoHash[:6])
name := sanitizePathPart(tf.Name) + "-" + hashPrefix + ".part"
partPath := filepath.Join(partsDir, name)
f, err := os.OpenFile(partPath, os.O_CREATE|os.O_RDWR|os.O_TRUNC, 0o644)
if err != nil {
return "", nil, err
}
if err := f.Truncate(int64(tf.Length)); err != nil {
f.Close()
return "", nil, err
}
return partPath, f, nil
}
func materializeDownloadedFiles(tf *torrentfile.TorrentFile, partPath, outputRoot string) error {
partFile, err := os.Open(partPath)
if err != nil {
return err
}
defer partFile.Close()
if err := os.MkdirAll(outputRoot, 0o755); err != nil {
return err
}
offset := int64(0)
for _, f := range tf.Files {
targetPath, err := safeOutputPath(outputRoot, f.Path)
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(targetPath), 0o755); err != nil {
return err
}
out, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o644)
if err != nil {
return err
}
_, copyErr := io.CopyN(out, io.NewSectionReader(partFile, offset, int64(f.Length)), int64(f.Length))
closeErr := out.Close()
if copyErr != nil {
return copyErr
}
if closeErr != nil {
return closeErr
}
offset += int64(f.Length)
}
if offset != int64(tf.Length) {
return fmt.Errorf("written bytes mismatch: expected %d, wrote %d", tf.Length, offset)
}
return nil
}
func safeOutputPath(root, rel string) (string, error) {
clean := filepath.Clean(filepath.FromSlash(rel))
if clean == "." || clean == "" {
return "", errors.New("invalid output file path")
}
if filepath.IsAbs(clean) || strings.HasPrefix(clean, "..") {
return "", fmt.Errorf("unsafe output path %q", rel)
}
return filepath.Join(root, clean), nil
}
func buildOutputPath(tf *torrentfile.TorrentFile, outputRoot string) string {
if len(tf.Files) == 1 {
if path, err := safeOutputPath(outputRoot, tf.Files[0].Path); err == nil {
return path
}
}
return filepath.Join(outputRoot, sanitizePathPart(tf.Name))
}
func sanitizePathPart(s string) string {
s = strings.TrimSpace(s)
if s == "" {
return "torrent"
}
var b strings.Builder
for _, r := range s {
switch {
case unicode.IsLetter(r), unicode.IsDigit(r), r == '-', r == '_', r == '.':
b.WriteRune(r)
default:
b.WriteByte('_')
}
}
out := strings.Trim(b.String(), "._")
if out == "" {
return "torrent"
}
return out
}