diff --git a/lualib/dns.lua b/lualib/dns.lua new file mode 100644 index 00000000..4636bf6d --- /dev/null +++ b/lualib/dns.lua @@ -0,0 +1,290 @@ +--[[ + lua dns resolver library + See https://github.com/xjdrew/levent/blob/master/levent/dns.lua for more detail + +-- resource record type: +-- TYPE value and meaning +-- A 1 a host address +-- NS 2 an authoritative name server +-- MD 3 a mail destination (Obsolete - use MX) +-- MF 4 a mail forwarder (Obsolete - use MX) +-- CNAME 5 the canonical name for an alias +-- SOA 6 marks the start of a zone of authority +-- MB 7 a mailbox domain name (EXPERIMENTAL) +-- MG 8 a mail group member (EXPERIMENTAL) +-- MR 9 a mail rename domain name (EXPERIMENTAL) +-- NULL 10 a null RR (EXPERIMENTAL) +-- WKS 11 a well known service description +-- PTR 12 a domain name pointer +-- HINFO 13 host information +-- MINFO 14 mailbox or mail list information +-- MX 15 mail exchange +-- TXT 16 text strings +-- AAAA 28 a ipv6 host address +-- only appear in the question section: +-- AXFR 252 A request for a transfer of an entire zone +-- MAILB 253 A request for mailbox-related records (MB, MG or MR) +-- MAILA 254 A request for mail agent RRs (Obsolete - see MX) +-- * 255 A request for all records +-- +-- resource recode class: +-- IN 1 the Internet +-- CS 2 the CSNET class (Obsolete - used only for examples in some obsolete RFCs) +-- CH 3 the CHAOS class +-- HS 4 Hesiod [Dyer 87] +-- only appear in the question section: +-- * 255 any class +-- ]] + +--[[ +-- struct header { +-- uint16_t tid # identifier assigned by the program that generates any kind of query. +-- uint16_t flags # flags +-- uint16_t qdcount # the number of entries in the question section. +-- uint16_t ancount # the number of resource records in the answer section. +-- uint16_t nscount # the number of name server resource records in the authority records section. +-- uint16_t arcount # the number of resource records in the additional records section. +-- } +-- +-- request body: +-- struct request { +-- string name +-- uint16_t atype +-- uint16_t class +-- } +-- +-- response body: +-- struct response { +-- string name +-- uint16_t atype +-- uint16_t class +-- uint16_t ttl +-- uint16_t rdlength +-- string rdata +-- } +--]] + +local skynet = require "skynet" +local socket = require "socket" + +local MAX_DOMAIN_LEN = 1024 +local MAX_LABEL_LEN = 63 +local MAX_PACKET_LEN = 2048 +local DNS_HEADER_LEN = 12 + +local QTYPE = { + A = 1, + CNAME = 5, + AAAA = 28, +} + +local QCLASS = { + IN = 1, +} + +local dns = {} + +local function verify_domain_name(name) + if #name > MAX_DOMAIN_LEN then + return false + end + if not name:match("^[%l%d-%.]+$") then + return false + end + for w in name:gmatch("([%w-]+)%.?") do + if #w > MAX_LABEL_LEN then + return false + end + end + return true +end + +local next_tid = 1 +local function gen_tid() + local tid = next_tid + next_tid = next_tid + 1 + return tid +end + +local function pack_header(t) + return string.pack(">HHHHHH", + t.tid, t.flags, t.qdcount, t.ancount or 0, t.nscount or 0, t.arcount or 0) +end + +local function pack_question(name, qtype, qclass) + local labels = {} + for w in name:gmatch("([%w-]+)%.?") do + table.insert(labels, string.pack("s1",w)) + end + table.insert(labels, '\0') + table.insert(labels, string.pack(">HH", qtype, qclass)) + return table.concat(labels) +end + +local function unpack_header(chunk) + local tid, flags, qdcount, ancount, nscount, arcount, left = string.unpack(">HHHHHH", chunk) + return { + tid = tid, + flags = flags, + qdcount = qdcount, + ancount = ancount, + nscount = nscount, + arcount = arcount + }, left +end + +-- unpack a resource name +local function unpack_name(chunk, left) + local t = {} + local jump_pointer + local tag, offset, label + while true do + tag, left = string.unpack("B", chunk, left) + if tag & 0xc0 == 0xc0 then + -- pointer + offset,left = string.unpack(">H", chunk, left - 1) + offset = offset & 0x3fff + if not jump_pointer then + jump_pointer = left + end + -- offset is base 0, need to plus 1 + left = offset + 1 + elseif tag == 0 then + break + else + label, left = string.unpack("s1", chunk, left - 1) + t[#t+1] = label + end + end + return table.concat(t, "."), jump_pointer or left +end + +local function unpack_question(chunk, left) + local name, left = unpack_name(chunk, left) + local atype, class, left = string.unpack(">HH", chunk, left) + return { + name = name, + atype = atype, + class = class + }, left +end + +local function unpack_answer(chunk, left) + local name, left = unpack_name(chunk, left) + local atype, class, ttl, rdata, left = string.unpack(">HHI4s2", chunk, left) + return { + name = name, + atype = atype, + class = class, + ttl = ttl, + rdata = rdata + },left +end + +local function unpack_rdata(qtype, chunk) + if qtype == QTYPE.A then + local a,b,c,d = string.unpack("BBBB", chunk) + return string.format("%d.%d.%d.%d", a,b,c,d) + elseif qtype == QTYPE.AAAA then + local a,b,c,d,e,f,g,h = string.unpack(">HHHHHHHH", chunk) + return string.format("%x:%x:%x:%x:%x:%x:%x:%x", a, b, c, d, e, f, g, h) + else + error("Error qtype " .. qtype) + end +end + +local dns_server +local request_pool = {} + +local function resolve(content) + if #content < DNS_HEADER_LEN then + -- drop + skynet.error("Recv an invalid package when dns query") + return + end + local answer_header,left = unpack_header(content) + -- verify answer + assert(answer_header.qdcount == 1, "malformed packet") + + local resp = request_pool[answer_header.tid] + if not resp then + skynet.error("Recv an invalid tid when dns query") + return + end + + local question,left = unpack_question(content, left) + if question.name ~= resp.name then + skynet.error("Recv an invalid name when dns query") + return + end + + local ttl + local answer + local answers = {} + for i=1, answer_header.ancount do + answer, left = unpack_answer(content, left) + -- only extract qtype address + if answer.atype == resp.qtype then + local ip = unpack_rdata(resp.qtype, answer.rdata) + ttl = ttl and math.min(ttl, answer.ttl) or answer.ttl + answers[#answers+1] = ip + end + end + + if #answers > 0 then + resp.answers = answers + end + + skynet.wakeup(resp.co) +end + +function dns.server(server, port) + if not server then + local f = assert(io.open "/etc/resolv.conf") + for line in f:lines() do + server = line:match("%s*nameserver%s+([^%s]+)") + if server then + break + end + end + assert(server, "Can't get nameserver") + end + dns_server = socket.udp(function(data, sz, from) + resolve(skynet.tostring(data,sz)) + end) + socket.udp_connect(dns_server, server, port or 53) + return server +end + +local function suspend(tid, name, qtype) + local req = { + name = name, + tid = tid, + qtype = qtype, + time = skynet.now(), -- for timeout + co = coroutine.running(), + } + request_pool[tid] = req + skynet.wait() + local answers = request_pool[tid].answers + request_pool[tid] = nil + assert(answers, "no ip") + return answers[1], answers +end + +function dns.resolve(name, ipv6) + local qtype = ipv6 and QTYPE.AAAA or QTYPE.A + local name = name:lower() + assert(verify_domain_name(name) , "illegal name") + local question_header = { + tid = gen_tid(), + flags = 0x100, -- flags: 00000001 00000000, set RD + qdcount = 1, + } + local req = pack_header(question_header) .. pack_question(name, qtype, QCLASS.IN) + assert(dns_server, "Call dns.server fist") + socket.write(dns_server, req) + return suspend(question_header.tid, name, qtype) +end + +return dns diff --git a/test/testdns.lua b/test/testdns.lua new file mode 100644 index 00000000..1d767ce0 --- /dev/null +++ b/test/testdns.lua @@ -0,0 +1,11 @@ +local skynet = require "skynet" +local dns = require "dns" + +skynet.start(function() + print("nameserver:", dns.server()) -- set nameserver + -- you can specify the server like dns.server("8.8.4.4", 53) + local ip, ips = dns.resolve "github.com" + for k,v in ipairs(ips) do + print("github.com",v) + end +end)