#include <sys/param.h>
#include <sys/systm.h>
#include <sys/kernel.h>
#include <sys/malloc.h>
#include <sys/msgport.h>
#include <sys/protosw.h>
#include <sys/socket.h>
#include <sys/socketvar.h>
#include <sys/socketops.h>
#include <sys/thread.h>
#include <sys/msgport2.h>
#include <sys/spinlock2.h>
#include <sys/sysctl.h>
#include <sys/mbuf.h>
#include <vm/pmap.h>
#include <net/netmsg2.h>
#include <net/netisr2.h>
#include <sys/socketvar2.h>
#include <net/netisr.h>
#include <net/netmsg.h>
static int async_rcvd_drop_race = 0;
SYSCTL_INT(_kern_ipc, OID_AUTO, async_rcvd_drop_race, CTLFLAG_RW,
&async_rcvd_drop_race, 0, "# of asynchronized pru_rcvd msg drop races");
void
so_pru_abort_async(struct socket *so)
{
struct netmsg_pru_abort *msg;
msg = kmalloc(sizeof(*msg), M_LWKTMSG, M_WAITOK | M_ZERO);
netmsg_init(&msg->base, so, &netisr_afree_free_so_rport,
0, so->so_proto->pr_usrreqs->pru_abort);
lwkt_sendmsg(so->so_port, &msg->base.lmsg);
}
void
so_pru_abort_direct(struct socket *so)
{
struct netmsg_pru_abort msg;
netisr_fn_t func = so->so_proto->pr_usrreqs->pru_abort;
netmsg_init(&msg.base, so, &netisr_adone_rport, 0, func);
msg.base.lmsg.ms_flags &= ~(MSGF_REPLY | MSGF_DONE);
msg.base.lmsg.ms_flags |= MSGF_SYNC;
func((netmsg_t)&msg);
KKASSERT(msg.base.lmsg.ms_flags & MSGF_DONE);
sofree(msg.base.nm_so);
}
int
so_pru_accept(struct socket *so, struct sockaddr **nam)
{
struct netmsg_pru_accept msg;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_accept);
msg.nm_nam = nam;
return lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
}
int
so_pru_attach(struct socket *so, int proto, struct pru_attach_info *ai)
{
struct netmsg_pru_attach msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_attach);
msg.nm_proto = proto;
msg.nm_ai = ai;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pru_attach_direct(struct socket *so, int proto, struct pru_attach_info *ai)
{
struct netmsg_pru_attach msg;
netisr_fn_t func = so->so_proto->pr_usrreqs->pru_attach;
netmsg_init(&msg.base, so, &netisr_adone_rport, 0, func);
msg.base.lmsg.ms_flags &= ~(MSGF_REPLY | MSGF_DONE);
msg.base.lmsg.ms_flags |= MSGF_SYNC;
msg.nm_proto = proto;
msg.nm_ai = ai;
func((netmsg_t)&msg);
KKASSERT(msg.base.lmsg.ms_flags & MSGF_DONE);
return(msg.base.lmsg.ms_error);
}
int
so_pru_attach_fast(struct socket *so, int proto, struct pru_attach_info *ai)
{
struct netmsg_pru_attach *msg;
int error;
error = so->so_proto->pr_usrreqs->pru_preattach(so, proto, ai);
if (error)
return error;
msg = kmalloc(sizeof(*msg), M_LWKTMSG, M_WAITOK | M_NULLOK);
if (msg == NULL) {
return so_pru_attach(so, proto, NULL );
}
netmsg_init(&msg->base, so, &netisr_afree_rport, 0,
so->so_proto->pr_usrreqs->pru_attach);
msg->nm_proto = proto;
msg->nm_ai = NULL;
if (so->so_port == netisr_curport())
lwkt_sendmsg_oncpu(so->so_port, &msg->base.lmsg);
else
lwkt_sendmsg(so->so_port, &msg->base.lmsg);
return 0;
}
int
so_pru_bind(struct socket *so, struct sockaddr *nam, struct thread *td)
{
struct netmsg_pru_bind msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_bind);
msg.nm_nam = nam;
msg.nm_td = td;
msg.nm_flags = 0;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pru_connect(struct socket *so, struct sockaddr *nam, struct thread *td)
{
struct netmsg_pru_connect msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_connect);
msg.nm_nam = nam;
msg.nm_td = td;
msg.nm_m = NULL;
msg.nm_sndflags = 0;
msg.nm_flags = 0;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pru_connect_async(struct socket *so, struct sockaddr *nam, struct thread *td)
{
struct netmsg_pru_connect *msg;
int error, flags;
KASSERT(so->so_proto->pr_usrreqs->pru_preconnect != NULL,
("async pru_connect is not supported"));
msg = kmalloc(sizeof(*msg) + nam->sa_len, M_LWKTMSG,
M_WAITOK | M_NULLOK);
if (msg == NULL) {
return so_pru_connect(so, nam, td);
}
error = so->so_proto->pr_usrreqs->pru_preconnect(so, nam, td);
if (error) {
kfree(msg, M_LWKTMSG);
return error;
}
flags = PRUC_ASYNC;
if (td != NULL && (so->so_proto->pr_flags & PR_ACONN_HOLDTD)) {
lwkt_hold(td);
flags |= PRUC_HELDTD;
}
netmsg_init(&msg->base, so, &netisr_afree_rport, 0,
so->so_proto->pr_usrreqs->pru_connect);
msg->nm_nam = (struct sockaddr *)(msg + 1);
memcpy(msg->nm_nam, nam, nam->sa_len);
msg->nm_td = td;
msg->nm_m = NULL;
msg->nm_sndflags = 0;
msg->nm_flags = flags;
if (so->so_port == netisr_curport())
lwkt_sendmsg_oncpu(so->so_port, &msg->base.lmsg);
else
lwkt_sendmsg(so->so_port, &msg->base.lmsg);
return 0;
}
int
so_pru_connect2(struct socket *so1, struct socket *so2, struct ucred *cred)
{
struct netmsg_pru_connect2 msg;
int error;
netmsg_init(&msg.base, so1, &curthread->td_msgport,
0, so1->so_proto->pr_usrreqs->pru_connect2);
msg.nm_so1 = so1;
msg.nm_so2 = so2;
msg.nm_cred = cred;
error = lwkt_domsg(so1->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pru_control_direct(struct socket *so, u_long cmd, caddr_t data,
struct ifnet *ifp)
{
struct netmsg_pru_control msg;
netisr_fn_t func = so->so_proto->pr_usrreqs->pru_control;
netmsg_init(&msg.base, so, &netisr_adone_rport, 0, func);
msg.base.lmsg.ms_flags &= ~(MSGF_REPLY | MSGF_DONE);
msg.base.lmsg.ms_flags |= MSGF_SYNC;
msg.nm_cmd = cmd;
msg.nm_data = data;
msg.nm_ifp = ifp;
msg.nm_td = curthread;
func((netmsg_t)&msg);
KKASSERT(msg.base.lmsg.ms_flags & MSGF_DONE);
return(msg.base.lmsg.ms_error);
}
int
so_pru_detach(struct socket *so)
{
struct netmsg_pru_detach msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_detach);
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pru_detach_direct(struct socket *so)
{
struct netmsg_pru_detach msg;
netisr_fn_t func = so->so_proto->pr_usrreqs->pru_detach;
netmsg_init(&msg.base, so, &netisr_adone_rport, 0, func);
msg.base.lmsg.ms_flags &= ~(MSGF_REPLY | MSGF_DONE);
msg.base.lmsg.ms_flags |= MSGF_SYNC;
func((netmsg_t)&msg);
KKASSERT(msg.base.lmsg.ms_flags & MSGF_DONE);
return(msg.base.lmsg.ms_error);
}
int
so_pru_disconnect(struct socket *so)
{
struct netmsg_pru_disconnect msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_disconnect);
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
void
so_pru_disconnect_direct(struct socket *so)
{
struct netmsg_pru_disconnect msg;
netisr_fn_t func = so->so_proto->pr_usrreqs->pru_disconnect;
netmsg_init(&msg.base, so, &netisr_adone_rport, 0, func);
msg.base.lmsg.ms_flags &= ~(MSGF_REPLY | MSGF_DONE);
msg.base.lmsg.ms_flags |= MSGF_SYNC;
func((netmsg_t)&msg);
KKASSERT(msg.base.lmsg.ms_flags & MSGF_DONE);
}
int
so_pru_listen(struct socket *so, struct thread *td)
{
struct netmsg_pru_listen msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_listen);
msg.nm_td = td;
msg.nm_flags = 0;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pru_peeraddr(struct socket *so, struct sockaddr **nam)
{
struct netmsg_pru_peeraddr msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_peeraddr);
msg.nm_nam = nam;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pru_rcvd(struct socket *so, int flags)
{
struct netmsg_pru_rcvd msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_rcvd);
msg.nm_flags = flags;
msg.nm_pru_flags = 0;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
void
so_pru_rcvd_async(struct socket *so)
{
lwkt_msg_t lmsg = &so->so_rcvd_msg.base.lmsg;
KASSERT(so->so_proto->pr_flags & PR_ASYNC_RCVD,
("async pru_rcvd is not supported"));
spin_lock(&so->so_rcvd_spin);
if ((so->so_rcvd_msg.nm_pru_flags & PRUR_DEAD) == 0) {
if (lmsg->ms_flags & MSGF_DONE) {
lwkt_sendmsg_prepare(so->so_port, lmsg);
spin_unlock(&so->so_rcvd_spin);
if (so->so_port == netisr_curport())
lwkt_sendmsg_start_oncpu(so->so_port, lmsg);
else
lwkt_sendmsg_start(so->so_port, lmsg);
} else {
spin_unlock(&so->so_rcvd_spin);
}
} else {
spin_unlock(&so->so_rcvd_spin);
}
}
int
so_pru_rcvoob(struct socket *so, struct mbuf *m, int flags)
{
struct netmsg_pru_rcvoob msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_rcvoob);
msg.nm_m = m;
msg.nm_flags = flags;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pru_send(struct socket *so, int flags, struct mbuf *m,
struct sockaddr *addr, struct mbuf *control, struct thread *td)
{
struct netmsg_pru_send msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_send);
msg.nm_flags = flags;
msg.nm_m = m;
msg.nm_addr = addr;
msg.nm_control = control;
msg.nm_td = td;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
void
so_pru_sync(struct socket *so)
{
struct netmsg_base msg;
netmsg_init(&msg, so, &curthread->td_msgport, 0,
netmsg_sync_handler);
lwkt_domsg(so->so_port, &msg.lmsg, 0);
}
void
so_pru_send_async(struct socket *so, int flags, struct mbuf *m,
struct sockaddr *addr0, struct mbuf *control, struct thread *td)
{
struct netmsg_pru_send *msg;
struct sockaddr *addr = NULL;
KASSERT(so->so_proto->pr_flags & PR_ASYNC_SEND,
("async pru_send is not supported"));
if (addr0 != NULL) {
addr = kmalloc(addr0->sa_len, M_SONAME, M_WAITOK | M_NULLOK);
if (addr == NULL) {
so_pru_send(so, flags, m, addr0, control, td);
return;
}
memcpy(addr, addr0, addr0->sa_len);
flags |= PRUS_FREEADDR;
}
flags |= PRUS_NOREPLY;
if (td != NULL && (so->so_proto->pr_flags & PR_ASEND_HOLDTD)) {
lwkt_hold(td);
flags |= PRUS_HELDTD;
}
msg = &m->m_hdr.mh_sndmsg;
netmsg_init(&msg->base, so, &netisr_apanic_rport,
0, so->so_proto->pr_usrreqs->pru_send);
msg->nm_flags = flags;
msg->nm_m = m;
msg->nm_addr = addr;
msg->nm_control = control;
msg->nm_td = td;
if (so->so_port == netisr_curport())
lwkt_sendmsg_oncpu(so->so_port, &msg->base.lmsg);
else
lwkt_sendmsg(so->so_port, &msg->base.lmsg);
}
int
so_pru_sense(struct socket *so, struct stat *sb)
{
struct netmsg_pru_sense msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_sense);
msg.nm_stat = sb;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pru_shutdown(struct socket *so)
{
struct netmsg_pru_shutdown msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_shutdown);
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pru_sockaddr(struct socket *so, struct sockaddr **nam)
{
struct netmsg_pru_sockaddr msg;
int error;
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_usrreqs->pru_sockaddr);
msg.nm_nam = nam;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
int
so_pr_ctloutput(struct socket *so, struct sockopt *sopt)
{
struct netmsg_pr_ctloutput msg;
int error;
KKASSERT(!sopt->sopt_val || kva_p(sopt->sopt_val));
if (sopt->sopt_dir == SOPT_SET && so->so_proto->pr_ctloutmsg != NULL) {
struct netmsg_pr_ctloutput *amsg;
amsg = so->so_proto->pr_ctloutmsg(sopt);
if (amsg != NULL) {
netmsg_init(&amsg->base, so, &netisr_afree_rport, 0,
so->so_proto->pr_ctloutput);
if (so->so_port == netisr_curport()) {
lwkt_sendmsg_oncpu(so->so_port,
&amsg->base.lmsg);
} else {
lwkt_sendmsg(so->so_port, &amsg->base.lmsg);
}
return 0;
}
}
netmsg_init(&msg.base, so, &curthread->td_msgport,
0, so->so_proto->pr_ctloutput);
msg.nm_flags = 0;
msg.nm_sopt = sopt;
error = lwkt_domsg(so->so_port, &msg.base.lmsg, 0);
return (error);
}
struct lwkt_port *
so_pr_ctlport(struct protosw *pr, int cmd, struct sockaddr *arg,
void *extra, int *cpuid)
{
if (pr->pr_ctlport == NULL)
return NULL;
KKASSERT(pr->pr_ctlinput != NULL);
return pr->pr_ctlport(cmd, arg, extra, cpuid);
}
void
so_pr_ctlinput(struct protosw *pr, int cmd, struct sockaddr *arg, void *extra)
{
struct netmsg_pr_ctlinput msg;
lwkt_port_t port;
int cpuid;
port = so_pr_ctlport(pr, cmd, arg, extra, &cpuid);
if (port == NULL)
return;
netmsg_init(&msg.base, NULL, &curthread->td_msgport,
0, pr->pr_ctlinput);
msg.nm_cmd = cmd;
msg.nm_direct = 0;
msg.nm_arg = arg;
msg.nm_extra = extra;
lwkt_domsg(port, &msg.base.lmsg, 0);
}
void
so_pr_ctlinput_direct(struct protosw *pr, int cmd, struct sockaddr *arg,
void *extra)
{
struct netmsg_pr_ctlinput msg;
netisr_fn_t func;
lwkt_port_t port;
int cpuid;
port = so_pr_ctlport(pr, cmd, arg, extra, &cpuid);
if (port == NULL)
return;
if (cpuid != netisr_ncpus && cpuid != mycpuid)
return;
func = pr->pr_ctlinput;
netmsg_init(&msg.base, NULL, &netisr_adone_rport, 0, func);
msg.base.lmsg.ms_flags &= ~(MSGF_REPLY | MSGF_DONE);
msg.base.lmsg.ms_flags |= MSGF_SYNC;
msg.nm_cmd = cmd;
msg.nm_direct = 1;
msg.nm_arg = arg;
msg.nm_extra = extra;
func((netmsg_t)&msg);
KKASSERT(msg.base.lmsg.ms_flags & MSGF_DONE);
}
void
netmsg_so_notify(netmsg_t msg)
{
struct socket *so = msg->base.nm_so;
struct signalsockbuf *ssb;
ssb = (msg->notify.nm_etype & NM_REVENT) ? &so->so_rcv : &so->so_snd;
lwkt_getpooltoken(so);
atomic_set_int(&ssb->ssb_flags, SSB_MEVENT);
if (msg->notify.nm_predicate(&msg->notify)) {
if (TAILQ_EMPTY(&ssb->ssb_mlist))
atomic_clear_int(&ssb->ssb_flags, SSB_MEVENT);
lwkt_relpooltoken(so);
lwkt_replymsg(&msg->base.lmsg,
msg->base.lmsg.ms_error);
} else {
TAILQ_INSERT_TAIL(&ssb->ssb_mlist, &msg->notify, nm_list);
atomic_set_int(&ssb->ssb_flags, SSB_MEVENT);
lwkt_relpooltoken(so);
}
}
void
netmsg_so_notify_doabort(lwkt_msg_t lmsg)
{
struct netmsg_so_notify_abort msg;
if ((lmsg->ms_flags & (MSGF_DONE | MSGF_REPLY)) == 0) {
const struct netmsg_base *nmsg =
(const struct netmsg_base *)lmsg;
netmsg_init(&msg.base, nmsg->nm_so, &curthread->td_msgport,
0, netmsg_so_notify_abort);
msg.nm_notifymsg = (void *)lmsg;
lwkt_domsg(lmsg->ms_target_port, &msg.base.lmsg, 0);
}
}
void
netmsg_so_notify_abort(netmsg_t msg)
{
struct netmsg_so_notify_abort *abrtmsg = &msg->notify_abort;
struct netmsg_so_notify *nmsg = abrtmsg->nm_notifymsg;
struct signalsockbuf *ssb;
lwkt_getpooltoken(nmsg->base.nm_so);
if ((nmsg->base.lmsg.ms_flags & (MSGF_DONE | MSGF_REPLY)) == 0) {
ssb = (nmsg->nm_etype & NM_REVENT) ?
&nmsg->base.nm_so->so_rcv :
&nmsg->base.nm_so->so_snd;
TAILQ_REMOVE(&ssb->ssb_mlist, nmsg, nm_list);
lwkt_relpooltoken(nmsg->base.nm_so);
lwkt_replymsg(&nmsg->base.lmsg, EINTR);
} else {
lwkt_relpooltoken(nmsg->base.nm_so);
}
lwkt_replymsg(&abrtmsg->base.lmsg, 0);
}
void
so_async_rcvd_reply(struct socket *so)
{
spin_lock(&so->so_rcvd_spin);
lwkt_replymsg(&so->so_rcvd_msg.base.lmsg, 0);
spin_unlock(&so->so_rcvd_spin);
}
void
so_async_rcvd_drop(struct socket *so)
{
lwkt_msg_t lmsg = &so->so_rcvd_msg.base.lmsg;
spin_lock(&so->so_rcvd_spin);
so->so_rcvd_msg.nm_pru_flags |= PRUR_DEAD;
again:
lwkt_dropmsg(lmsg);
if ((lmsg->ms_flags & MSGF_DONE) == 0) {
++async_rcvd_drop_race;
ssleep(so, &so->so_rcvd_spin, 0, "soadrop", 1);
goto again;
}
spin_unlock(&so->so_rcvd_spin);
}