#include "naut/extension.h" #include "naut/hash.h" #include "naut/peer.h" #include #include #include #include #include #include #include #include #include #define EXT_RESERVED 0x0000000000100000ULL static bool send_all(int fd, const void *data, size_t len) { const uint8_t *p = data; while (len) { ssize_t n = send(fd, p, len, MSG_NOSIGNAL); if (n < 0) { if (errno == EINTR) continue; return false; } if (n == 0) return false; p += n; len -= (size_t)n; } return true; } static bool recv_exact(int fd, void *data, size_t len) { uint8_t *p = data; while (len) { ssize_t n = recv(fd, p, len, 0); if (n < 0) { if (errno == EINTR) continue; return false; } if (n == 0) return false; p += n; len -= (size_t)n; } return true; } static int connect_peer(const naut_peer_addr *peer) { int fd = socket(AF_INET, SOCK_STREAM, 0); if (fd < 0) return -1; struct sockaddr_in addr; memset(&addr, 0, sizeof addr); addr.sin_family = AF_INET; addr.sin_port = htons(peer->port); memcpy(&addr.sin_addr, peer->ip, sizeof peer->ip); if (connect(fd, (struct sockaddr *)&addr, sizeof addr) != 0) { close(fd); return -1; } int one = 1; setsockopt(fd, IPPROTO_TCP, TCP_NODELAY, &one, sizeof one); struct timeval timeout = { .tv_sec = 10 }; setsockopt(fd, SOL_SOCKET, SO_RCVTIMEO, &timeout, sizeof timeout); setsockopt(fd, SOL_SOCKET, SO_SNDTIMEO, &timeout, sizeof timeout); return fd; } naut_err naut_metadata_fetch(const naut_peer_addr *peer, const uint8_t info_hash[20], const uint8_t peer_id[20], uint8_t **info, size_t *info_len) { if (!peer || !info_hash || !peer_id || !info || !info_len) return NAUT_ERR_INVAL; *info = NULL; *info_len = 0; int fd = connect_peer(peer); if (fd < 0) return NAUT_ERR_IO; naut_err result = NAUT_ERR_IO; uint8_t *metadata = NULL, *frame = NULL, *buffer = NULL; bool *received = NULL; uint8_t handshake[NAUT_HANDSHAKE_LEN]; naut_peer_handshake_build(handshake, info_hash, peer_id, EXT_RESERVED); if (!send_all(fd, handshake, sizeof handshake) || !recv_exact(fd, handshake, sizeof handshake)) goto done; uint8_t remote_hash[20], remote_id[20]; uint64_t reserved = 0; if (!naut_peer_handshake_parse(handshake, remote_hash, remote_id, &reserved) || memcmp(remote_hash, info_hash, 20) != 0 || (reserved & EXT_RESERVED) == 0) { result = NAUT_ERR_PROTO; goto done; } size_t frame_len = 0; result = naut_ext_build_handshake(NAUT_EXT_UT_METADATA, NAUT_EXT_UT_PEX, 0, 0, &frame, &frame_len); if (result != NAUT_OK || !send_all(fd, frame, frame_len)) { result = NAUT_ERR_IO; goto done; } free(frame); frame = NULL; size_t cap = 128 * 1024, len = 0; buffer = malloc(cap); if (!buffer) { result = NAUT_ERR_NOMEM; goto done; } naut_ext_handshake remote_ext = {0}; uint32_t piece_count = 0, received_count = 0; while (!metadata || received_count < piece_count) { if (len == cap) { if (cap >= NAUT_METADATA_MAX + (1u << 20)) { result = NAUT_ERR_PROTO; goto done; } size_t next_cap = cap * 2; uint8_t *next = realloc(buffer, next_cap); if (!next) { result = NAUT_ERR_NOMEM; goto done; } buffer = next; cap = next_cap; } ssize_t n = recv(fd, buffer + len, cap - len, 0); if (n < 0) { if (errno == EINTR) continue; result = NAUT_ERR_IO; goto done; } if (n == 0) { result = NAUT_ERR_IO; goto done; } len += (size_t)n; size_t pos = 0; for (;;) { naut_msg msg; int consumed = naut_peer_msg_parse(buffer + pos, len - pos, &msg); if (consumed == 0) break; if (consumed < 0) { result = NAUT_ERR_PROTO; goto done; } pos += (size_t)consumed; if (msg.type != NAUT_MSG_EXTENDED || msg.payload_len < 1) continue; uint8_t ext_id = msg.payload[0]; const uint8_t *payload = msg.payload + 1; size_t payload_len = msg.payload_len - 1; if (ext_id == 0) { result = naut_ext_parse_handshake(payload, payload_len, &remote_ext); if (result != NAUT_OK || remote_ext.ut_metadata == 0 || remote_ext.metadata_size == 0) { result = NAUT_ERR_PROTO; goto done; } if (!metadata) { metadata = malloc(remote_ext.metadata_size); piece_count = (remote_ext.metadata_size + NAUT_METADATA_BLOCK - 1) / NAUT_METADATA_BLOCK; received = calloc(piece_count, sizeof(*received)); if (!metadata || !received) { result = NAUT_ERR_NOMEM; goto done; } for (uint32_t piece = 0; piece < piece_count; piece++) { result = naut_metadata_build( remote_ext.ut_metadata, NAUT_METADATA_REQUEST, piece, 0, NULL, 0, &frame, &frame_len); if (result != NAUT_OK || !send_all(fd, frame, frame_len)) { result = NAUT_ERR_IO; goto done; } free(frame); frame = NULL; } } } else if (metadata && (ext_id == NAUT_EXT_UT_METADATA || ext_id == remote_ext.ut_metadata)) { naut_metadata_msg metadata_msg; result = naut_metadata_parse(payload, payload_len, &metadata_msg); if (result != NAUT_OK || metadata_msg.type == NAUT_METADATA_REJECT || metadata_msg.total_size != remote_ext.metadata_size || metadata_msg.piece >= piece_count) { result = NAUT_ERR_PROTO; goto done; } if (metadata_msg.type == NAUT_METADATA_DATA && !received[metadata_msg.piece]) { memcpy(metadata + (size_t)metadata_msg.piece * NAUT_METADATA_BLOCK, metadata_msg.data, metadata_msg.data_len); received[metadata_msg.piece] = true; received_count++; } } } memmove(buffer, buffer + pos, len - pos); len -= pos; } uint8_t digest[20]; naut_sha1(metadata, remote_ext.metadata_size, digest); if (memcmp(digest, info_hash, 20) != 0) { result = NAUT_ERR_PROTO; goto done; } *info = metadata; *info_len = remote_ext.metadata_size; metadata = NULL; result = NAUT_OK; done: close(fd); free(metadata); free(received); free(frame); free(buffer); return result; }