This commit is contained in:
itexpert228 2026-03-06 13:48:38 +03:00
parent 27bb592659
commit ebba1d5b61
5 changed files with 1008 additions and 153 deletions

Binary file not shown.

Binary file not shown.

View file

@ -3,25 +3,41 @@ package torrent
import ( import (
"context" "context"
"crypto/rand" "crypto/rand"
"encoding/hex"
"errors" "errors"
"fmt" "fmt"
"io"
"net" "net"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
"sync" "sync"
"time" "time"
"unicode"
"github.com/veggiedefender/torrent-client/internal/torrentfile" "github.com/veggiedefender/torrent-client/internal/torrentfile"
"github.com/veggiedefender/torrent-client/internal/tracker" "github.com/veggiedefender/torrent-client/internal/tracker"
) )
type Engine struct { type Engine struct {
mu sync.RWMutex mu sync.RWMutex
torrent *torrentfile.TorrentFile torrent *torrentfile.TorrentFile
peers []PeerStatus peers []PeerStatus
trackers []TrackerStatus trackers []TrackerStatus
peerID [20]byte peerID [20]byte
cancel context.CancelFunc
lastErr error cancel context.CancelFunc
phase string
lastErr error
phase string
downloadedBytes int64
completedPieces int
totalPieces int
outputPath string
} }
type Status struct { type Status struct {
@ -30,21 +46,32 @@ type Status struct {
Length int Length int
PieceLength int PieceLength int
PieceCount int PieceCount int
Files []torrentfile.File
Announce string CompletedPieces int
PeerCount int DownloadedBytes int64
Peers []PeerStatus TotalBytes int64
Trackers []TrackerStatus OutputPath string
PeerID string
Progress float64 Files []torrentfile.File
Phase string Announce string
LastError string
PeerCount int
Peers []PeerStatus
Trackers []TrackerStatus
PeerID string
Progress float64
Phase string
LastError string
} }
type PeerStatus struct { type PeerStatus struct {
Address string Address string
Port uint16 Port uint16
Source string Source string
State string
Error string
DownloadedPieces int
} }
type TrackerStatus struct { type TrackerStatus struct {
@ -54,6 +81,15 @@ type TrackerStatus struct {
Error string 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 { func NewEngine() *Engine {
return &Engine{ return &Engine{
peerID: generatePeerID(), peerID: generatePeerID(),
@ -62,7 +98,10 @@ func NewEngine() *Engine {
} }
func (e *Engine) LoadTorrent(path string) error { func (e *Engine) LoadTorrent(path string) error {
e.Stop()
e.resetStateForNewLoad()
e.setPhase("loading_metadata") e.setPhase("loading_metadata")
tf, err := torrentfile.Open(path) tf, err := torrentfile.Open(path)
if err != nil { if err != nil {
e.setError(err) e.setError(err)
@ -70,77 +109,62 @@ func (e *Engine) LoadTorrent(path string) error {
return err return err
} }
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) downloadCtx, cancel := context.WithCancel(context.Background())
e.swapCancel(cancel) e.swapCancel(cancel)
defer e.clearCancel()
queryCtx, queryCancel := context.WithTimeout(downloadCtx, 30*time.Second)
defer queryCancel()
e.setPhase("querying_trackers") e.setPhase("querying_trackers")
peers, trackerStatuses := e.queryTrackers(ctx, tf, tracker.AnnounceOptions{ peers, trackerStatuses := e.queryTrackers(queryCtx, tf, tracker.AnnounceOptions{
PeerID: e.peerID, PeerID: e.peerID,
Port: 6881, Port: trackerPort,
NumWant: 200, NumWant: trackerNumWant,
Timeout: 6 * time.Second, 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.mu.Lock()
e.torrent = tf e.torrent = tf
e.peers = peers
e.trackers = trackerStatuses e.trackers = trackerStatuses
e.lastErr = finalErr e.peers = peers
if len(peers) > 0 { e.totalPieces = len(tf.PieceHashes)
e.phase = "ready" 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 { } else {
e.phase = "waiting_peers" e.lastErr = nil
} }
e.phase = "ready"
e.mu.Unlock() e.mu.Unlock()
if finalErr != nil { go e.runDownload(downloadCtx, tf)
e.mu.Lock()
if e.phase == "waiting_peers" {
e.phase = "tracker_errors"
}
e.mu.Unlock()
return finalErr
}
return nil return nil
} }
func (e *Engine) Stop() { func (e *Engine) Stop() {
e.mu.Lock() e.mu.Lock()
defer e.mu.Unlock() cancel := e.cancel
if e.cancel != nil { e.cancel = nil
e.cancel() if cancel != nil {
e.cancel = nil e.phase = "stopped"
}
e.mu.Unlock()
if cancel != nil {
cancel()
} }
} }
func (e *Engine) Progress() float64 { func (e *Engine) Progress() float64 {
e.mu.RLock() e.mu.RLock()
defer e.mu.RUnlock() defer e.mu.RUnlock()
switch e.phase { return e.progressUnsafe()
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
}
} }
func (e *Engine) Status() Status { func (e *Engine) Status() Status {
@ -148,17 +172,22 @@ func (e *Engine) Status() Status {
defer e.mu.RUnlock() defer e.mu.RUnlock()
status := Status{ status := Status{
Loaded: e.torrent != nil, Loaded: e.torrent != nil,
PeerCount: len(e.peers), PeerCount: len(e.peers),
PeerID: string(e.peerID[:]), Peers: append([]PeerStatus(nil), e.peers...),
Progress: e.progressUnsafe(), Trackers: append([]TrackerStatus(nil), e.trackers...),
Phase: e.phase, PeerID: string(e.peerID[:]),
Peers: append([]PeerStatus(nil), e.peers...), Progress: e.progressUnsafe(),
Trackers: append([]TrackerStatus(nil), e.trackers...), Phase: e.phase,
CompletedPieces: e.completedPieces,
DownloadedBytes: e.downloadedBytes,
OutputPath: e.outputPath,
} }
if e.torrent != nil { if e.torrent != nil {
status.Name = e.torrent.Name status.Name = e.torrent.Name
status.Length = e.torrent.Length status.Length = e.torrent.Length
status.TotalBytes = int64(e.torrent.Length)
status.PieceLength = e.torrent.PieceLength status.PieceLength = e.torrent.PieceLength
status.PieceCount = len(e.torrent.PieceHashes) status.PieceCount = len(e.torrent.PieceHashes)
status.Files = append(status.Files, e.torrent.Files...) 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) { 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) 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) perTrackerCtx, cancel := context.WithTimeout(ctx, opts.Timeout)
peers, err := tracker.GetPeersFromURL(perTrackerCtx, announce, tf, opts) peers, err := tracker.GetPeersFromURL(perTrackerCtx, announce, tf, opts)
cancel() cancel()
st := TrackerStatus{ st := TrackerStatus{URL: announce}
URL: announce,
}
if err != nil { if err != nil {
st.State = "error" st.State = "error"
st.Error = err.Error() st.Error = err.Error()
@ -208,112 +235,267 @@ func (e *Engine) queryTrackers(ctx context.Context, tf *torrentfile.TorrentFile,
for _, peer := range peerMap { for _, peer := range peerMap {
peerList = append(peerList, peer) 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 return peerList, statuses
} }
func addPeer(peerMap map[string]PeerStatus, peer tracker.Peer, source string) { func (e *Engine) runDownload(ctx context.Context, tf *torrentfile.TorrentFile) {
key := net.JoinHostPort(peer.IP.String(), fmt.Sprintf("%d", peer.Port)) defer e.clearCancel()
if _, exists := peerMap[key]; exists {
e.setPhase("preparing_download")
partPath, partFile, err := createPartFile(tf)
if err != nil {
e.setTerminalError("failed", fmt.Errorf("create temp file: %w", err))
return return
} }
peerMap[key] = PeerStatus{ defer partFile.Close()
Address: peer.IP.String(),
Port: peer.Port,
Source: source,
}
}
func collectTrackersToTry(tf *torrentfile.TorrentFile) []string { pieceDone := make([]bool, len(tf.PieceHashes))
const fallbackTrackers = ` e.setPhase("downloading")
udp://open.stealth.si:80/announce lastProgressAt := time.Now()
udp://tracker.opentrackr.org:1337/announce lastReannounceAt := time.Now()
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{}) for !allPiecesDone(pieceDone) {
list := make([]string, 0, len(tf.Trackers)+8) if err := ctx.Err(); err != nil {
e.setPhase("stopped")
add := func(tr string) { e.setError(nil)
if tr == "" {
return return
} }
if _, ok := seen[tr]; ok {
return
}
seen[tr] = struct{}{}
list = append(list, tr)
}
for _, tr := range tf.Trackers { peers := e.snapshotPeers()
add(tr) if len(peers) == 0 {
} e.setPhase("querying_trackers")
for _, tr := range splitLines(fallbackTrackers) { refreshCtx, cancel := context.WithTimeout(ctx, 20*time.Second)
add(tr) _ = e.refreshPeersFromTrackers(refreshCtx, tf)
} cancel()
lastReannounceAt = time.Now()
return list e.setPhase("downloading")
} peers = e.snapshotPeers()
func splitLines(s string) []string {
lines := make([]string, 0)
current := make([]byte, 0, len(s))
flush := func() {
if len(current) == 0 {
return
} }
lines = append(lines, string(current))
current = current[:0]
}
for i := 0; i < len(s); i++ { madeProgress := false
if s[i] == '\n' { for _, peer := range peers {
flush() if allPiecesDone(pieceDone) {
continue break
} }
if s[i] == '\r' { if err := ctx.Err(); err != nil {
continue e.setPhase("stopped")
} e.setError(nil)
if s[i] == ' ' || s[i] == '\t' { return
if len(current) == 0 { }
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 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 { func (e *Engine) snapshotPeers() []PeerStatus {
for _, st := range statuses { e.mu.RLock()
if st.State == "error" && st.Error != "" { defer e.mu.RUnlock()
return errors.New(st.Error) 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 { 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 { switch e.phase {
case "idle": case "idle", "failed", "stopped", "stalled", "tracker_errors":
return 0 return 0
case "loading_metadata": case "loading_metadata":
return 0.1 return 0.05
case "querying_trackers": case "querying_trackers":
return 0.25 return 0.12
case "ready": case "ready":
return 0.4 return 0.2
case "waiting_peers", "tracker_errors": case "preparing_download":
return 0.25
case "downloading":
return 0.3 return 0.3
case "writing_files":
return 0.95
default: default:
return 0 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) { func (e *Engine) setError(err error) {
e.mu.Lock() e.mu.Lock()
defer e.mu.Unlock() defer e.mu.Unlock()
@ -365,3 +547,207 @@ func generatePeerID() [20]byte {
return id 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
}

View file

@ -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<<bitOffset) != 0
}
func pieceSizeForIndex(tf *torrentfile.TorrentFile, pieceIndex int) int {
if pieceIndex < 0 || pieceIndex >= 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()
}

View file

@ -36,6 +36,8 @@ type desktopUI struct {
sizeLabel *widget.Label sizeLabel *widget.Label
pieceLabel *widget.Label pieceLabel *widget.Label
pieceCountLabel *widget.Label pieceCountLabel *widget.Label
downloadedLabel *widget.Label
outputLabel *widget.Label
peerCountLabel *widget.Label peerCountLabel *widget.Label
mu sync.RWMutex mu sync.RWMutex
@ -75,6 +77,8 @@ func newDesktopUI(controller *appcore.Controller) *desktopUI {
sizeLabel: widget.NewLabel("-"), sizeLabel: widget.NewLabel("-"),
pieceLabel: widget.NewLabel("-"), pieceLabel: widget.NewLabel("-"),
pieceCountLabel: widget.NewLabel("-"), pieceCountLabel: widget.NewLabel("-"),
downloadedLabel: widget.NewLabel("0 / 0"),
outputLabel: widget.NewLabel("-"),
peerCountLabel: widget.NewLabel("0"), peerCountLabel: widget.NewLabel("0"),
} }
@ -84,11 +88,13 @@ func newDesktopUI(controller *appcore.Controller) *desktopUI {
loadButton := widget.NewButton("Load Torrent", ui.onLoad) loadButton := widget.NewButton("Load Torrent", ui.onLoad)
topRow := container.NewBorder(nil, nil, nil, container.NewHBox(browseButton, loadButton), ui.pathEntry) 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("Name", ui.nameLabel),
wrapMetric("Size (bytes)", ui.sizeLabel), wrapMetric("Size (bytes)", ui.sizeLabel),
wrapMetric("Piece length", ui.pieceLabel), 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("Discovered peers", ui.peerCountLabel),
wrapMetric("Phase", ui.phaseMetricLabel), wrapMetric("Phase", ui.phaseMetricLabel),
) )
@ -204,7 +210,13 @@ func (ui *desktopUI) updateFromStatus(status torrent.Status) {
ui.nameLabel.SetText(status.Name) ui.nameLabel.SetText(status.Name)
ui.sizeLabel.SetText(strconv.Itoa(status.Length)) ui.sizeLabel.SetText(strconv.Itoa(status.Length))
ui.pieceLabel.SetText(strconv.Itoa(status.PieceLength)) 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.peerCountLabel.SetText(strconv.Itoa(status.PeerCount))
ui.mu.Lock() ui.mu.Lock()
@ -219,6 +231,8 @@ func (ui *desktopUI) updateFromStatus(status torrent.Status) {
if status.LastError != "" { if status.LastError != "" {
ui.errorLabel.SetText(status.LastError) ui.errorLabel.SetText(status.LastError)
} else {
ui.errorLabel.SetText("")
} }
}) })
} }
@ -292,7 +306,7 @@ func (ui *desktopUI) newPeersTable() *widget.Table {
func() (int, int) { func() (int, int) {
ui.mu.RLock() ui.mu.RLock()
defer ui.mu.RUnlock() defer ui.mu.RUnlock()
return len(ui.peers) + 1, 3 return len(ui.peers) + 1, 6
}, },
func() fyne.CanvasObject { func() fyne.CanvasObject {
return widget.NewLabel("") return widget.NewLabel("")
@ -306,7 +320,13 @@ func (ui *desktopUI) newPeersTable() *widget.Table {
case 1: case 1:
label.SetText("Port") label.SetText("Port")
case 2: case 2:
label.SetText("State")
case 3:
label.SetText("Pieces")
case 4:
label.SetText("Source tracker") label.SetText("Source tracker")
case 5:
label.SetText("Error")
} }
return return
} }
@ -320,13 +340,22 @@ func (ui *desktopUI) newPeersTable() *widget.Table {
case 1: case 1:
label.SetText(strconv.Itoa(int(peer.Port))) label.SetText(strconv.Itoa(int(peer.Port)))
case 2: case 2:
label.SetText(peer.State)
case 3:
label.SetText(strconv.Itoa(peer.DownloadedPieces))
case 4:
label.SetText(peer.Source) label.SetText(peer.Source)
case 5:
label.SetText(peer.Error)
} }
}, },
) )
table.SetColumnWidth(0, 180) table.SetColumnWidth(0, 170)
table.SetColumnWidth(1, 80) 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 return table
} }