diff --git a/sys/netproto/mpls/mpls_var.h b/sys/netproto/mpls/mpls_var.h --- a/sys/netproto/mpls/mpls_var.h +++ b/sys/netproto/mpls/mpls_var.h @@ -56,8 +56,8 @@ void mpls_init(void); void mpls_hashfn(struct mbuf **, int); void mpls_input(struct mbuf *); -int mpls_output(struct mbuf *, struct rtentry *); -boolean_t mpls_output_process(struct mbuf *, struct rtentry *); +int mpls_output(struct mbuf **, struct rtentry *); +boolean_t mpls_output_process(struct mbuf **, struct rtentry *); #endif /* _KERNEL */ diff --git a/sys/netproto/mpls/mpls_output.c b/sys/netproto/mpls/mpls_output.c --- a/sys/netproto/mpls/mpls_output.c +++ b/sys/netproto/mpls/mpls_output.c @@ -43,17 +43,18 @@ static int mpls_push(struct mbuf **, mpls_label_t, mpls_s_t, mpls_exp_t, mpls_ttl_t); -static int mpls_swap(struct mbuf *, mpls_label_t); -static int mpls_pop(struct mbuf *, mpls_s_t *); +static int mpls_swap(struct mbuf **, mpls_label_t); +static int mpls_pop(struct mbuf **, mpls_s_t *); int -mpls_output(struct mbuf *m, struct rtentry *rt) +mpls_output(struct mbuf **mp, struct rtentry *rt) { struct sockaddr_mpls *smpls = NULL; int error = 0, i; mpls_s_t stackempty; mpls_ttl_t ttl = 255; struct ip *ip; + struct mbuf *m = *mp; M_ASSERTPKTHDR(m); @@ -85,7 +86,7 @@ 0, ttl); if (error) - return (error); + goto out; stackempty = 0; m->m_flags |= M_MPLSLABELED; break; @@ -94,24 +95,28 @@ * Operation is only permmited if label stack * is not empty. */ - if (stackempty) - return (ENOTSUP); + if (stackempty) { + error = ENOTSUP; + goto out; + } KKASSERT(m->m_flags & M_MPLSLABELED); - error = mpls_swap(m, ntohl(smpls->smpls_label)); + error = mpls_swap(&m, ntohl(smpls->smpls_label)); if (error) - return (error); + goto out; break; case MPLSLOP_POP: /* * Operation is only permmited if label stack * is not empty. */ - if (stackempty) - return (ENOTSUP); + if (stackempty) { + error = ENOTSUP; + goto out; + } KKASSERT(m->m_flags & M_MPLSLABELED); - error = mpls_pop(m, &stackempty); + error = mpls_pop(&m, &stackempty); if (error) - return (error); + goto out; /* * If we are popping out the last label then * mark the mbuf as ~M_MPLSLABELED. @@ -121,10 +126,19 @@ break; default: /* Unknown label operation */ - return (ENOTSUP); + error = ENOTSUP; + goto out; } } +out: + /* + * Propagate the (possibly new) mbuf head back to the caller. + * mpls_push/mpls_swap/mpls_pop may reallocate the head mbuf via + * M_PREPEND/m_pullup; without this write-back the caller would + * dereference a stale (possibly freed) pointer. (DF-0753/DF-0754) + */ + *mp = m; return (error); } @@ -132,7 +146,7 @@ * Returns FALSE if no further output processing required. */ boolean_t -mpls_output_process(struct mbuf *m, struct rtentry *rt) +mpls_output_process(struct mbuf **mp, struct rtentry *rt) { int error; @@ -140,9 +154,9 @@ if (!(rt->rt_flags & RTF_MPLSOPS)) return TRUE; - error = mpls_output(m, rt); + error = mpls_output(mp, rt); if (error) { - m_freem(m); + m_freem(*mp); return FALSE; } @@ -169,14 +183,20 @@ } static int -mpls_swap(struct mbuf *m, mpls_label_t label) { +mpls_swap(struct mbuf **mp, mpls_label_t label) { struct mpls *mpls; u_int32_t buf; mpls_ttl_t ttl; + struct mbuf *m = *mp; - if (m->m_len < sizeof(struct mpls) && - (m = m_pullup(m, sizeof(struct mpls))) == NULL) - return (ENOBUFS); + if (m->m_len < sizeof(struct mpls)) { + m = m_pullup(m, sizeof(struct mpls)); + if (m == NULL) { + *mp = NULL; + return (ENOBUFS); + } + *mp = m; + } mpls = mtod(m, struct mpls *); buf = ntohl(mpls->mpls_shim); @@ -194,14 +214,18 @@ } static int -mpls_pop(struct mbuf *m, mpls_s_t *sbit) { +mpls_pop(struct mbuf **mp, mpls_s_t *sbit) { struct mpls *mpls; u_int32_t buf; + struct mbuf *m = *mp; if (m->m_len < sizeof(struct mpls)) { m = m_pullup(m, sizeof(struct mpls)); - if (m == NULL) + if (m == NULL) { + *mp = NULL; return (ENOBUFS); + } + *mp = m; } mpls = mtod(m, struct mpls *); buf = ntohl(mpls->mpls_shim); diff --git a/sys/netproto/mpls/mpls_input.c b/sys/netproto/mpls/mpls_input.c --- a/sys/netproto/mpls/mpls_input.c +++ b/sys/netproto/mpls/mpls_input.c @@ -205,7 +205,7 @@ ifp = cache_rt->ro_rt->rt_ifp; dst = cache_rt->ro_rt->rt_gateway; - error = mpls_output(m, cache_rt->ro_rt); + error = mpls_output(&m, cache_rt->ro_rt); if (error) goto bad; error = (*ifp->if_output)(ifp, m, dst, cache_rt->ro_rt); diff --git a/sys/netinet/ip_output.c b/sys/netinet/ip_output.c --- a/sys/netinet/ip_output.c +++ b/sys/netinet/ip_output.c @@ -692,7 +692,7 @@ #endif #ifdef MPLS - if (!mpls_output_process(m, ro->ro_rt)) + if (!mpls_output_process(&m, ro->ro_rt)) goto done; #endif error = ifp->if_output(ifp, m, (struct sockaddr *)dst, @@ -736,7 +736,7 @@ m->m_pkthdr.len); } #ifdef MPLS - if (!mpls_output_process(m, ro->ro_rt)) + if (!mpls_output_process(&m, ro->ro_rt)) continue; #endif error = ifp->if_output(ifp, m, (struct sockaddr *)dst,