Legend:
Page
Library
Module
Module type
Parameter
Class
Class type
Source
Source file dns_forward_lwt_unix.ml
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480(*
* Copyright (C) 2016 David Scott <dave@recoil.org>
*
* Permission to use, copy, modify, and distribute this software for any
* purpose with or without fee is hereby granted, provided that the above
* copyright notice and this permission notice appear in all copies.
*
* THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES
* WITH REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF
* MERCHANTABILITY AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR
* ANY SPECIAL, DIRECT, INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES
* WHATSOEVER RESULTING FROM LOSS OF USE, DATA OR PROFITS, WHETHER IN AN
* ACTION OF CONTRACT, NEGLIGENCE OR OTHER TORTIOUS ACTION, ARISING OUT OF
* OR IN CONNECTION WITH THE USE OR PERFORMANCE OF THIS SOFTWARE.
*
*)openLwt.Infixletsrc=letsrc=Logs.Src.create"Dns_forward_lwt_unix"~doc:"Lwt_unix-based I/O"inLogs.Src.set_levelsrc(SomeLogs.Debug);srcmoduleLog=(valLogs.src_logsrc:Logs.LOG)letdefault_read_buffer_size=65536letmax_udp_length=65507(* IP datagram (65535) - IP header(20) - UDP header(8) *)letstring_of_sockaddr=function|Lwt_unix.ADDR_INET(ip,port)->Unix.string_of_inet_addrip^":"^(string_of_intport)|Lwt_unix.ADDR_UNIXpath->pathmoduleCommon=struct(** Both UDP and TCP *)typeerror=[`Msgofstring]typewrite_error=Mirage_flow.write_errorletpp_errorppf(`Msgx)=Fmt.stringppfxletpp_write_error=Mirage_flow.pp_write_errorleterrorffmt=Printf.ksprintf(funs->Lwt.return(Error(`Msgs)))fmttypeaddress=Ipaddr.t*inttypebuffer=Cstruct.tletsockaddr_of_address(dst,dst_port)=Unix.ADDR_INET(Unix.inet_addr_of_string@@Ipaddr.to_stringdst,dst_port)letaddress_of_sockaddr=function|Lwt_unix.ADDR_INET(ip,port)->(trySome(Ipaddr.of_string_exn@@Unix.string_of_inet_addrip,port)with_->None)|_->Noneletstring_of_address(dst,dst_port)=Ipaddr.to_stringdst^":"^(string_of_intdst_port)type'aio='aLwt.tletgetsocknamefn_namefd_opt=matchfd_optwith|None->failwith(fn_name^": socket is closed")|Somefd->beginmatchLwt_unix.getsocknamefdwith|Lwt_unix.ADDR_INET(iaddr,port)->Ipaddr.V4(Ipaddr.V4.of_string_exn(Unix.string_of_inet_addriaddr)),port|_->invalid_arg(fn_name^": passed a non-TCP socket")endendmoduleTcp=structincludeCommontypeflow={mutablefd:Lwt_unix.file_descroption;read_buffer_size:int;mutableread_buffer:Cstruct.t;address:address;}letof_fd~read_buffer_sizeaddressfd=letread_buffer=Cstruct.createread_buffer_sizein{fd=Somefd;read_buffer_size;read_buffer;address}letstring_of_flowflow=Printf.sprintf"tcp -> %s"(string_of_addressflow.address)letconnect?(read_buffer_size=default_read_buffer_size)address=letdescription=Printf.sprintf"tcp -> %s"(string_of_addressaddress)inLog.debug(funf->f"%s: connect"description);letsockaddr=sockaddr_of_addressaddressinletfd=Lwt_unix.socketLwt_unix.PF_INETLwt_unix.SOCK_STREAM0inLwt.catch(fun()->Lwt_unix.connectfdsockaddr>>=fun()->Lwt.return(Ok(of_fd~read_buffer_sizeaddressfd)))(fune->Lwt_unix.closefd>>=fun()->errorf"%s: Lwt_unix.connect: caught %s"description(Printexc.to_stringe))letreadt=matcht.fdwith|None->Lwt.return(Ok`Eof)|Somefd->ifCstruct.lent.read_buffer=0thent.read_buffer<-Cstruct.createt.read_buffer_size;Lwt.catch(fun()->Lwt_bytes.readfdt.read_buffer.Cstruct.buffert.read_buffer.Cstruct.offt.read_buffer.Cstruct.len>>=function|0->Lwt.return(Ok`Eof)|n->letresults=Cstruct.subt.read_buffer0nint.read_buffer<-Cstruct.shiftt.read_buffern;Lwt.return(Ok(`Dataresults)))(fune->Log.err(funf->f"%s: read caught %s returning Eof"(string_of_flowt)(Printexc.to_stringe));Lwt.return(Ok`Eof))letwritetbuf=matcht.fdwith|None->Lwt.return(Error`Closed)|Somefd->Lwt.catch(fun()->Lwt_cstruct.(complete(writefd)buf)>>=fun()->Lwt.return(Ok()))(function|Unix.Unix_error(Unix.ECONNRESET,_,_)->Lwt.return(Error`Closed)|e->Log.err(funf->f"%s: write caught %s returning Eof"(string_of_flowt)(Printexc.to_stringe));Lwt.return(Error`Closed))letwritevtbufs=matcht.fdwith|None->Lwt.return(Error`Closed)|Somefd->Lwt.catch(fun()->letrecloop=function|[]->Lwt.return(Ok())|buf::bufs->Lwt_cstruct.(complete(writefd)buf)>>=fun()->loopbufsinloopbufs)(fun_e->Lwt.return(Error`Closed))letcloset=matcht.fdwith|None->Lwt.return_unit|Somefd->t.fd<-None;Log.debug(funf->f"%s: Tcp.close"(string_of_flowt));Lwt_unix.closefdletshutdown_readt=matcht.fdwith|None->Lwt.return_unit|Somefd->Lwt.catch(fun()->Lwt_unix.shutdownfdUnix.SHUTDOWN_RECEIVE;Lwt.return_unit)(function|Unix.Unix_error(Unix.ENOTCONN,_,_)->Lwt.return_unit|e->Log.err(funf->f"%s: Lwt_unix.shutdown receive caught %s"(string_of_flowt)(Printexc.to_stringe));Lwt.return_unit)letshutdown_writet=matcht.fdwith|None->Lwt.return_unit|Somefd->Lwt.catch(fun()->Lwt_unix.shutdownfdUnix.SHUTDOWN_SEND;Lwt.return_unit)(function|Unix.Unix_error(Unix.ENOTCONN,_,_)->Lwt.return_unit|e->Log.err(funf->f"%s: Lwt_unix.shutdown send caught %s"(string_of_flowt)(Printexc.to_stringe));Lwt.return_unit)typeserver={mutableserver_fd:Lwt_unix.file_descroption;read_buffer_size:int;address:address;}letstring_of_servert=Printf.sprintf"listen:tcp <- %s"(string_of_addresst.address)letbindaddress=letfd=Lwt_unix.socketLwt_unix.PF_INETLwt_unix.SOCK_STREAM0inLwt.catch(fun()->Lwt_unix.setsockoptfdLwt_unix.SO_REUSEADDRtrue;Lwt_unix.bindfd(sockaddr_of_addressaddress)>|=fun()->Ok{server_fd=Somefd;read_buffer_size=default_read_buffer_size;address})(fune->Lwt_unix.closefd>>=fun()->errorf"listen:tcp <- %s caught %s"(string_of_addressaddress)(Printexc.to_stringe))letgetsocknameserver=getsockname"Tcp.getsockname"server.server_fdletshutdownserver=matchserver.server_fdwith|None->Lwt.return_unit|Somefd->server.server_fd<-None;Log.debug(funf->f"%s: close server socket"(string_of_serverserver));Lwt_unix.closefdletlisten(server:server)cb=letrecloopfd=Lwt_unix.acceptfd>>=fun(client,sockaddr)->letread_buffer_size=server.read_buffer_sizeinLwt.async(fun()->Lwt.catch(fun()->(matchaddress_of_sockaddrsockaddrwith|Someaddress->Lwt.returnaddress|_->Lwt.fail(Failure"unknown incoming socket address"))>>=funaddress->Lwt.return(Some(of_fd~read_buffer_sizeaddressclient)))(fun_e->Lwt_unix.closeclient>>=fun()->Lwt.return_none)>>=function|None->Lwt.return_unit|Someflow->Lwt.finalize(fun()->Lwt.catch(fun()->cbflow)(fune->Log.info(funf->f"tcp:%s <- %s: caught %s so closing flow"(string_of_serverserver)(string_of_sockaddrsockaddr)(Printexc.to_stringe));Lwt.return_unit))(fun()->closeflow));loopfdinmatchserver.server_fdwith|None->()|Somefd->Lwt.async(fun()->Lwt.catch(fun()->Lwt.finalize(fun()->Lwt_unix.listenfd32;loopfd)(fun()->shutdownserver))(fune->Log.info(funf->f"%s: caught %s so shutting down server"(string_of_serverserver)(Printexc.to_stringe));Lwt.return_unit))endmoduleUdp=structincludeCommontypeflow={mutablefd:Lwt_unix.file_descroption;read_buffer_size:int;mutablealready_read:Cstruct.toption;sockaddr:Unix.sockaddr;address:address;}letstring_of_flowt=Printf.sprintf"udp -> %s"(string_of_addresst.address)letof_fd?(read_buffer_size=max_udp_length)?(already_read=None)sockaddraddressfd={fd=Somefd;read_buffer_size;already_read;sockaddr;address}letconnect?read_buffer_sizeaddress=Log.debug(funf->f"udp -> %s: connect"(string_of_addressaddress));letfd=Lwt_unix.socketLwt_unix.PF_INETLwt_unix.SOCK_DGRAM0in(* Win32 requires all sockets to be bound however macOS and Linux don't *)Lwt.catch(fun()->Lwt_unix.bindfd(Lwt_unix.ADDR_INET(Unix.inet_addr_any,0)))(fun_->Lwt.return())>|=fun()->letsockaddr=sockaddr_of_addressaddressinOk(of_fd?read_buffer_sizesockaddraddressfd)letreadt=matcht.fd,t.already_readwith|None,_->Lwt.return(Ok`Eof)|Some_,SomedatawhenCstruct.lendata>0->t.already_read<-Some(Cstruct.subdata00);(* next read is `Eof *)Lwt.return(Ok(`Datadata))|Some_,Some_->Lwt.return(Ok`Eof)|Somefd,None->letbuffer=Cstruct.createt.read_buffer_sizeinletbytes=Bytes.maket.read_buffer_size'\000'inLwt.catch(fun()->(* Lwt on Win32 doesn't support Lwt_bytes.recvfrom *)Lwt_unix.recvfromfdbytes0(Bytes.lengthbytes)[]>>=fun(n,_)->Cstruct.blit_from_bytesbytes0buffer0n;letresponse=Cstruct.subbuffer0ninLwt.return(Ok(`Dataresponse)))(fune->Log.err(funf->f"%s: recvfrom caught %s returning Eof"(string_of_flowt)(Printexc.to_stringe));Lwt.return(Ok`Eof))letwritetbuf=matcht.fdwith|None->Lwt.return(Error`Closed)|Somefd->Lwt.catch(fun()->(* Lwt on Win32 doesn't support Lwt_bytes.sendto *)letbytes=Bytes.make(Cstruct.lenbuf)'\000'inCstruct.blit_to_bytesbuf0bytes0(Cstruct.lenbuf);Lwt_unix.sendtofdbytes0(Bytes.lengthbytes)[]t.sockaddr>|=fun_n->Ok())(fune->Log.err(funf->f"%s: sendto caught %s returning Eof"(string_of_flowt)(Printexc.to_stringe));Lwt.return(Error`Closed))letwritevtbufs=writet(Cstruct.concatbufs)letcloset=matcht.fdwith|None->Lwt.return_unit|Somefd->t.fd<-None;Log.debug(funf->f"%s: close"(string_of_flowt));Lwt_unix.closefdletshutdown_read_t=Lwt.return_unitletshutdown_write_t=Lwt.return_unittypeserver={mutableserver_fd:Lwt_unix.file_descroption;address:address;}letstring_of_servert=Printf.sprintf"listen udp:%s"(string_of_addresst.address)letgetsocknameserver=getsockname"Udp.getsockname"server.server_fdletbindaddress=letfd=Lwt_unix.socketLwt_unix.PF_INETLwt_unix.SOCK_DGRAM0intryletsockaddr=sockaddr_of_addressaddressinLwt_unix.bindfdsockaddr>|=fun()->Ok{server_fd=Somefd;address}with|e->errorf"udp:%s: bind caught %s"(string_of_addressaddress)(Printexc.to_stringe)letshutdownt=matcht.server_fdwith|None->Lwt.return_unit|Somefd->t.server_fd<-None;Log.debug(funf->f"%s: close"(string_of_servert));Lwt_unix.closefdletlistentflow_cb=letbuffer=Cstruct.createmax_udp_lengthinletbytes=Bytes.makemax_udp_length'\000'inmatcht.server_fdwith|None->()|Somefd->letrecloop()=Lwt.catch(fun()->(* Lwt on Win32 doesn't support Lwt_bytes.recvfrom *)Lwt_unix.recvfromfdbytes0(Bytes.lengthbytes)[]>>=fun(n,sockaddr)->Cstruct.blit_from_bytesbytes0buffer0n;letdata=Cstruct.subbuffer0nin(* construct a flow with this buffer available for reading *)(matchaddress_of_sockaddrsockaddrwith|Someaddress->Lwt.returnaddress|None->Lwt.fail(Failure"failed to discover incoming socket address"))>>=funaddress->letflow=of_fd~read_buffer_size:0~already_read:(Somedata)sockaddraddressfdinLwt.async(fun()->Lwt.catch(fun()->flow_cbflow)(fune->Log.info(funf->f"%s: listen callback caught: %s"(string_of_servert)(Printexc.to_stringe));Lwt.return_unit));Lwt.returntrue)(fune->Log.err(funf->f"%s: listen caught %s shutting down server"(string_of_servert)(Printexc.to_stringe));Lwt.returnfalse)>>=function|false->Lwt.return_unit|true->loop()inLwt.asyncloopendmoduleTime=structtype'aio='aLwt.tletsleep_nsns=Lwt_unix.sleep(Duration.to_fns)endmoduleClock=MclockmoduleR=structopenDns_forwardmoduleUdp_client=Rpc.Client.Make(Udp)(Framing.Udp(Udp))(Time)moduleUdp=Resolver.Make(Udp_client)(Time)(Clock)moduleTcp_client=Rpc.Client.Make(Tcp)(Framing.Tcp(Tcp))(Time)moduleTcp=Resolver.Make(Tcp_client)(Time)(Clock)endmoduleServer=structopenDns_forwardmoduleUdp_server=Rpc.Server.Make(Udp)(Framing.Udp(Udp))(Time)moduleUdp=Server.Make(Udp_server)(R.Udp)moduleTcp_server=Rpc.Server.Make(Tcp)(Framing.Tcp(Tcp))(Time)moduleTcp=Server.Make(Tcp_server)(R.Tcp)endmoduleResolver=R