diff --git a/dist/ztorrent-darwin-amd64 b/dist/ztorrent-darwin-amd64 index 4c1770a..de1d9d5 100755 Binary files a/dist/ztorrent-darwin-amd64 and b/dist/ztorrent-darwin-amd64 differ diff --git a/dist/ztorrent-darwin-arm64 b/dist/ztorrent-darwin-arm64 index 140a88a..0e43d70 100755 Binary files a/dist/ztorrent-darwin-arm64 and b/dist/ztorrent-darwin-arm64 differ diff --git a/internal/torrent/engine.go b/internal/torrent/engine.go index 4229d4b..200ee14 100644 --- a/internal/torrent/engine.go +++ b/internal/torrent/engine.go @@ -3,25 +3,41 @@ package torrent import ( "context" "crypto/rand" + "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 + mu sync.RWMutex + torrent *torrentfile.TorrentFile peers []PeerStatus trackers []TrackerStatus peerID [20]byte - cancel context.CancelFunc - lastErr error - phase string + + cancel context.CancelFunc + + lastErr error + phase string + + downloadedBytes int64 + completedPieces int + totalPieces int + outputPath string } type Status struct { @@ -30,21 +46,32 @@ type Status struct { Length int PieceLength int PieceCount int - Files []torrentfile.File - Announce string - PeerCount int - Peers []PeerStatus - Trackers []TrackerStatus - PeerID string - Progress float64 - Phase string - LastError string + + 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 + Address string + Port uint16 + Source string + State string + Error string + DownloadedPieces int } type TrackerStatus struct { @@ -54,6 +81,15 @@ type TrackerStatus struct { 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(), @@ -62,7 +98,10 @@ func NewEngine() *Engine { } 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) @@ -70,77 +109,62 @@ func (e *Engine) LoadTorrent(path string) error { return err } - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + downloadCtx, cancel := context.WithCancel(context.Background()) e.swapCancel(cancel) - defer e.clearCancel() + + queryCtx, queryCancel := context.WithTimeout(downloadCtx, 30*time.Second) + defer queryCancel() e.setPhase("querying_trackers") - peers, trackerStatuses := e.queryTrackers(ctx, tf, tracker.AnnounceOptions{ + peers, trackerStatuses := e.queryTrackers(queryCtx, tf, tracker.AnnounceOptions{ PeerID: e.peerID, - Port: 6881, - NumWant: 200, - Timeout: 6 * time.Second, + Port: trackerPort, + NumWant: trackerNumWant, + Timeout: trackerAnnounceTimeout, }) - var finalErr error - if len(peers) == 0 { - if err := firstTrackerError(trackerStatuses); err != nil { - finalErr = fmt.Errorf("no peers found: %w", err) - } else { - finalErr = fmt.Errorf("no peers found via %d trackers", len(trackerStatuses)) - } - } - e.mu.Lock() e.torrent = tf - e.peers = peers e.trackers = trackerStatuses - e.lastErr = finalErr - if len(peers) > 0 { - e.phase = "ready" + 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.phase = "waiting_peers" + e.lastErr = nil } + e.phase = "ready" e.mu.Unlock() - if finalErr != nil { - e.mu.Lock() - if e.phase == "waiting_peers" { - e.phase = "tracker_errors" - } - e.mu.Unlock() - return finalErr - } - + go e.runDownload(downloadCtx, tf) return nil } func (e *Engine) Stop() { e.mu.Lock() - defer e.mu.Unlock() - if e.cancel != nil { - e.cancel() - e.cancel = nil + 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() - switch e.phase { - case "idle": - return 0 - case "loading_metadata": - return 0.1 - case "querying_trackers": - return 0.25 - case "ready": - return 0.4 - case "waiting_peers", "tracker_errors": - return 0.3 - default: - return 0 - } + return e.progressUnsafe() } func (e *Engine) Status() Status { @@ -148,17 +172,22 @@ func (e *Engine) Status() Status { defer e.mu.RUnlock() status := Status{ - Loaded: e.torrent != nil, - PeerCount: len(e.peers), - PeerID: string(e.peerID[:]), - Progress: e.progressUnsafe(), - Phase: e.phase, - Peers: append([]PeerStatus(nil), e.peers...), - Trackers: append([]TrackerStatus(nil), e.trackers...), + 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...) @@ -172,18 +201,16 @@ func (e *Engine) Status() Status { } func (e *Engine) queryTrackers(ctx context.Context, tf *torrentfile.TorrentFile, opts tracker.AnnounceOptions) ([]PeerStatus, []TrackerStatus) { - trackers := collectTrackersToTry(tf) + trackersToTry := collectTrackersToTry(tf) peerMap := make(map[string]PeerStatus) - statuses := make([]TrackerStatus, 0, len(trackers)) + statuses := make([]TrackerStatus, 0, len(trackersToTry)) - for _, announce := range trackers { + for _, announce := range trackersToTry { perTrackerCtx, cancel := context.WithTimeout(ctx, opts.Timeout) peers, err := tracker.GetPeersFromURL(perTrackerCtx, announce, tf, opts) cancel() - st := TrackerStatus{ - URL: announce, - } + st := TrackerStatus{URL: announce} if err != nil { st.State = "error" st.Error = err.Error() @@ -208,112 +235,267 @@ func (e *Engine) queryTrackers(ctx context.Context, tf *torrentfile.TorrentFile, 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 addPeer(peerMap map[string]PeerStatus, peer tracker.Peer, source string) { - key := net.JoinHostPort(peer.IP.String(), fmt.Sprintf("%d", peer.Port)) - if _, exists := peerMap[key]; exists { +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 } - peerMap[key] = PeerStatus{ - Address: peer.IP.String(), - Port: peer.Port, - Source: source, - } -} + defer partFile.Close() -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 -` + pieceDone := make([]bool, len(tf.PieceHashes)) + e.setPhase("downloading") + lastProgressAt := time.Now() + lastReannounceAt := time.Now() - seen := make(map[string]struct{}) - list := make([]string, 0, len(tf.Trackers)+8) - - add := func(tr string) { - if tr == "" { + for !allPiecesDone(pieceDone) { + if err := ctx.Err(); err != nil { + e.setPhase("stopped") + e.setError(nil) 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 { - lines := make([]string, 0) - current := make([]byte, 0, len(s)) - flush := func() { - if len(current) == 0 { - return + peers := e.snapshotPeers() + if len(peers) == 0 { + e.setPhase("querying_trackers") + refreshCtx, cancel := context.WithTimeout(ctx, 20*time.Second) + _ = e.refreshPeersFromTrackers(refreshCtx, tf) + cancel() + lastReannounceAt = time.Now() + e.setPhase("downloading") + peers = e.snapshotPeers() } - lines = append(lines, string(current)) - current = current[:0] - } - for i := 0; i < len(s); i++ { - if s[i] == '\n' { - flush() - continue - } - if s[i] == '\r' { - continue - } - if s[i] == ' ' || s[i] == '\t' { - if len(current) == 0 { + madeProgress := false + for _, peer := range peers { + if allPiecesDone(pieceDone) { + break + } + if err := ctx.Err(); err != nil { + e.setPhase("stopped") + e.setError(nil) + return + } + + e.setPeerState(peer.Address, peer.Port, "connecting", "") + pieces, _, err := downloadFromPeer(ctx, tf, peer, e.peerID, partFile, pieceDone, func(pieceIndex int, pieceSize int) { + e.recordPieceComplete(peer.Address, peer.Port, pieceSize) + }) + + if err != nil { + e.setPeerState(peer.Address, peer.Port, "error", err.Error()) continue } + + if pieces > 0 { + madeProgress = true + e.setPeerState(peer.Address, peer.Port, "active", "") + } else { + e.setPeerState(peer.Address, peer.Port, "idle", "") + } + } + + if madeProgress { + lastProgressAt = time.Now() + continue + } + + 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") + } + + if time.Since(lastProgressAt) >= downloadStallTimeout { + e.setTerminalError("stalled", fmt.Errorf("download stalled: no piece progress for %s", downloadStallTimeout.Round(time.Second))) + return + } + + if err := sleepWithContext(ctx, idleRetryDelay); err != nil { + e.setPhase("stopped") + e.setError(nil) + return } - current = append(current, s[i]) } - flush() - return lines + + 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 firstTrackerError(statuses []TrackerStatus) error { - for _, st := range statuses { - if st.State == "error" && st.Error != "" { - return errors.New(st.Error) +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 } } - return nil +} + +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": + case "idle", "failed", "stopped", "stalled", "tracker_errors": return 0 case "loading_metadata": - return 0.1 + return 0.05 case "querying_trackers": - return 0.25 + return 0.12 case "ready": - return 0.4 - case "waiting_peers", "tracker_errors": + 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() @@ -365,3 +547,207 @@ func generatePeerID() [20]byte { 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 allPiecesDone(pieceDone []bool) bool { + for _, done := range pieceDone { + if !done { + return false + } + } + return true +} + +func sleepWithContext(ctx context.Context, d time.Duration) error { + timer := time.NewTimer(d) + defer timer.Stop() + + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + 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 +} diff --git a/internal/torrent/peerwire.go b/internal/torrent/peerwire.go new file mode 100644 index 0000000..ae2ab7b --- /dev/null +++ b/internal/torrent/peerwire.go @@ -0,0 +1,440 @@ +package torrent + +import ( + "bytes" + "context" + "crypto/sha1" + "encoding/binary" + "errors" + "fmt" + "io" + "net" + "os" + "strconv" + "time" + + "github.com/veggiedefender/torrent-client/internal/torrentfile" +) + +const ( + wireProtocolString = "BitTorrent protocol" + + msgChoke = 0 + msgUnchoke = 1 + msgInterested = 2 + msgNotInterested = 3 + msgHave = 4 + msgBitfield = 5 + msgRequest = 6 + msgPiece = 7 + + maxWireMessageSize = 2 * 1024 * 1024 + requestBlockSize = 16 * 1024 + + peerConnectTimeout = 8 * time.Second + peerReadTimeout = 15 * time.Second + peerWriteTimeout = 10 * time.Second +) + +type wireMessage struct { + ID int + Payload []byte +} + +type peerClient struct { + conn net.Conn + + have []bool + hasPieceInfo bool + peerIsChoked bool +} + +func downloadFromPeer( + ctx context.Context, + tf *torrentfile.TorrentFile, + peer PeerStatus, + peerID [20]byte, + partFile *os.File, + pieceDone []bool, + onPiece func(pieceIndex int, pieceSize int), +) (int, int, error) { + peerAddr := net.JoinHostPort(peer.Address, strconv.Itoa(int(peer.Port))) + client, err := newPeerClient(ctx, peerAddr, tf.InfoHash, peerID, len(tf.PieceHashes)) + if err != nil { + return 0, 0, err + } + defer client.Close() + + piecesDownloaded := 0 + bytesDownloaded := 0 + + for pieceIndex := 0; pieceIndex < len(tf.PieceHashes); pieceIndex++ { + if ctx.Err() != nil { + return piecesDownloaded, bytesDownloaded, ctx.Err() + } + + if pieceDone[pieceIndex] { + continue + } + if !client.HasPiece(pieceIndex) { + continue + } + + size := pieceSizeForIndex(tf, pieceIndex) + pieceData, err := client.DownloadPiece(ctx, pieceIndex, size) + if err != nil { + if piecesDownloaded > 0 { + return piecesDownloaded, bytesDownloaded, nil + } + return 0, 0, err + } + + sum := sha1.Sum(pieceData) + if sum != tf.PieceHashes[pieceIndex] { + if piecesDownloaded > 0 { + return piecesDownloaded, bytesDownloaded, nil + } + return 0, 0, fmt.Errorf("piece %d hash mismatch", pieceIndex) + } + + offset := int64(pieceIndex * tf.PieceLength) + if _, err := partFile.WriteAt(pieceData, offset); err != nil { + return piecesDownloaded, bytesDownloaded, err + } + + pieceDone[pieceIndex] = true + piecesDownloaded++ + bytesDownloaded += len(pieceData) + if onPiece != nil { + onPiece(pieceIndex, len(pieceData)) + } + } + + return piecesDownloaded, bytesDownloaded, nil +} + +func newPeerClient(ctx context.Context, addr string, infoHash [20]byte, peerID [20]byte, pieceCount int) (*peerClient, error) { + dialer := net.Dialer{Timeout: peerConnectTimeout} + conn, err := dialer.DialContext(ctx, "tcp", addr) + if err != nil { + return nil, err + } + + pc := &peerClient{ + conn: conn, + have: make([]bool, pieceCount), + peerIsChoked: true, + } + + if err := pc.sendHandshake(ctx, infoHash, peerID); err != nil { + conn.Close() + return nil, err + } + if err := pc.readHandshake(ctx, infoHash); err != nil { + conn.Close() + return nil, err + } + + if err := pc.sendMessage(ctx, msgInterested, nil); err != nil { + conn.Close() + return nil, err + } + + _ = pc.readInitialMessages(ctx) + return pc, nil +} + +func (pc *peerClient) Close() error { + return pc.conn.Close() +} + +func (pc *peerClient) HasPiece(index int) bool { + if index < 0 || index >= len(pc.have) { + return false + } + if !pc.hasPieceInfo { + return true + } + return pc.have[index] +} + +func (pc *peerClient) DownloadPiece(ctx context.Context, pieceIndex int, pieceLength int) ([]byte, error) { + if pieceLength <= 0 { + return nil, fmt.Errorf("invalid piece length %d", pieceLength) + } + + if err := pc.waitForUnchoke(ctx); err != nil { + return nil, err + } + + piece := make([]byte, pieceLength) + for offset := 0; offset < pieceLength; { + blockLength := requestBlockSize + if remaining := pieceLength - offset; remaining < blockLength { + blockLength = remaining + } + + if err := pc.sendRequest(ctx, pieceIndex, offset, blockLength); err != nil { + return nil, err + } + + block, err := pc.readPieceBlock(ctx, pieceIndex, offset, blockLength) + if err != nil { + return nil, err + } + copy(piece[offset:], block) + offset += len(block) + } + + return piece, nil +} + +func (pc *peerClient) readInitialMessages(ctx context.Context) error { + if err := pc.setReadDeadlineFromContext(ctx, 1200*time.Millisecond); err != nil { + return err + } + defer pc.conn.SetReadDeadline(time.Time{}) + + for { + msg, err := readWireMessage(pc.conn) + if err != nil { + if isTimeout(err) { + return nil + } + return err + } + pc.consumeMessage(msg) + } +} + +func (pc *peerClient) waitForUnchoke(ctx context.Context) error { + if !pc.peerIsChoked { + return nil + } + + deadline := time.Now().Add(12 * time.Second) + for { + if err := ctx.Err(); err != nil { + return err + } + if time.Now().After(deadline) { + return errors.New("peer did not unchoke") + } + + if err := pc.setReadDeadlineFromContext(ctx, peerReadTimeout); err != nil { + return err + } + msg, err := readWireMessage(pc.conn) + if err != nil { + if isTimeout(err) { + continue + } + return err + } + pc.consumeMessage(msg) + if msg.ID == msgUnchoke { + return nil + } + } +} + +func (pc *peerClient) sendRequest(ctx context.Context, pieceIndex, begin, length int) error { + payload := make([]byte, 12) + binary.BigEndian.PutUint32(payload[0:4], uint32(pieceIndex)) + binary.BigEndian.PutUint32(payload[4:8], uint32(begin)) + binary.BigEndian.PutUint32(payload[8:12], uint32(length)) + return pc.sendMessage(ctx, msgRequest, payload) +} + +func (pc *peerClient) readPieceBlock(ctx context.Context, pieceIndex, begin, expectedLen int) ([]byte, error) { + for { + if err := ctx.Err(); err != nil { + return nil, err + } + if err := pc.setReadDeadlineFromContext(ctx, peerReadTimeout); err != nil { + return nil, err + } + + msg, err := readWireMessage(pc.conn) + if err != nil { + if isTimeout(err) { + continue + } + return nil, err + } + + switch msg.ID { + case msgPiece: + if len(msg.Payload) < 8 { + continue + } + gotIndex := int(binary.BigEndian.Uint32(msg.Payload[0:4])) + gotBegin := int(binary.BigEndian.Uint32(msg.Payload[4:8])) + block := msg.Payload[8:] + if gotIndex != pieceIndex || gotBegin != begin { + pc.consumeMessage(msg) + continue + } + if len(block) == 0 { + continue + } + if len(block) > expectedLen { + block = block[:expectedLen] + } + return block, nil + case msgChoke: + pc.consumeMessage(msg) + return nil, errors.New("peer choked") + default: + pc.consumeMessage(msg) + } + } +} + +func (pc *peerClient) consumeMessage(msg wireMessage) { + switch msg.ID { + case msgChoke: + pc.peerIsChoked = true + case msgUnchoke: + pc.peerIsChoked = false + case msgHave: + if len(msg.Payload) < 4 { + return + } + idx := int(binary.BigEndian.Uint32(msg.Payload[:4])) + if idx >= 0 && idx < len(pc.have) { + pc.have[idx] = true + pc.hasPieceInfo = true + } + case msgBitfield: + pc.hasPieceInfo = true + for i := range pc.have { + pc.have[i] = bitfieldHasPiece(msg.Payload, i) + } + } +} + +func (pc *peerClient) sendHandshake(ctx context.Context, infoHash [20]byte, peerID [20]byte) error { + payload := make([]byte, 49+len(wireProtocolString)) + payload[0] = byte(len(wireProtocolString)) + copy(payload[1:1+len(wireProtocolString)], wireProtocolString) + copy(payload[1+len(wireProtocolString)+8:1+len(wireProtocolString)+8+20], infoHash[:]) + copy(payload[1+len(wireProtocolString)+8+20:], peerID[:]) + + if err := pc.setWriteDeadlineFromContext(ctx, peerWriteTimeout); err != nil { + return err + } + _, err := pc.conn.Write(payload) + return err +} + +func (pc *peerClient) readHandshake(ctx context.Context, expectedInfoHash [20]byte) error { + head := make([]byte, 1) + if err := pc.setReadDeadlineFromContext(ctx, peerReadTimeout); err != nil { + return err + } + if _, err := io.ReadFull(pc.conn, head); err != nil { + return err + } + + pstrlen := int(head[0]) + if pstrlen <= 0 || pstrlen > 64 { + return fmt.Errorf("invalid handshake pstrlen %d", pstrlen) + } + + rest := make([]byte, pstrlen+48) + if _, err := io.ReadFull(pc.conn, rest); err != nil { + return err + } + + if string(rest[:pstrlen]) != wireProtocolString { + return errors.New("invalid peer protocol string") + } + + infoHashOffset := pstrlen + 8 + if !bytes.Equal(rest[infoHashOffset:infoHashOffset+20], expectedInfoHash[:]) { + return errors.New("peer info_hash mismatch") + } + + return nil +} + +func (pc *peerClient) sendMessage(ctx context.Context, msgID int, payload []byte) error { + if err := pc.setWriteDeadlineFromContext(ctx, peerWriteTimeout); err != nil { + return err + } + + length := uint32(1 + len(payload)) + buf := make([]byte, 4+length) + binary.BigEndian.PutUint32(buf[0:4], length) + buf[4] = byte(msgID) + copy(buf[5:], payload) + + _, err := pc.conn.Write(buf) + return err +} + +func readWireMessage(r io.Reader) (wireMessage, error) { + var lengthBuf [4]byte + if _, err := io.ReadFull(r, lengthBuf[:]); err != nil { + return wireMessage{}, err + } + + length := binary.BigEndian.Uint32(lengthBuf[:]) + if length == 0 { + return wireMessage{ID: -1}, nil + } + if length > maxWireMessageSize { + return wireMessage{}, fmt.Errorf("wire message too large: %d", length) + } + + msg := make([]byte, length) + if _, err := io.ReadFull(r, msg); err != nil { + return wireMessage{}, err + } + return wireMessage{ID: int(msg[0]), Payload: msg[1:]}, nil +} + +func bitfieldHasPiece(bitfield []byte, index int) bool { + byteIndex := index / 8 + if byteIndex < 0 || byteIndex >= len(bitfield) { + return false + } + bitOffset := 7 - (index % 8) + return bitfield[byteIndex]&(1<= len(tf.PieceHashes) { + return 0 + } + if pieceIndex == len(tf.PieceHashes)-1 { + used := tf.PieceLength * (len(tf.PieceHashes) - 1) + return tf.Length - used + } + return tf.PieceLength +} + +func (pc *peerClient) setReadDeadlineFromContext(ctx context.Context, fallback time.Duration) error { + deadline := time.Now().Add(fallback) + if ctxDeadline, ok := ctx.Deadline(); ok && ctxDeadline.Before(deadline) { + deadline = ctxDeadline + } + return pc.conn.SetReadDeadline(deadline) +} + +func (pc *peerClient) setWriteDeadlineFromContext(ctx context.Context, fallback time.Duration) error { + deadline := time.Now().Add(fallback) + if ctxDeadline, ok := ctx.Deadline(); ok && ctxDeadline.Before(deadline) { + deadline = ctxDeadline + } + return pc.conn.SetWriteDeadline(deadline) +} + +func isTimeout(err error) bool { + if err == nil { + return false + } + var netErr net.Error + return errors.As(err, &netErr) && netErr.Timeout() +} diff --git a/ui/window.go b/ui/window.go index d95735e..5366979 100644 --- a/ui/window.go +++ b/ui/window.go @@ -36,6 +36,8 @@ type desktopUI struct { sizeLabel *widget.Label pieceLabel *widget.Label pieceCountLabel *widget.Label + downloadedLabel *widget.Label + outputLabel *widget.Label peerCountLabel *widget.Label mu sync.RWMutex @@ -75,6 +77,8 @@ func newDesktopUI(controller *appcore.Controller) *desktopUI { sizeLabel: widget.NewLabel("-"), pieceLabel: widget.NewLabel("-"), pieceCountLabel: widget.NewLabel("-"), + downloadedLabel: widget.NewLabel("0 / 0"), + outputLabel: widget.NewLabel("-"), peerCountLabel: widget.NewLabel("0"), } @@ -84,11 +88,13 @@ func newDesktopUI(controller *appcore.Controller) *desktopUI { loadButton := widget.NewButton("Load Torrent", ui.onLoad) topRow := container.NewBorder(nil, nil, nil, container.NewHBox(browseButton, loadButton), ui.pathEntry) - metaGrid := container.NewGridWithColumns(3, + metaGrid := container.NewGridWithColumns(4, wrapMetric("Name", ui.nameLabel), wrapMetric("Size (bytes)", ui.sizeLabel), wrapMetric("Piece length", ui.pieceLabel), - wrapMetric("Pieces", ui.pieceCountLabel), + wrapMetric("Completed pieces", ui.pieceCountLabel), + wrapMetric("Downloaded", ui.downloadedLabel), + wrapMetric("Output", ui.outputLabel), wrapMetric("Discovered peers", ui.peerCountLabel), wrapMetric("Phase", ui.phaseMetricLabel), ) @@ -204,7 +210,13 @@ func (ui *desktopUI) updateFromStatus(status torrent.Status) { ui.nameLabel.SetText(status.Name) ui.sizeLabel.SetText(strconv.Itoa(status.Length)) ui.pieceLabel.SetText(strconv.Itoa(status.PieceLength)) - ui.pieceCountLabel.SetText(strconv.Itoa(status.PieceCount)) + if status.PieceCount > 0 { + ui.pieceCountLabel.SetText(fmt.Sprintf("%d / %d", status.CompletedPieces, status.PieceCount)) + } else { + ui.pieceCountLabel.SetText("-") + } + ui.downloadedLabel.SetText(fmt.Sprintf("%d / %d", status.DownloadedBytes, status.TotalBytes)) + ui.outputLabel.SetText(status.OutputPath) ui.peerCountLabel.SetText(strconv.Itoa(status.PeerCount)) ui.mu.Lock() @@ -219,6 +231,8 @@ func (ui *desktopUI) updateFromStatus(status torrent.Status) { if status.LastError != "" { ui.errorLabel.SetText(status.LastError) + } else { + ui.errorLabel.SetText("") } }) } @@ -292,7 +306,7 @@ func (ui *desktopUI) newPeersTable() *widget.Table { func() (int, int) { ui.mu.RLock() defer ui.mu.RUnlock() - return len(ui.peers) + 1, 3 + return len(ui.peers) + 1, 6 }, func() fyne.CanvasObject { return widget.NewLabel("") @@ -306,7 +320,13 @@ func (ui *desktopUI) newPeersTable() *widget.Table { case 1: label.SetText("Port") case 2: + label.SetText("State") + case 3: + label.SetText("Pieces") + case 4: label.SetText("Source tracker") + case 5: + label.SetText("Error") } return } @@ -320,13 +340,22 @@ func (ui *desktopUI) newPeersTable() *widget.Table { case 1: label.SetText(strconv.Itoa(int(peer.Port))) case 2: + label.SetText(peer.State) + case 3: + label.SetText(strconv.Itoa(peer.DownloadedPieces)) + case 4: label.SetText(peer.Source) + case 5: + label.SetText(peer.Error) } }, ) - table.SetColumnWidth(0, 180) + table.SetColumnWidth(0, 170) table.SetColumnWidth(1, 80) - table.SetColumnWidth(2, 700) + table.SetColumnWidth(2, 100) + table.SetColumnWidth(3, 70) + table.SetColumnWidth(4, 310) + table.SetColumnWidth(5, 440) return table }