/*
 * DF-2905 — proof that md_done() frees only the current record's m_next
 * chain and leaks every m_nextpkt record still attached to the mdchain.
 *
 * Compile the real sys/kern/libmchain/subr_mchain.c into this KLD (stock
 * GENERIC kernels do not contain libmchain, so there is no symbol clash).
 *
 * Two modes selected at build time:
 *   default   TAILTEST  — md_done() on a 3-record mdchain; records 2 and 3
 *                         (2 * NMB mbufs) must leak on unpatched code.
 *   -DWALKTEST WALKTEST — md_next_record() x3 + md_done(); must free all
 *                         mbufs on BOTH unpatched and patched code (guards
 *                         against the fix breaking record walking or
 *                         double-freeing).
 */
#include <sys/param.h>
#include <sys/systm.h>
#include <sys/kernel.h>
#include <sys/malloc.h>
#include <sys/module.h>
#include <sys/mbuf.h>
#include <sys/mchain.h>

#include "walkflag.h"		/* generated: "#define WALKTEST 1" or empty */

#define NREC	3		/* records in the mdchain */
#define NMB	4		/* mbufs per record chain  */

static struct mbuf *
mkrec(void)
{
	struct mbuf *head = NULL, **tp = &head;
	int i;

	for (i = 0; i < NMB; i++) {
		struct mbuf *n = m_get(M_WAITOK, MT_DATA);
		n->m_len = 8;
		n->m_next = NULL;
		n->m_nextpkt = NULL;
		*tp = n;
		tp = &n->m_next;
	}
	return head;
}

#ifndef WALKTEST
static void
run_tailtest(void)
{
	struct mdchain md;
	struct mbuf *recs[NREC];
	int r;

	for (r = 0; r < NREC; r++)
		recs[r] = mkrec();
	md_initm(&md, recs[0]);
	for (r = 1; r < NREC; r++)
		md_append_record(&md, recs[r]);

	kprintf("mchainleak: TAILTEST: attached %d records x %d mbufs "
	    "(rec2=%p rec3=%p); calling md_done()\n",
	    NREC, NMB, recs[1], recs[2]);

	md_done(&md);		/* <-- the call under test */

	kprintf("mchainleak: TAILTEST: md_done() returned. "
	    "%d of %d mbufs freed (see `netstat -m` / mbread delta: "
	    "+%d leaked on vulnerable code, +0 on fixed code)\n",
	    NMB, NREC * NMB, (NREC - 1) * NMB);
}
#endif /* !WALKTEST */

#ifdef WALKTEST
static void
run_walktest(void)
{
	struct mdchain md;
	struct mbuf *recs[NREC];
	int r, rc;
	for (r = 0; r < NREC; r++)
		recs[r] = mkrec();
	md_initm(&md, recs[0]);
	for (r = 1; r < NREC; r++)
		md_append_record(&md, recs[r]);

	rc = md_next_record(&md);
	kprintf("mchainleak: WALKTEST: next#1 rc=%d md_top=%p (want rec2=%p)\n",
	    rc, md.md_top, recs[1]);
	rc = md_next_record(&md);
	kprintf("mchainleak: WALKTEST: next#2 rc=%d md_top=%p (want rec3=%p)\n",
	    rc, md.md_top, recs[2]);
	rc = md_next_record(&md);
	kprintf("mchainleak: WALKTEST: next#3 rc=%d (want ENOENT=%d) "
	    "md_top=%p (want NULL)\n", rc, ENOENT, md.md_top);
	md_done(&md);
	kprintf("mchainleak: WALKTEST: complete; mbread delta must be 0 "
	    "on both vulnerable and fixed code\n");
}
#endif /* WALKTEST */

static int
mchainleak_modevent(module_t mod, int type, void *data)
{
	switch (type) {
	case MOD_LOAD:
#ifdef WALKTEST
		run_walktest();
#else
		run_tailtest();
#endif
		return 0;
	case MOD_UNLOAD:
		return 0;
	default:
		return EOPNOTSUPP;
	}
}

static moduledata_t mchainleak_mod = {
	"mchainleak",
	mchainleak_modevent,
	NULL
};

DECLARE_MODULE(mchainleak, mchainleak_mod, SI_SUB_DRIVERS, SI_ORDER_MIDDLE);
